diff --git a/src/RouterProvider.php b/src/RouterProvider.php index 3d8717c..e5e49de 100644 --- a/src/RouterProvider.php +++ b/src/RouterProvider.php @@ -2,7 +2,7 @@ declare(strict_types=1); -namespace NeuronAI\Router; +namespace NeuronAI; use Closure; use Generator; @@ -15,17 +15,21 @@ use NeuronAI\Providers\AIProviderInterface; use NeuronAI\Providers\MessageMapperInterface; use NeuronAI\Providers\ToolMapperInterface; -use NeuronAI\Router\Rules\RoutingRuleInterface; +use NeuronAI\Rules\RoundRobinRule; +use NeuronAI\Rules\RoutingRuleInterface; use NeuronAI\StaticConstructor; use NeuronAI\Tools\ToolInterface; -use function implode; -use function is_array; -use function array_keys; -use function in_array; use function array_key_last; +use function array_keys; use function array_values; +use function count; +use function implode; +use function in_array; +use function is_array; +// Added a built-in RouterProvider so this fork can route requests across +// multiple providers, with optional weighted load balancing on RoundRobinRule. class RouterProvider implements AIProviderInterface { use StaticConstructor; @@ -58,14 +62,38 @@ class RouterProvider implements AIProviderInterface */ protected array $tools = []; - public function addProvider(string $name, AIProviderInterface $provider): self - { + /** + * @var array + */ + protected array $providerRoutingConfig = []; + + public function addProvider( + string $name, + AIProviderInterface $provider, + bool $useLoadBalancing = false, + ?int $utilizationRate = null, + ): self { + if ($useLoadBalancing && ($utilizationRate === null || $utilizationRate <= 0)) { + throw new ProviderException( + "RouterProvider: provider '{$name}' has load balancing enabled but utilizationRate is missing or invalid.", + ); + } + $this->providers[$name] = $provider; + $this->providerRoutingConfig[$name] = [ + 'useLoadBalancing' => $useLoadBalancing, + 'utilizationRate' => $utilizationRate, + ]; + return $this; } public function setRule(RoutingRuleInterface $rule): self { + if ($rule instanceof RoundRobinRule) { + $this->configureRoundRobinLoadBalancing($rule); + } + $this->rule = $rule; return $this; } @@ -141,7 +169,7 @@ public function chat(Message ...$messages): Message return $this->withFallback( 'chat', $messages, - fn (AIProviderInterface $provider): \NeuronAI\Chat\Messages\Message => $provider->chat(...$messages), + fn (AIProviderInterface $provider): Message => $provider->chat(...$messages), ); } @@ -213,7 +241,7 @@ public function structured(array|Message $messages, string $class, array $respon return $this->withFallback( 'structured', is_array($messages) ? $messages : [$messages], - fn (AIProviderInterface $provider): \NeuronAI\Chat\Messages\Message => $provider->structured($messages, $class, $response_schema), + fn (AIProviderInterface $provider): Message => $provider->structured($messages, $class, $response_schema), ); } @@ -222,7 +250,7 @@ public function structured(array|Message $messages, string $class, array $respon */ public function messageMapper(): MessageMapperInterface { - if (!$this->resolvedProvider instanceof \NeuronAI\Providers\AIProviderInterface) { + if (!$this->resolvedProvider instanceof AIProviderInterface) { throw new ProviderException( 'RouterProvider: no provider available for delegation. Call setDefaultProvider() or make an inference call first.', ); @@ -235,7 +263,7 @@ public function messageMapper(): MessageMapperInterface */ public function toolPayloadMapper(): ToolMapperInterface { - if (!$this->resolvedProvider instanceof \NeuronAI\Providers\AIProviderInterface) { + if (!$this->resolvedProvider instanceof AIProviderInterface) { throw new ProviderException( 'RouterProvider: no provider available for delegation. Call setDefaultProvider() or make an inference call first.', ); @@ -378,4 +406,54 @@ protected function defaultFallbackStrategy(Throwable $e): bool return $status === null || $status === 429 || $status >= 500; } + + /** + * Reads per-provider routing flags and, when present for all providers in the + * round-robin rule, upgrades it to weighted load balancing. + * + * @throws ProviderException + */ + protected function configureRoundRobinLoadBalancing(RoundRobinRule $rule): void + { + $providers = $rule->getProviders(); + $weights = []; + + foreach ($providers as $name) { + if (!isset($this->providers[$name])) { + throw new ProviderException( + "RouterProvider: unknown provider '{$name}' in RoundRobinRule. Available: " . implode(', ', array_keys($this->providers)), + ); + } + + $config = $this->providerRoutingConfig[$name] ?? [ + 'useLoadBalancing' => false, + 'utilizationRate' => null, + ]; + + if ($config['useLoadBalancing']) { + if ($config['utilizationRate'] === null || $config['utilizationRate'] <= 0) { + throw new ProviderException( + "RouterProvider: provider '{$name}' has invalid utilizationRate for load balancing.", + ); + } + + $weights[$name] = $config['utilizationRate']; + } + } + + if ($weights === []) { + $rule->setUseLoadBalancing(false); + return; + } + + if (count($weights) !== count($providers)) { + throw new ProviderException( + 'RouterProvider: when enabling load balancing in RoundRobinRule, all listed providers must define useLoadBalancing=true and utilizationRate.', + ); + } + + $rule + ->setUseLoadBalancing(true) + ->setProviderWeights($weights); + } } diff --git a/src/Rules/RoundRobinRule.php b/src/Rules/RoundRobinRule.php index 07d30d1..72d373c 100644 --- a/src/Rules/RoundRobinRule.php +++ b/src/Rules/RoundRobinRule.php @@ -1,30 +1,128 @@ -providers = array_values($providers); - } - - public function resolveProvider(string $method, array $messages, array $tools): string - { - $name = $this->providers[$this->index]; - - $this->index = ($this->index + 1) % count($this->providers); - - return $name; - } -} + */ + protected array $providerWeights = []; + + /** @var string[] */ + protected array $weightedProviders = []; + + /** + * @param string[] $providers + * @param array $providerWeights + */ + public function __construct(array $providers, bool $useLoadBalancing = false, array $providerWeights = []) + { + $this->providers = array_values($providers); + $this->useLoadBalancing = $useLoadBalancing; + + if ($providerWeights !== []) { + $this->setProviderWeights($providerWeights); + } + } + + /** + * @return string[] + */ + public function getProviders(): array + { + return $this->providers; + } + + public function setUseLoadBalancing(bool $enabled): self + { + $this->useLoadBalancing = $enabled; + $this->index = 0; + + if (!$enabled) { + $this->weightedProviders = []; + } + + return $this; + } + + /** + * @param array $providerWeights + */ + public function setProviderWeights(array $providerWeights): self + { + $this->assertValidWeights($providerWeights); + + $weightedProviders = []; + + foreach ($this->providers as $name) { + $weight = $providerWeights[$name]; + $weightedProviders = array_merge($weightedProviders, array_fill(0, $weight, $name)); + } + + $this->providerWeights = $providerWeights; + $this->weightedProviders = $weightedProviders; + $this->index = 0; + + return $this; + } + + public function resolveProvider(string $method, array $messages, array $tools): string + { + $pool = $this->useLoadBalancing && $this->weightedProviders !== [] + ? $this->weightedProviders + : $this->providers; + + if ($pool === []) { + throw new InvalidArgumentException('RoundRobinRule: providers list cannot be empty.'); + } + + $name = $pool[$this->index]; + $this->index = ($this->index + 1) % count($pool); + + return $name; + } + + /** + * @param array $providerWeights + */ + protected function assertValidWeights(array $providerWeights): void + { + if ($providerWeights === []) { + throw new InvalidArgumentException('RoundRobinRule: providerWeights cannot be empty when load balancing is enabled.'); + } + + foreach ($this->providers as $name) { + if (!array_key_exists($name, $providerWeights)) { + throw new InvalidArgumentException("RoundRobinRule: missing weight for provider '{$name}'."); + } + + if (!is_int($providerWeights[$name]) || $providerWeights[$name] <= 0) { + throw new InvalidArgumentException("RoundRobinRule: weight for provider '{$name}' must be a positive integer."); + } + } + + if (array_sum($providerWeights) !== 100) { + throw new InvalidArgumentException('RoundRobinRule: provider weights must sum to 100.'); + } + } +} diff --git a/src/Rules/RoutingRuleInterface.php b/src/Rules/RoutingRuleInterface.php index a216462..49b075b 100644 --- a/src/Rules/RoutingRuleInterface.php +++ b/src/Rules/RoutingRuleInterface.php @@ -1,17 +1,17 @@ -assertSame('a', $rule->resolveProvider('chat', [], [])); + $this->assertSame('b', $rule->resolveProvider('chat', [], [])); + $this->assertSame('a', $rule->resolveProvider('chat', [], [])); + } + + public function test_weighted_round_robin_respects_percentages(): void + { + $rule = new RoundRobinRule( + providers: ['anthropic', 'openai'], + useLoadBalancing: true, + providerWeights: ['anthropic' => 30, 'openai' => 70], + ); + + $counts = [ + 'anthropic' => 0, + 'openai' => 0, + ]; + + for ($i = 0; $i < 100; $i++) { + $counts[$rule->resolveProvider('chat', [], [])]++; + } + + $this->assertSame(30, $counts['anthropic']); + $this->assertSame(70, $counts['openai']); + } + + public function test_weighted_round_robin_throws_when_weights_do_not_sum_to_100(): void + { + $this->expectException(InvalidArgumentException::class); + $this->expectExceptionMessage('must sum to 100'); + + new RoundRobinRule( + providers: ['anthropic', 'openai'], + useLoadBalancing: true, + providerWeights: ['anthropic' => 40, 'openai' => 40], + ); + } +} diff --git a/tests/RouterProviderLoadBalancingTest.php b/tests/RouterProviderLoadBalancingTest.php new file mode 100644 index 0000000..766f609 --- /dev/null +++ b/tests/RouterProviderLoadBalancingTest.php @@ -0,0 +1,66 @@ +fakeProviderWithResponses('anthropic', 100); + $openai = $this->fakeProviderWithResponses('openai', 100); + + $router = RouterProvider::make() + ->addProvider('anthropic', $anthropic, useLoadBalancing: true, utilizationRate: 30) + ->addProvider('openai', $openai, useLoadBalancing: true, utilizationRate: 70) + ->setRule(new RoundRobinRule(['anthropic', 'openai'])); + + foreach (range(1, 100) as $_) { + $router->chat(UserMessage::make('hello')); + } + + $this->assertSame(30, $anthropic->getCallCount()); + $this->assertSame(70, $openai->getCallCount()); + } + + public function test_router_provider_keeps_standard_round_robin_when_not_configured(): void + { + $providerA = $this->fakeProviderWithResponses('a', 20); + $providerB = $this->fakeProviderWithResponses('b', 20); + + $router = RouterProvider::make() + ->addProvider('a', $providerA) + ->addProvider('b', $providerB) + ->setRule(new RoundRobinRule(['a', 'b'])); + + foreach (range(1, 10) as $_) { + $router->chat(UserMessage::make('hi')); + } + + $this->assertSame(5, $providerA->getCallCount()); + $this->assertSame(5, $providerB->getCallCount()); + } + + protected function fakeProviderWithResponses(string $label, int $count): FakeAIProvider + { + $messages = []; + + foreach (range(1, $count) as $index) { + $messages[] = new AssistantMessage("{$label}-{$index}"); + } + + return new FakeAIProvider(...$messages); + } +}