mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
refactor: clean up fresh tech debt from 2026-09-29 (#43830)
* refactor: clean up fresh tech debt from 2026-09-29 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(routing): pin usage-based routing Redis reads through the proxy and SDK Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
04fa760bf2
commit
6bc17f98d7
5 changed files with 482 additions and 38 deletions
|
|
@ -489,10 +489,6 @@ class RequestRedisBatches:
|
|||
def batches(self) -> tuple[RedisBatch, ...]:
|
||||
return tuple(self._batches.values())
|
||||
|
||||
@property
|
||||
def post_call_batches(self) -> tuple[RedisBatch, ...]:
|
||||
return tuple(self._post_call.values())
|
||||
|
||||
|
||||
_active_request_batches: Final[ContextVar[RequestRedisBatches | None]] = ContextVar(
|
||||
"request_redis_batches", default=None
|
||||
|
|
|
|||
|
|
@ -454,18 +454,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
def usage_counter_keys(self, healthy_deployments: list) -> tuple[list[str], list[str]]:
|
||||
"""The `<id>:<model>:tpm:<HH-MM>` and `<id>:<model>:rpm:<HH-MM>` counter keys selection reads."""
|
||||
current_minute: Final = get_utc_datetime().strftime("%H-%M")
|
||||
|
||||
tpm_keys: Final[list[str]] = []
|
||||
rpm_keys: Final[list[str]] = []
|
||||
for m in healthy_deployments:
|
||||
if isinstance(m, dict):
|
||||
id = m.get("model_info", {}).get(
|
||||
"id"
|
||||
) # a deployment should always have an 'id'. this is set in router.py
|
||||
deployment_name = m.get("litellm_params", {}).get("model")
|
||||
tpm_keys.append(f"{id}:{deployment_name}:tpm:{current_minute}")
|
||||
rpm_keys.append(f"{id}:{deployment_name}:rpm:{current_minute}")
|
||||
return tpm_keys, rpm_keys
|
||||
prefixes: Final = tuple(
|
||||
f"{m.get('model_info', {}).get('id')}:{m.get('litellm_params', {}).get('model')}"
|
||||
for m in healthy_deployments
|
||||
if isinstance(m, dict)
|
||||
)
|
||||
return (
|
||||
[f"{prefix}:tpm:{current_minute}" for prefix in prefixes],
|
||||
[f"{prefix}:rpm:{current_minute}" for prefix in prefixes],
|
||||
)
|
||||
|
||||
async def async_get_available_deployments(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from contextlib import contextmanager
|
|||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
|
@ -24,15 +24,9 @@ from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
|
|||
from litellm.router_utils.cooldown_cache import CooldownCache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
from opentelemetry.trace import Span
|
||||
|
||||
from litellm.router import Router as _Router
|
||||
|
||||
LitellmRouter = _Router
|
||||
Span = _Span
|
||||
else:
|
||||
LitellmRouter = Any
|
||||
Span = Any
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
_PREFETCH_SLOT: Final = "routing_read"
|
||||
|
|
@ -89,7 +83,7 @@ class RoutingPrefetch:
|
|||
|
||||
@staticmethod
|
||||
def arm(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
litellm_router_instance: "Router",
|
||||
usage_selector: LowestTPMLoggingHandler_v2 | None,
|
||||
deployments: list,
|
||||
) -> None:
|
||||
|
|
@ -180,9 +174,9 @@ class RoutingReadBatch:
|
|||
|
||||
async def async_get_cooldown_deployments(
|
||||
self,
|
||||
litellm_router_instance: LitellmRouter,
|
||||
litellm_router_instance: "Router",
|
||||
healthy_deployments: list,
|
||||
parent_otel_span: Span | None,
|
||||
parent_otel_span: "Span | None",
|
||||
) -> list[str]:
|
||||
"""
|
||||
`_async_get_cooldown_deployments`, with the strategy's tpm/rpm counters for
|
||||
|
|
@ -190,19 +184,23 @@ class RoutingReadBatch:
|
|||
"""
|
||||
model_ids: Final = litellm_router_instance.get_model_ids()
|
||||
cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
reads: Final[list[tuple[DualCache, list[str]]]] = [ # mutable-ok: the usage read is appended below
|
||||
(litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys)
|
||||
]
|
||||
usage_keys: list[str] = [] # mutable-ok: DualCache batch reads take a list
|
||||
if self.usage_selector is not None:
|
||||
tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments)
|
||||
usage_keys = tpm_keys + rpm_keys
|
||||
reads.append((self.usage_selector.router_cache, usage_keys))
|
||||
selector: Final = self.usage_selector
|
||||
usage_keys: Final = (
|
||||
() if selector is None else tuple(itertools.chain(*selector.usage_counter_keys(healthy_deployments)))
|
||||
)
|
||||
reads: Final = (
|
||||
(litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys),
|
||||
*(
|
||||
()
|
||||
if selector is None
|
||||
else ((selector.router_cache, list(usage_keys)),) # mutable-ok: DualCache batch reads take a list
|
||||
),
|
||||
)
|
||||
results: Final = await self._read_prefetched(reads) or await DualCache.async_batch_get_cache_shared(
|
||||
reads, parent_otel_span=parent_otel_span
|
||||
)
|
||||
cooldown_results: Final = results[0]
|
||||
if self.usage_selector is not None:
|
||||
if selector is not None:
|
||||
usage_values: Final = results[1]
|
||||
self.prefetched_usage = PrefetchedUsage(
|
||||
keys=frozenset(usage_keys),
|
||||
|
|
@ -217,7 +215,7 @@ class RoutingReadBatch:
|
|||
|
||||
@staticmethod
|
||||
async def _read_prefetched(
|
||||
reads: list[tuple[DualCache, list[str]]],
|
||||
reads: Sequence[tuple[DualCache, list[str]]],
|
||||
) -> list[list[object | None] | None] | None:
|
||||
"""Serve the reads from the request's armed `RoutingPrefetch`, backfilling each cache's memory tier as
|
||||
its own batch read would. None when nothing usable was armed or the prefetch failed."""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,262 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import shlex
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.redis_process import owned_redis
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from redis import Redis
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
|
||||
OPENAI_MODEL: Final = "gpt-4o-mini"
|
||||
MASTER_KEY: Final = "sk-integration-usage-routing-redis-reads"
|
||||
API_KEY: Final = "synthetic-usage-routing-key"
|
||||
ENDPOINT_PATHS: Final = MappingProxyType(
|
||||
{
|
||||
"/v1/chat/completions": ("/v1/chat/completions", "/v1/chat/completions"),
|
||||
"/v1/messages": ("/v1/responses", "/v1/responses"),
|
||||
"/v1/responses": ("/v1/responses", "/v1/responses"),
|
||||
}
|
||||
)
|
||||
CHAT_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_usage_routing_redis_reads",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": OPENAI_MODEL,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "redis read contract"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
RESPONSES_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "resp_usage_routing_redis_reads",
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": OPENAI_MODEL,
|
||||
"output": [
|
||||
{
|
||||
"id": "msg_usage_routing_redis_reads",
|
||||
"type": "message",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "redis read contract", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _request_object(body: bytes) -> dict[str, JsonValue]:
|
||||
return JSON_OBJECT.validate_json(body)
|
||||
|
||||
|
||||
def _deployment_list(
|
||||
model_name: str, api_base: str, deployment_ids: tuple[str, str]
|
||||
) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{OPENAI_MODEL}",
|
||||
"api_base": api_base,
|
||||
"api_key": API_KEY,
|
||||
"rpm": 1,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
for deployment_id in deployment_ids
|
||||
]
|
||||
|
||||
|
||||
def _request_payload(endpoint: str, model_name: str, marker: str) -> dict[str, JsonValue]:
|
||||
if endpoint == "/v1/responses":
|
||||
return {"model": model_name, "input": marker, "max_output_tokens": 16, "store": False}
|
||||
return {"model": model_name, "messages": [{"role": "user", "content": marker}], "max_tokens": 16}
|
||||
|
||||
|
||||
def _expected_wire_body(endpoint: str, marker: str) -> dict[str, JsonValue]:
|
||||
if endpoint == "/v1/messages":
|
||||
return {
|
||||
"model": OPENAI_MODEL,
|
||||
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": marker}]}],
|
||||
"include": ["reasoning.encrypted_content"],
|
||||
"max_output_tokens": 16,
|
||||
}
|
||||
if endpoint == "/v1/responses":
|
||||
return {"model": OPENAI_MODEL, "input": marker, "max_output_tokens": 16, "store": False}
|
||||
return {"model": OPENAI_MODEL, "messages": [{"role": "user", "content": marker}], "max_tokens": 16}
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
if request.target == "/v1/models":
|
||||
return Reply(body=json.dumps({"object": "list", "data": [{"id": OPENAI_MODEL, "object": "model"}]}).encode())
|
||||
return Reply(body=RESPONSES_RESPONSE if request.target == "/v1/responses" else CHAT_RESPONSE)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
|
||||
commands: Final = SimpleQueue[str]()
|
||||
started: Final = threading.Event()
|
||||
armed: Final = threading.Event()
|
||||
stopped: Final = threading.Event()
|
||||
ready_marker: Final = f"monitor-ready-{uuid.uuid4().hex}"
|
||||
stop_marker: Final = f"monitor-stop-{uuid.uuid4().hex}"
|
||||
|
||||
def capture() -> None:
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
with client.monitor() as monitor:
|
||||
started.set()
|
||||
stream: Final = iter(monitor.listen())
|
||||
while not stopped.is_set():
|
||||
try:
|
||||
record: Final = MONITOR_COMMAND.validate_python(next(stream))
|
||||
except RedisTimeoutError:
|
||||
continue
|
||||
command: Final = record.get("command")
|
||||
if not isinstance(command, str):
|
||||
continue
|
||||
commands.put(command)
|
||||
if ready_marker in command:
|
||||
armed.set()
|
||||
|
||||
thread: Final = threading.Thread(target=capture, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
assert started.wait(timeout=5), "Redis MONITOR did not start"
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(ready_marker, "ready", ex=1)
|
||||
assert armed.wait(timeout=5), "Redis MONITOR did not capture its readiness command"
|
||||
yield commands
|
||||
finally:
|
||||
stopped.set()
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(stop_marker, "stop", ex=1)
|
||||
thread.join(timeout=5)
|
||||
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
|
||||
|
||||
|
||||
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
|
||||
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
|
||||
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
|
||||
return tuple((line, arguments) for line, arguments in parsed if arguments and arguments[0] == "MGET")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint",
|
||||
("/v1/chat/completions", "/v1/messages", "/v1/responses"),
|
||||
ids=("chat-completions", "messages", "responses"),
|
||||
)
|
||||
def test_proxy_usage_routing_reads_cooldown_tpm_then_rpm_from_redis(
|
||||
endpoint: str, tmp_path: Path
|
||||
) -> None:
|
||||
with owned_redis(tmp_path) as cache, wire_server(_reply) as wire:
|
||||
run_id: Final = uuid.uuid4().hex
|
||||
model_name: Final = f"usage-redis-{run_id}"
|
||||
deployment_ids: Final = (f"dep-a-{run_id[:8]}", f"dep-b-{run_id[:8]}")
|
||||
configuration: Final = JSON_OBJECT.validate_python(
|
||||
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
)
|
||||
config: Final = {
|
||||
**configuration,
|
||||
"model_list": _deployment_list(model_name, f"{wire.url}/v1", deployment_ids),
|
||||
"router_settings": {
|
||||
"routing_strategy": "usage-based-routing-v2",
|
||||
"redis_host": cache.host,
|
||||
"redis_port": cache.port,
|
||||
},
|
||||
}
|
||||
config_path: Final = tmp_path / "usage-routing.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with httpx.Client(base_url=wire.url, timeout=15, trust_env=False) as bootstrap_client:
|
||||
bootstrap: Final = Gateway(bootstrap_client, MASTER_KEY, wire.url)
|
||||
with owned_proxy(
|
||||
bootstrap,
|
||||
tmp_path,
|
||||
{"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)},
|
||||
config=config_path,
|
||||
) as candidate:
|
||||
eventually(
|
||||
lambda: wire.received.qsize(),
|
||||
lambda received: received >= len(deployment_ids),
|
||||
seconds=15,
|
||||
)
|
||||
wire.drain()
|
||||
eventually(
|
||||
lambda: datetime.now(UTC),
|
||||
lambda current: current.second < 40,
|
||||
seconds=65,
|
||||
)
|
||||
minute: Final = datetime.now(UTC).strftime("%H-%M")
|
||||
markers: Final = tuple(f"{run_id}-{index}" for index in range(3))
|
||||
payloads: Final = tuple(_request_payload(endpoint, model_name, marker) for marker in markers)
|
||||
request_headers: Final = (
|
||||
{"anthropic-version": "2023-06-01"} if endpoint == "/v1/messages" else {}
|
||||
)
|
||||
with _capture_redis_commands(cache.host, cache.port) as commands:
|
||||
responses: Final = tuple(
|
||||
candidate.request("POST", endpoint, payload, headers=request_headers) for payload in payloads
|
||||
)
|
||||
assert tuple(response.status_code for response in responses) == (200, 200, 429), [
|
||||
response.text for response in responses
|
||||
]
|
||||
assert "No deployments available" in responses[2].text
|
||||
served_ids: Final = tuple(response.headers["x-litellm-model-id"] for response in responses[:2])
|
||||
assert set(served_ids) == set(deployment_ids), served_ids
|
||||
rpm_keys: Final = tuple(
|
||||
f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}" for deployment_id in deployment_ids
|
||||
)
|
||||
with Redis(host=cache.host, port=cache.port, decode_responses=True) as redis_client:
|
||||
rpm_values: Final = eventually(
|
||||
lambda: tuple(redis_client.get(key) for key in rpm_keys),
|
||||
lambda values: values == ("1", "1"),
|
||||
seconds=15,
|
||||
)
|
||||
assert rpm_values == ("1", "1")
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 2
|
||||
assert tuple(request.method for request in received) == ("POST", "POST")
|
||||
assert tuple(request.target for request in received) == ENDPOINT_PATHS[endpoint]
|
||||
observed_bodies: Final = tuple(_request_object(request.body) for request in received)
|
||||
expected_bodies: Final = tuple(_expected_wire_body(endpoint, marker) for marker in markers[:2])
|
||||
assert observed_bodies == expected_bodies, observed_bodies
|
||||
expected_mget: Final = (
|
||||
"MGET",
|
||||
f"deployment:{deployment_ids[0]}:cooldown",
|
||||
f"deployment:{deployment_ids[1]}:cooldown",
|
||||
f"{deployment_ids[0]}:openai/{OPENAI_MODEL}:tpm:{minute}",
|
||||
f"{deployment_ids[1]}:openai/{OPENAI_MODEL}:tpm:{minute}",
|
||||
*(
|
||||
f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}"
|
||||
for deployment_id in deployment_ids
|
||||
),
|
||||
)
|
||||
mgets: Final = _drain_mgets(commands)
|
||||
assert any(arguments == expected_mget for _, arguments in mgets), mgets
|
||||
raw_mgets: Final = tuple(line for line, _ in mgets)
|
||||
print(f"proxy {endpoint} MGETs: {raw_mgets}")
|
||||
|
|
@ -0,0 +1,191 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import shlex
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from integration._support.client import eventually
|
||||
from integration._support.redis_process import owned_redis
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from litellm import Router
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from redis import Redis
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
MONITOR_COMMAND: Final = TypeAdapter(dict[str, JsonValue])
|
||||
OPENAI_MODEL: Final = "gpt-4o-mini"
|
||||
API_KEY: Final = "synthetic-usage-routing-key"
|
||||
CHAT_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_usage_routing_sdk_redis_reads",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": OPENAI_MODEL,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "redis read contract"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _request_object(body: bytes) -> dict[str, JsonValue]:
|
||||
return JSON_OBJECT.validate_json(body)
|
||||
|
||||
|
||||
def _deployment_list(
|
||||
model_name: str, api_base: str, deployment_ids: tuple[str, str]
|
||||
) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{OPENAI_MODEL}",
|
||||
"api_base": api_base,
|
||||
"api_key": API_KEY,
|
||||
"rpm": 1,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
for deployment_id in deployment_ids
|
||||
]
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
return Reply(body=CHAT_RESPONSE)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _capture_redis_commands(host: str, port: int) -> Iterator[SimpleQueue[str]]:
|
||||
commands: Final = SimpleQueue[str]()
|
||||
started: Final = threading.Event()
|
||||
armed: Final = threading.Event()
|
||||
stopped: Final = threading.Event()
|
||||
ready_marker: Final = f"monitor-ready-{uuid.uuid4().hex}"
|
||||
stop_marker: Final = f"monitor-stop-{uuid.uuid4().hex}"
|
||||
|
||||
def capture() -> None:
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
with client.monitor() as monitor:
|
||||
started.set()
|
||||
stream: Final = iter(monitor.listen())
|
||||
while not stopped.is_set():
|
||||
try:
|
||||
record: Final = MONITOR_COMMAND.validate_python(next(stream))
|
||||
except RedisTimeoutError:
|
||||
continue
|
||||
command: Final = record.get("command")
|
||||
if not isinstance(command, str):
|
||||
continue
|
||||
commands.put(command)
|
||||
if ready_marker in command:
|
||||
armed.set()
|
||||
|
||||
thread: Final = threading.Thread(target=capture, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
assert started.wait(timeout=5), "Redis MONITOR did not start"
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(ready_marker, "ready", ex=1)
|
||||
assert armed.wait(timeout=5), "Redis MONITOR did not capture its readiness command"
|
||||
yield commands
|
||||
finally:
|
||||
stopped.set()
|
||||
with Redis(host=host, port=port, socket_timeout=1, decode_responses=True) as client:
|
||||
client.set(stop_marker, "stop", ex=1)
|
||||
thread.join(timeout=5)
|
||||
assert not thread.is_alive(), "Redis MONITOR thread survived cleanup"
|
||||
|
||||
|
||||
def _drain_mgets(commands: SimpleQueue[str]) -> tuple[tuple[str, tuple[str, ...]], ...]:
|
||||
captured: Final = tuple(commands.get_nowait() for _ in range(commands.qsize()))
|
||||
parsed: Final = tuple((line, tuple(shlex.split(line))) for line in captured)
|
||||
return tuple((line, arguments) for line, arguments in parsed if arguments and arguments[0] == "MGET")
|
||||
|
||||
|
||||
def _model_id(response: object) -> str:
|
||||
response_params: Final = getattr(response, "_hidden_params")
|
||||
hidden_params: Final = JSON_OBJECT.validate_python(response_params)
|
||||
model_id: Final = hidden_params.get("model_id")
|
||||
assert isinstance(model_id, str), hidden_params
|
||||
return model_id
|
||||
|
||||
|
||||
async def _exercise_router(router: Router, model_name: str, markers: tuple[str, str, str]) -> tuple[str, str]:
|
||||
first: Final = await router.acompletion(
|
||||
model=model_name, messages=[{"role": "user", "content": markers[0]}], max_tokens=8
|
||||
)
|
||||
second: Final = await router.acompletion(
|
||||
model=model_name, messages=[{"role": "user", "content": markers[1]}], max_tokens=8
|
||||
)
|
||||
with pytest.raises(litellm.RateLimitError, match="No deployments available"):
|
||||
await router.acompletion(
|
||||
model=model_name, messages=[{"role": "user", "content": markers[2]}], max_tokens=8
|
||||
)
|
||||
return _model_id(first), _model_id(second)
|
||||
|
||||
|
||||
def test_sdk_usage_routing_reads_tpm_then_rpm_from_redis(tmp_path: Path) -> None:
|
||||
with owned_redis(tmp_path) as cache, wire_server(_reply) as wire:
|
||||
run_id: Final = uuid.uuid4().hex
|
||||
model_name: Final = f"usage-redis-{run_id}"
|
||||
deployment_ids: Final = (f"dep-a-{run_id[:8]}", f"dep-b-{run_id[:8]}")
|
||||
router: Final = Router(
|
||||
model_list=_deployment_list(model_name, f"{wire.url}/v1", deployment_ids),
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
redis_host=cache.host,
|
||||
redis_port=cache.port,
|
||||
)
|
||||
try:
|
||||
eventually(
|
||||
lambda: datetime.now(UTC),
|
||||
lambda current: current.second < 40,
|
||||
seconds=65,
|
||||
)
|
||||
minute: Final = datetime.now(UTC).strftime("%H-%M")
|
||||
markers: Final = tuple(f"{run_id}-{index}" for index in range(3))
|
||||
with _capture_redis_commands(cache.host, cache.port) as commands:
|
||||
served_ids: Final = asyncio.run(_exercise_router(router, model_name, markers))
|
||||
assert set(served_ids) == set(deployment_ids), served_ids
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 2
|
||||
assert tuple(request.method for request in received) == ("POST", "POST")
|
||||
assert tuple(request.target for request in received) == ("/v1/chat/completions",) * 2
|
||||
observed_bodies: Final = tuple(_request_object(request.body) for request in received)
|
||||
expected_bodies: Final = tuple(
|
||||
{
|
||||
"model": OPENAI_MODEL,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
"max_tokens": 8,
|
||||
}
|
||||
for marker in markers[:2]
|
||||
)
|
||||
assert observed_bodies == expected_bodies, observed_bodies
|
||||
expected_mget: Final = (
|
||||
"MGET",
|
||||
f"deployment:{deployment_ids[0]}:cooldown",
|
||||
f"deployment:{deployment_ids[1]}:cooldown",
|
||||
*(f"{deployment_id}:openai/{OPENAI_MODEL}:tpm:{minute}" for deployment_id in deployment_ids),
|
||||
*(f"{deployment_id}:openai/{OPENAI_MODEL}:rpm:{minute}" for deployment_id in deployment_ids),
|
||||
)
|
||||
mgets: Final = _drain_mgets(commands)
|
||||
assert any(arguments == expected_mget for _, arguments in mgets), mgets
|
||||
raw_mgets: Final = tuple(line for line, _ in mgets)
|
||||
print(f"sdk MGETs: {raw_mgets}")
|
||||
finally:
|
||||
router.reset()
|
||||
Loading…
Add table
Reference in a new issue