59 lines
2.0 KiB
Python

"""Multi-endpoint load balancing: one logical model, N local ports.
from evalharness.model.pool import PooledAdapter
from evalharness.model.adapter import resolve_adapter
base = resolve_adapter('openai/http://127.0.0.1:8123/v1?Qwen3-8B')
pool = PooledAdapter([resolve_adapter(f'openai/http://127.0.0.1:{p}/v1?Qwen3-8B')
for p in range(8123, 8131)])
out = await pool.generate(...) # round-robin over instances
"""
import itertools
from typing import List, Optional
from ..data.sample import ChatMessage
from .adapter import ModelAdapter
from .output import ModelOutput, Usage
class PooledAdapter(ModelAdapter):
"""Round-robin over N equivalent backend instances."""
name = 'pool'
def __init__(self, adapters: List[ModelAdapter]):
if not adapters:
raise ValueError('PooledAdapter needs at least one backend')
super().__init__(model=adapters[0].model, api_base=adapters[0].api_base)
self.adapters = adapters
self._cycle = itertools.cycle(range(len(adapters)))
self.usage = Usage()
def _next(self) -> ModelAdapter:
return self.adapters[next(self._cycle)]
async def generate(self, messages: List[ChatMessage],
tools: Optional[list] = None, **kw) -> ModelOutput:
last_exc = None
for _ in range(len(self.adapters)): # try each instance once
adapter = self._next()
try:
out = await adapter.generate(messages, tools=tools, **kw)
self.usage = self.usage + out.usage
return out
except Exception as e: # dead instance -> next
last_exc = e
raise last_exc
async def close(self) -> None:
for a in self.adapters:
await a.close()
def pooled(specs: List[str]) -> PooledAdapter:
"""['openai/http://127.0.0.1:8123/v1?M', ...] -> PooledAdapter."""
from .adapter import resolve_adapter
return PooledAdapter([resolve_adapter(s) for s in specs])