Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
102 changes: 90 additions & 12 deletions src/RouterProvider.php
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

declare(strict_types=1);

namespace NeuronAI\Router;
namespace NeuronAI;

use Closure;
use Generator;
Expand All @@ -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;
Expand Down Expand Up @@ -58,14 +62,38 @@ class RouterProvider implements AIProviderInterface
*/
protected array $tools = [];

public function addProvider(string $name, AIProviderInterface $provider): self
{
/**
* @var array<string, array{useLoadBalancing: bool, utilizationRate: int|null}>
*/
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;
}
Expand Down Expand Up @@ -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),
);
}

Expand Down Expand Up @@ -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),
);
}

Expand All @@ -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.',
);
Expand All @@ -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.',
);
Expand Down Expand Up @@ -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);
}
}
158 changes: 128 additions & 30 deletions src/Rules/RoundRobinRule.php
Original file line number Diff line number Diff line change
@@ -1,30 +1,128 @@
<?php

declare(strict_types=1);

namespace NeuronAI\Router\Rules;

use function array_values;
use function count;

class RoundRobinRule implements RoutingRuleInterface
{
/** @var string[] */
protected array $providers;

protected int $index = 0;

public function __construct(array $providers)
{
$this->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;
}
}
<?php

declare(strict_types=1);

namespace NeuronAI\Rules;

use InvalidArgumentException;

use function array_fill;
use function array_key_exists;
use function array_merge;
use function array_sum;
use function array_values;
use function count;
use function is_int;

// Extended round-robin to optionally support weighted load balancing so traffic
// can be distributed by percentage while preserving legacy sequential behavior.
class RoundRobinRule implements RoutingRuleInterface
{
/** @var string[] */
protected array $providers;

protected int $index = 0;

protected bool $useLoadBalancing;

/** @var array<string, int> */
protected array $providerWeights = [];

/** @var string[] */
protected array $weightedProviders = [];

/**
* @param string[] $providers
* @param array<string, int> $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<string, int> $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<string, int> $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.');
}
}
}
34 changes: 17 additions & 17 deletions src/Rules/RoutingRuleInterface.php
Original file line number Diff line number Diff line change
@@ -1,17 +1,17 @@
<?php

declare(strict_types=1);

namespace NeuronAI\Router\Rules;

use NeuronAI\Chat\Messages\Message;
use NeuronAI\Tools\ToolInterface;

interface RoutingRuleInterface
{
/**
* @param Message[] $messages
* @param ToolInterface[] $tools
*/
public function resolveProvider(string $method, array $messages, array $tools): string;
}
<?php
declare(strict_types=1);
namespace NeuronAI\Rules;
use NeuronAI\Chat\Messages\Message;
use NeuronAI\Tools\ToolInterface;
interface RoutingRuleInterface
{
/**
* @param Message[] $messages
* @param ToolInterface[] $tools
*/
public function resolveProvider(string $method, array $messages, array $tools): string;
}
Loading