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:
devin-ai-integration[bot] 2026-09-30 02:40:37 -07:00 • committed by GitHub
parent 04fa760bf2
commit 6bc17f98d7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 482 additions and 38 deletions

View file

@ -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

View file

@ -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,

View file

@ -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."""

View file

@ -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}")

View file

@ -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()