mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_decrease_anys_fable5
# Conflicts: # litellm/proxy/common_request_processing.py # litellm/proxy/pass_through_endpoints/streaming_handler.py
This commit is contained in:
commit
d608948eef
13 changed files with 808 additions and 111 deletions
|
|
@ -7,6 +7,7 @@ import traceback
|
|||
from collections.abc import AsyncGenerator, Callable, Mapping
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload
|
||||
|
||||
import anyio
|
||||
|
|
@ -3114,12 +3115,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if maybe_modified is not None:
|
||||
return maybe_modified
|
||||
elif isinstance(chunk, (bytes, bytearray)):
|
||||
# Decode to str, inject, and rebuild as bytes
|
||||
try:
|
||||
s: Final = chunk.decode("utf-8", errors="ignore")
|
||||
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name)
|
||||
if maybe_mod is not None:
|
||||
return (maybe_mod + ("" if maybe_mod.endswith("\n\n") else "\n\n")).encode("utf-8")
|
||||
s: Final = chunk.decode("utf-8")
|
||||
if s.endswith(("\n\n", "\r\n\r\n")):
|
||||
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name)
|
||||
if maybe_mod is not None:
|
||||
return maybe_mod.encode("utf-8")
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(chunk, str):
|
||||
|
|
@ -3157,17 +3158,85 @@ class ProxyBaseLLMRequestProcessing:
|
|||
obj = json.loads(json_part)
|
||||
maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(obj, model_name)
|
||||
if maybe_modified is not None:
|
||||
# Replace just this line with updated JSON using safe_dumps
|
||||
lines[idx] = f"data: {safe_dumps(maybe_modified)}"
|
||||
lines[idx] = "data: " + safe_dumps(maybe_modified) + ("\r" if ln.endswith("\r") else "")
|
||||
return "\n".join(lines)
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _anthropic_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
prompt_tokens: Final = int(usage.get("input_tokens", 0) or 0)
|
||||
completion_tokens: Final = int(usage.get("output_tokens", 0) or 0)
|
||||
total_tokens: Final = int(
|
||||
usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
|
||||
)
|
||||
web_search_requests: Final = usage.get("web_search_requests")
|
||||
server_tool_use: Final = (
|
||||
ServerToolUse(web_search_requests=web_search_requests) if web_search_requests is not None else None
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("prompt_tokens", prompt_tokens),
|
||||
("completion_tokens", completion_tokens),
|
||||
("total_tokens", total_tokens),
|
||||
("completion_tokens_details", usage.get("completion_tokens_details")),
|
||||
("prompt_tokens_details", usage.get("prompt_tokens_details")),
|
||||
("cache_creation_input_tokens", usage.get("cache_creation_input_tokens")),
|
||||
("cache_read_input_tokens", usage.get("cache_read_input_tokens")),
|
||||
("server_tool_use", server_tool_use),
|
||||
)
|
||||
if value is not None
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _openai_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
prompt_tokens: Final = int(usage.get("prompt_tokens", 0) or 0)
|
||||
completion_tokens: Final = int(usage.get("completion_tokens", 0) or 0)
|
||||
total_tokens: Final = int(
|
||||
usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("prompt_tokens", prompt_tokens),
|
||||
("completion_tokens", completion_tokens),
|
||||
("total_tokens", total_tokens),
|
||||
("completion_tokens_details", usage.get("completion_tokens_details")),
|
||||
("prompt_tokens_details", usage.get("prompt_tokens_details")),
|
||||
)
|
||||
if value is not None
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _stream_usage_kwargs_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Mapping[str, Any] | None:
|
||||
if obj.get("type") == "message_delta":
|
||||
return ProxyBaseLLMRequestProcessing._anthropic_stream_usage_kwargs(usage)
|
||||
if obj.get("object") == "chat.completion.chunk":
|
||||
return ProxyBaseLLMRequestProcessing._openai_stream_usage_kwargs(usage)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _completion_cost_or_none(
|
||||
model_response: ModelResponse, model_name: str, service_tier: str | None
|
||||
) -> float | None:
|
||||
try:
|
||||
return litellm.completion_cost(
|
||||
completion_response=model_response, model=model_name, service_tier=service_tier
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _inject_cost_into_usage_dict(obj: dict, model_name: str) -> dict | None:
|
||||
"""
|
||||
Inject cost information into a usage dictionary for message_delta events.
|
||||
Inject cost information into the usage object of a streamed usage event
|
||||
(Anthropic ``message_delta`` or OpenAI ``chat.completion.chunk``).
|
||||
|
||||
Args:
|
||||
obj: Dictionary containing the SSE event data
|
||||
|
|
@ -3176,57 +3245,21 @@ class ProxyBaseLLMRequestProcessing:
|
|||
Returns:
|
||||
Modified dictionary with cost injected, or None if no modification needed
|
||||
"""
|
||||
if obj.get("type") == "message_delta" and isinstance(obj.get("usage"), dict):
|
||||
_usage: Final = obj["usage"]
|
||||
prompt_tokens: Final = int(_usage.get("input_tokens", 0) or 0)
|
||||
completion_tokens: Final = int(_usage.get("output_tokens", 0) or 0)
|
||||
total_tokens: Final = int(
|
||||
_usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
|
||||
)
|
||||
|
||||
# Extract additional usage fields
|
||||
cache_creation_input_tokens: Final = _usage.get("cache_creation_input_tokens")
|
||||
cache_read_input_tokens: Final = _usage.get("cache_read_input_tokens")
|
||||
web_search_requests: Final = _usage.get("web_search_requests")
|
||||
completion_tokens_details: Final = _usage.get("completion_tokens_details")
|
||||
prompt_tokens_details: Final = _usage.get("prompt_tokens_details")
|
||||
|
||||
usage_kwargs: Final[dict[str, Any]] = {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": total_tokens,
|
||||
}
|
||||
|
||||
# Add optional named parameters
|
||||
if completion_tokens_details is not None:
|
||||
usage_kwargs["completion_tokens_details"] = completion_tokens_details
|
||||
if prompt_tokens_details is not None:
|
||||
usage_kwargs["prompt_tokens_details"] = prompt_tokens_details
|
||||
|
||||
# Handle web_search_requests by wrapping in ServerToolUse
|
||||
if web_search_requests is not None:
|
||||
usage_kwargs["server_tool_use"] = ServerToolUse(web_search_requests=web_search_requests)
|
||||
|
||||
# Add cache-related fields to **params (handled by Usage.__init__)
|
||||
if cache_creation_input_tokens is not None:
|
||||
usage_kwargs["cache_creation_input_tokens"] = cache_creation_input_tokens
|
||||
if cache_read_input_tokens is not None:
|
||||
usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens
|
||||
|
||||
_mr: Final = ModelResponse(usage=Usage(**usage_kwargs))
|
||||
|
||||
try:
|
||||
cost_val = litellm.completion_cost(
|
||||
completion_response=_mr,
|
||||
model=model_name,
|
||||
)
|
||||
except Exception:
|
||||
cost_val = None
|
||||
|
||||
if cost_val is not None:
|
||||
obj.setdefault("usage", {})["cost"] = cost_val
|
||||
return obj
|
||||
return None
|
||||
usage: Final = obj.get("usage")
|
||||
if not isinstance(usage, dict):
|
||||
return None
|
||||
usage_kwargs: Final = ProxyBaseLLMRequestProcessing._stream_usage_kwargs_for_event(obj, usage)
|
||||
if usage_kwargs is None:
|
||||
return None
|
||||
service_tier: Final = obj.get("service_tier")
|
||||
cost_val: Final = ProxyBaseLLMRequestProcessing._completion_cost_or_none(
|
||||
ModelResponse(usage=Usage(**usage_kwargs)),
|
||||
model_name,
|
||||
service_tier if isinstance(service_tier, str) else None,
|
||||
)
|
||||
if cost_val is None:
|
||||
return None
|
||||
return {**obj, "usage": {**usage, "cost": cost_val}}
|
||||
|
||||
def maybe_get_model_id(self, _logging_obj: LiteLLMLoggingObj | None) -> str | None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from collections.abc import Coroutine
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -24,6 +25,21 @@ from .llm_provider_handlers.vertex_passthrough_logging_handler import (
|
|||
from .success_handler import PassThroughEndpointLogging
|
||||
|
||||
|
||||
class RouteStreamingLogging(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
passthrough_success_handler_obj: PassThroughEndpointLogging,
|
||||
url_route: str,
|
||||
request_body: dict,
|
||||
endpoint_type: EndpointType,
|
||||
start_time: datetime,
|
||||
raw_bytes: list[bytes],
|
||||
end_time: datetime,
|
||||
) -> Coroutine[None, None, None]: ...
|
||||
|
||||
|
||||
class PassThroughStreamingHandler:
|
||||
@staticmethod
|
||||
def _stamp_first_chunk_if_needed(litellm_logging_obj: LiteLLMLoggingObj) -> None:
|
||||
|
|
@ -39,7 +55,11 @@ class PassThroughStreamingHandler:
|
|||
start_time: datetime,
|
||||
passthrough_success_handler_obj: PassThroughEndpointLogging,
|
||||
url_route: str,
|
||||
route_streaming_logging: RouteStreamingLogging | None = None,
|
||||
):
|
||||
resolved_route_streaming_logging: Final[RouteStreamingLogging] = (
|
||||
route_streaming_logging or PassThroughStreamingHandler._route_streaming_logging_to_handler
|
||||
)
|
||||
raw_bytes: Final[list[bytes]] = []
|
||||
logging_scheduled = False
|
||||
model_name: Final = PassThroughStreamingHandler._extract_model_for_cost_injection(
|
||||
|
|
@ -56,7 +76,13 @@ class PassThroughStreamingHandler:
|
|||
cost_injection_active: Final = (
|
||||
bool(getattr(litellm, "include_cost_in_streaming_usage", False))
|
||||
and bool(model_name)
|
||||
and endpoint_type in (EndpointType.VERTEX_AI, EndpointType.ANTHROPIC)
|
||||
and (
|
||||
endpoint_type in (EndpointType.ANTHROPIC, EndpointType.OPENAI)
|
||||
or (
|
||||
endpoint_type == EndpointType.VERTEX_AI
|
||||
and ("streamRawPredict" in url_route or "rawPredict" in url_route)
|
||||
)
|
||||
)
|
||||
)
|
||||
try:
|
||||
if not cost_injection_active:
|
||||
|
|
@ -71,20 +97,19 @@ class PassThroughStreamingHandler:
|
|||
# -> ``str`` for the per-chunk call site.
|
||||
assert model_name is not None
|
||||
resolved_model_name: Final[str] = model_name
|
||||
pending = b""
|
||||
async for chunk in response.aiter_bytes():
|
||||
raw_bytes.append(chunk)
|
||||
PassThroughStreamingHandler._stamp_first_chunk_if_needed(litellm_logging_obj)
|
||||
if endpoint_type == EndpointType.VERTEX_AI:
|
||||
if "streamRawPredict" in url_route or "rawPredict" in url_route:
|
||||
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
chunk, resolved_model_name
|
||||
)
|
||||
else: # EndpointType.ANTHROPIC
|
||||
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
chunk, resolved_model_name
|
||||
complete_frames, pending = PassThroughStreamingHandler._split_complete_sse_frames(
|
||||
pending + chunk
|
||||
) # rebind-ok: SSE frame reassembly buffer across transport chunks
|
||||
if complete_frames:
|
||||
yield ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
complete_frames, resolved_model_name
|
||||
)
|
||||
|
||||
yield chunk
|
||||
if pending:
|
||||
yield pending
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error in chunk_processor: %s", e)
|
||||
raise
|
||||
|
|
@ -100,7 +125,7 @@ class PassThroughStreamingHandler:
|
|||
logging_scheduled = True
|
||||
try:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
async_coroutine=resolved_route_streaming_logging(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
|
|
@ -114,6 +139,17 @@ class PassThroughStreamingHandler:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error scheduling chunk_processor logging: %s", e)
|
||||
|
||||
@staticmethod
|
||||
def _split_complete_sse_frames(pending: bytes) -> tuple[bytes, bytes]:
|
||||
lf_boundary_end: Final = pending.rfind(b"\n\n") + 2
|
||||
crlf_boundary_end: Final = pending.rfind(b"\r\n\r\n") + 4
|
||||
boundary_end: Final = max(
|
||||
lf_boundary_end if lf_boundary_end >= 2 else 0, crlf_boundary_end if crlf_boundary_end >= 4 else 0
|
||||
)
|
||||
if boundary_end == 0:
|
||||
return b"", pending
|
||||
return pending[:boundary_end], pending[boundary_end:]
|
||||
|
||||
@staticmethod
|
||||
async def _route_streaming_logging_to_handler(
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
|
|
|
|||
|
|
@ -126,8 +126,8 @@ class AutoRouterBenchmarkGroup(AutoRouterBenchmarkTotals):
|
|||
description="Turns per tier, keyed by the tier name the routing decision recorded at "
|
||||
"request time (never re-derived at read time, since the tier-to-model mapping is "
|
||||
"mutable config). Tier names are scoped to this group's router_type and are not "
|
||||
"comparable across types: a complexity router reports 'simple'/'medium'/'complex'/"
|
||||
"'reasoning', a quality router reports its numeric quality tier, and an adaptive router "
|
||||
"comparable across types: a complexity router reports 'SIMPLE'/'MEDIUM'/'COMPLEX'/"
|
||||
"'REASONING', a quality router reports its numeric quality tier, and an adaptive router "
|
||||
"records no tier at all. Turns no tier served (the classifier fell back to default_model) "
|
||||
"are absent rather than pooled under a sentinel key, so the values may sum to less than turns",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,14 @@
|
|||
"""Regression tests for LIT-2642 — interrupted pass-through streams must still log usage."""
|
||||
"""Regression tests for PassThroughStreamingHandler.chunk_processor."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.proxy.pass_through_endpoints.streaming_handler import (
|
||||
PassThroughStreamingHandler,
|
||||
|
|
@ -361,6 +363,145 @@ async def test_chunk_processor_stamps_completion_start_time_on_cost_injection_pa
|
|||
mock_logging_obj._update_completion_start_time.assert_called_once()
|
||||
|
||||
|
||||
def _openai_passthrough_stream_chunks():
|
||||
return [
|
||||
(
|
||||
b'data: {"id":"chatcmpl-1","object":"chat.completion.chunk",'
|
||||
b'"choices":[{"index":0,"delta":{"content":"Hi"}}],"usage":null}\n\n'
|
||||
),
|
||||
b": keepalive\n\n",
|
||||
(
|
||||
b'data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[],'
|
||||
b'"usage":{"prompt_tokens":11,"completion_tokens":4,"total_tokens":15,'
|
||||
b'"prompt_tokens_details":{"cached_tokens":0,"audio_tokens":0},'
|
||||
b'"completion_tokens_details":{"reasoning_tokens":0,"audio_tokens":0,'
|
||||
b'"accepted_prediction_tokens":0,"rejected_prediction_tokens":0}}}\n\n'
|
||||
),
|
||||
b"data: [DONE]\n\n",
|
||||
]
|
||||
|
||||
|
||||
async def _collect_openai_passthrough_chunks(chunks, endpoint_type):
|
||||
response = _make_streaming_response(chunks)
|
||||
received = []
|
||||
async for chunk in PassThroughStreamingHandler.chunk_processor(
|
||||
response=response,
|
||||
request_body={"model": "gpt-4o-mini", "stream": True},
|
||||
litellm_logging_obj=MagicMock(),
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=datetime.now(),
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
url_route="/openai/v1/chat/completions",
|
||||
route_streaming_logging=AsyncMock(),
|
||||
):
|
||||
received.append(chunk)
|
||||
await asyncio.sleep(0)
|
||||
return received
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_injects_cost_into_openai_passthrough_usage_frame(monkeypatch):
|
||||
"""Regression: issue #36492 — with include_cost_in_streaming_usage on, the final
|
||||
OpenAI passthrough chat.completion.chunk usage frame must carry usage.cost, like
|
||||
every other streaming surface already does."""
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunks = _openai_passthrough_stream_chunks()
|
||||
|
||||
received = await _collect_openai_passthrough_chunks(chunks, EndpointType.OPENAI)
|
||||
|
||||
assert received[0] == chunks[0]
|
||||
assert received[1] == chunks[1]
|
||||
assert received[3] == chunks[3]
|
||||
final_payload = json.loads(received[2].decode("utf-8").split("data:", 1)[1].strip())
|
||||
pricing = litellm.model_cost["gpt-4o-mini"]
|
||||
expected_cost = 11 * pricing["input_cost_per_token"] + 4 * pricing["output_cost_per_token"]
|
||||
assert final_payload["usage"]["cost"] == pytest.approx(expected_cost)
|
||||
assert final_payload["usage"]["cost"] > 0
|
||||
assert final_payload["usage"]["prompt_tokens"] == 11
|
||||
assert final_payload["usage"]["completion_tokens"] == 4
|
||||
assert final_payload["usage"]["total_tokens"] == 15
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_injects_cost_into_usage_frame_fragmented_across_chunks(monkeypatch):
|
||||
"""Regression: an SSE usage frame split across transport chunks must still get
|
||||
cost injected once the frame completes, instead of passing through untouched."""
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
whole = _openai_passthrough_stream_chunks()
|
||||
usage_frame = whole[2]
|
||||
split_at = len(usage_frame) // 2
|
||||
chunks = [whole[0], whole[1], usage_frame[:split_at], usage_frame[split_at:], whole[3]]
|
||||
|
||||
received = await _collect_openai_passthrough_chunks(chunks, EndpointType.OPENAI)
|
||||
|
||||
reassembled = b"".join(received).decode("utf-8")
|
||||
usage_lines = [ln for ln in reassembled.split("\n") if '"total_tokens"' in ln]
|
||||
assert len(usage_lines) == 1
|
||||
final_payload = json.loads(usage_lines[0].split("data:", 1)[1].strip())
|
||||
assert final_payload["usage"]["cost"] > 0
|
||||
assert final_payload["usage"]["prompt_tokens"] == 11
|
||||
assert reassembled.endswith("data: [DONE]\n\n")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_streams_crlf_delimited_frames_live_and_injects_cost(monkeypatch):
|
||||
"""Regression: CRLF-delimited SSE frames must flow as they complete instead of
|
||||
buffering until EOF, and the usage frame must still get cost injected."""
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunks = [chunk.replace(b"\n\n", b"\r\n\r\n") for chunk in _openai_passthrough_stream_chunks()]
|
||||
|
||||
received = await _collect_openai_passthrough_chunks(chunks, EndpointType.OPENAI)
|
||||
|
||||
assert len(received) == len(chunks)
|
||||
assert received[0] == chunks[0]
|
||||
injected_usage_frame = received[2]
|
||||
assert injected_usage_frame.endswith(b"\r\n\r\n")
|
||||
assert b"\n" not in injected_usage_frame.replace(b"\r\n", b"")
|
||||
reassembled = b"".join(received).decode("utf-8")
|
||||
usage_lines = [ln for ln in reassembled.replace("\r\n", "\n").split("\n") if '"total_tokens"' in ln]
|
||||
assert len(usage_lines) == 1
|
||||
final_payload = json.loads(usage_lines[0].split("data:", 1)[1].strip())
|
||||
assert final_payload["usage"]["cost"] > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_flag_off_leaves_openai_passthrough_stream_byte_identical(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False)
|
||||
chunks = _openai_passthrough_stream_chunks()
|
||||
|
||||
received = await _collect_openai_passthrough_chunks(chunks, EndpointType.OPENAI)
|
||||
|
||||
assert received == chunks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_flag_on_leaves_openai_frames_without_usage_untouched(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunks = [
|
||||
(
|
||||
b'data: {"id":"chatcmpl-1","object":"chat.completion.chunk",'
|
||||
b'"choices":[{"index":0,"delta":{"content":"Hi"}}],"usage":null}\n\n'
|
||||
),
|
||||
b": keepalive\n\n",
|
||||
b"not json at all\n\n",
|
||||
b"data: [DONE]\n\n",
|
||||
]
|
||||
|
||||
received = await _collect_openai_passthrough_chunks(chunks, EndpointType.OPENAI)
|
||||
|
||||
assert received == chunks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_processor_flag_on_leaves_generic_passthrough_untouched(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunks = _openai_passthrough_stream_chunks()
|
||||
|
||||
received = await _collect_openai_passthrough_chunks(chunks, EndpointType.GENERIC)
|
||||
|
||||
assert received == chunks
|
||||
|
||||
|
||||
def test_convert_raw_bytes_survives_truncated_multibyte_sequence():
|
||||
"""A stream cut mid-multibyte-sequence (client disconnect) must still decode
|
||||
via errors="replace" so the usage events already received are logged, instead
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import datetime
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import AsyncGenerator, Callable, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -5746,3 +5747,132 @@ class TestPerRequestModelGroupAlias:
|
|||
)
|
||||
|
||||
assert merged_for == ["group-b"]
|
||||
|
||||
|
||||
class TestInjectCostIntoUsageDict:
|
||||
@staticmethod
|
||||
def _expected_cost(model, prompt_tokens, completion_tokens):
|
||||
pricing = litellm.model_cost[model]
|
||||
return prompt_tokens * pricing["input_cost_per_token"] + completion_tokens * pricing["output_cost_per_token"]
|
||||
|
||||
def test_openai_chat_completion_chunk_usage_gets_cost(self):
|
||||
event = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion.chunk",
|
||||
"choices": [],
|
||||
"usage": {
|
||||
"prompt_tokens": 11,
|
||||
"completion_tokens": 4,
|
||||
"total_tokens": 15,
|
||||
"prompt_tokens_details": {"cached_tokens": 0, "audio_tokens": 0},
|
||||
"completion_tokens_details": {
|
||||
"reasoning_tokens": 0,
|
||||
"audio_tokens": 0,
|
||||
"accepted_prediction_tokens": 0,
|
||||
"rejected_prediction_tokens": 0,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini")
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == pytest.approx(self._expected_cost("gpt-4o-mini", 11, 4))
|
||||
assert result["usage"]["cost"] > 0
|
||||
assert result["usage"]["prompt_tokens"] == 11
|
||||
assert result["id"] == "chatcmpl-1"
|
||||
assert "cost" not in event["usage"]
|
||||
|
||||
def test_anthropic_message_delta_usage_still_gets_cost(self):
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "claude-haiku-4-5")
|
||||
|
||||
assert result is not None
|
||||
assert result["usage"]["cost"] == pytest.approx(self._expected_cost("claude-haiku-4-5", 11, 4))
|
||||
assert result["usage"]["cost"] > 0
|
||||
assert result["usage"]["output_tokens"] == 4
|
||||
|
||||
def test_openai_chunk_with_flex_service_tier_uses_flex_pricing(self):
|
||||
event = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion.chunk",
|
||||
"service_tier": "flex",
|
||||
"choices": [],
|
||||
"usage": {"prompt_tokens": 1000, "completion_tokens": 100, "total_tokens": 1100},
|
||||
}
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-5-mini")
|
||||
|
||||
assert result is not None
|
||||
pricing = litellm.model_cost["gpt-5-mini"]
|
||||
expected_flex_cost = 1000 * pricing["input_cost_per_token_flex"] + 100 * pricing["output_cost_per_token_flex"]
|
||||
assert result["usage"]["cost"] == pytest.approx(expected_flex_cost)
|
||||
assert result["usage"]["cost"] < self._expected_cost("gpt-5-mini", 1000, 100)
|
||||
|
||||
def test_openai_chunk_with_null_usage_is_not_modified(self):
|
||||
event = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion.chunk",
|
||||
"choices": [{"index": 0, "delta": {"content": "Hi"}}],
|
||||
"usage": None,
|
||||
}
|
||||
|
||||
assert ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini") is None
|
||||
|
||||
def test_unrecognized_event_shape_with_usage_is_not_modified(self):
|
||||
event = {"kind": "custom", "usage": {"prompt_tokens": 11, "completion_tokens": 4}}
|
||||
|
||||
assert ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(event, "gpt-4o-mini") is None
|
||||
|
||||
def test_sse_frame_with_coalesced_done_line_injects_into_usage_frame(self):
|
||||
frame = (
|
||||
'data: {"object":"chat.completion.chunk","choices":[],'
|
||||
'"usage":{"prompt_tokens":11,"completion_tokens":4,"total_tokens":15}}\n\n'
|
||||
"data: [DONE]\n\n"
|
||||
)
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(frame, "gpt-4o-mini")
|
||||
|
||||
assert result is not None
|
||||
assert "data: [DONE]" in result
|
||||
injected = json.loads(result.split("\n")[0].split("data:", 1)[1].strip())
|
||||
assert injected["usage"]["cost"] == pytest.approx(self._expected_cost("gpt-4o-mini", 11, 4))
|
||||
|
||||
|
||||
class TestProcessChunkWithCostInjection:
|
||||
def test_complete_usage_frame_chunk_is_injected(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunk = (
|
||||
b'data: {"object":"chat.completion.chunk","choices":[],'
|
||||
b'"usage":{"prompt_tokens":11,"completion_tokens":4,"total_tokens":15}}\n\n'
|
||||
)
|
||||
|
||||
result = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, "gpt-4o-mini")
|
||||
|
||||
assert result != chunk
|
||||
assert result.endswith(b"\n\n")
|
||||
payload = json.loads(result.decode("utf-8").split("data:", 1)[1].strip())
|
||||
assert payload["usage"]["cost"] > 0
|
||||
|
||||
def test_chunk_ending_in_partial_frame_passes_through_byte_identical(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunk = (
|
||||
b'data: {"object":"chat.completion.chunk","choices":[],'
|
||||
b'"usage":{"prompt_tokens":11,"completion_tokens":4,"total_tokens":15}}\n\ndata: [DO'
|
||||
)
|
||||
|
||||
assert ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, "gpt-4o-mini") == chunk
|
||||
|
||||
def test_chunk_with_invalid_utf8_passes_through_byte_identical(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
|
||||
chunk = (
|
||||
b'\xa8data: {"object":"chat.completion.chunk","choices":[],'
|
||||
b'"usage":{"prompt_tokens":11,"completion_tokens":4,"total_tokens":15}}\n\n'
|
||||
)
|
||||
|
||||
assert ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, "gpt-4o-mini") == chunk
|
||||
|
|
|
|||
|
|
@ -1,10 +1,15 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { fireEvent, render, screen } from "@testing-library/react";
|
||||
import React from "react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
import type { AutoRouterDeployment } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { ApiError } from "@/lib/http/client";
|
||||
|
||||
vi.mock("./useAutoRouterBenchmarks", () => ({ useAutoRouterBenchmarks: vi.fn() }));
|
||||
vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ useAutoRouters: vi.fn() }));
|
||||
|
||||
import { useAutoRouters } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
|
||||
import AutoRouterBenchmarksTab from "./AutoRouterBenchmarksTab";
|
||||
import type {
|
||||
|
|
@ -16,6 +21,10 @@ import { useAutoRouterBenchmarks } from "./useAutoRouterBenchmarks";
|
|||
|
||||
type HookResult = ReturnType<typeof useAutoRouterBenchmarks>;
|
||||
|
||||
const mockAutoRouters = (deployments: AutoRouterDeployment[] = []) => {
|
||||
vi.mocked(useAutoRouters).mockReturnValue({ data: deployments } as unknown as ReturnType<typeof useAutoRouters>);
|
||||
};
|
||||
|
||||
const mockHook = (result: { data?: AutoRouterBenchmarksResponse; isPending?: boolean; error?: Error }) => {
|
||||
vi.mocked(useAutoRouterBenchmarks).mockReturnValue({
|
||||
data: result.data,
|
||||
|
|
@ -71,9 +80,20 @@ const response = (groups: AutoRouterBenchmarkGroup[], shared: Totals = totals())
|
|||
groups,
|
||||
});
|
||||
|
||||
const renderTab = () => render(<AutoRouterBenchmarksTab accessToken="sk-test" />);
|
||||
const renderTab = () => {
|
||||
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||
return render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<AutoRouterBenchmarksTab accessToken="sk-test" />
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
};
|
||||
|
||||
describe("AutoRouterBenchmarksTab", () => {
|
||||
beforeEach(() => {
|
||||
mockAutoRouters();
|
||||
});
|
||||
|
||||
it("leads with total estimated savings, before the three session-shape metrics", () => {
|
||||
mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) });
|
||||
renderTab();
|
||||
|
|
@ -97,7 +117,7 @@ describe("AutoRouterBenchmarksTab", () => {
|
|||
expect(screen.getByText("-86%")).toBeInTheDocument();
|
||||
expect(screen.getByText("Actual auto-router spend")).toBeInTheDocument();
|
||||
expect(screen.getByText("$359.86")).toBeInTheDocument();
|
||||
expect(screen.getByText("Estimated spend at highest-cost model")).toBeInTheDocument();
|
||||
expect(screen.getByText("Estimated spend at highest-tier model")).toBeInTheDocument();
|
||||
expect(screen.getByText("$2,534.45")).toBeInTheDocument();
|
||||
expect(screen.getByText("32.7")).toBeInTheDocument();
|
||||
expect(screen.getByText("2.1h")).toBeInTheDocument();
|
||||
|
|
@ -108,12 +128,9 @@ describe("AutoRouterBenchmarksTab", () => {
|
|||
mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) });
|
||||
renderTab();
|
||||
|
||||
expect(screen.getByText("Total sessions")).toBeInTheDocument();
|
||||
expect(screen.getByText("94")).toBeInTheDocument();
|
||||
expect(screen.getByText("Total turns")).toBeInTheDocument();
|
||||
expect(screen.getByText("3,073")).toBeInTheDocument();
|
||||
expect(screen.getByText("Avg saved per session")).toBeInTheDocument();
|
||||
expect(screen.getByText("$23.13")).toBeInTheDocument();
|
||||
expect(screen.getByText("across 94 sessions")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows a cost increase as a positive delta rather than a saving", () => {
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
|
||||
import React, { useState } from "react";
|
||||
|
||||
import type { AutoRouterDeployment } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { useAutoRouters } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
|
|
@ -29,6 +31,7 @@ import {
|
|||
type BucketRow,
|
||||
} from "./autoRouterBenchmarks";
|
||||
import { usd } from "./costOptimizationUtils";
|
||||
import TierTurnsChart from "./TierTurnsChart";
|
||||
import { useAutoRouterBenchmarks } from "./useAutoRouterBenchmarks";
|
||||
|
||||
const Message: React.FC<{ children: React.ReactNode }> = ({ children }) => (
|
||||
|
|
@ -51,7 +54,7 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => {
|
|||
const cheaper = stats.saved_spend >= 0;
|
||||
return (
|
||||
<Card className="overflow-hidden py-0">
|
||||
<div className="grid md:grid-cols-[4fr_3fr_5fr]">
|
||||
<div className="grid md:grid-cols-[1fr_1fr]">
|
||||
<div className="flex flex-col justify-center gap-3 p-6">
|
||||
<p className="text-sm text-muted-foreground">Total estimated savings</p>
|
||||
<div className="flex flex-wrap items-center gap-3">
|
||||
|
|
@ -64,38 +67,22 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => {
|
|||
{Math.abs(stats.saved_pct).toFixed(0)}%
|
||||
</Badge>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="flex flex-col justify-center px-6 pb-6 md:py-6">
|
||||
<dl className="divide-y text-sm">
|
||||
<div className="flex items-baseline justify-between gap-6 py-3">
|
||||
<dt className="text-muted-foreground">Actual auto-router spend</dt>
|
||||
<dd className="font-medium tabular-nums text-foreground">{usd(stats.spend)}</dd>
|
||||
</div>
|
||||
<div className="flex items-baseline justify-between gap-6 py-3">
|
||||
<dt className="text-muted-foreground">Estimated spend at highest-cost model</dt>
|
||||
<dt className="text-muted-foreground">Estimated spend at highest-tier model</dt>
|
||||
<dd className="font-medium tabular-nums text-foreground">{usd(stats.baseline_spend)}</dd>
|
||||
</div>
|
||||
</dl>
|
||||
</div>
|
||||
|
||||
<div className="flex flex-col border-t md:border-t-0 md:border-l">
|
||||
<div className="grid flex-1 grid-cols-2 divide-x">
|
||||
<div className="flex flex-col justify-center gap-1 px-6 py-4">
|
||||
<p className="text-[11px] uppercase tracking-wide text-muted-foreground">Total sessions</p>
|
||||
<p className="text-3xl font-semibold text-foreground">{stats.sessions.toLocaleString()}</p>
|
||||
</div>
|
||||
<div className="flex flex-col justify-center gap-1 px-6 py-4">
|
||||
<p className="text-[11px] uppercase tracking-wide text-muted-foreground">Total turns</p>
|
||||
<p className="text-3xl font-semibold text-foreground">{stats.turns.toLocaleString()}</p>
|
||||
</div>
|
||||
</div>
|
||||
<dl className="flex flex-col divide-y border-t text-sm">
|
||||
<div className="flex items-center justify-between gap-2 px-6 py-3">
|
||||
<dt className="text-[11px] uppercase tracking-wide text-muted-foreground">Avg saved per session</dt>
|
||||
<dd className="text-lg font-semibold tabular-nums text-foreground">{usd(stats.saved_per_session)}</dd>
|
||||
</div>
|
||||
</dl>
|
||||
<div className="flex flex-col items-center justify-center gap-2 border-t p-6 md:border-t-0 md:border-l">
|
||||
<p className="text-[11px] uppercase tracking-wide text-muted-foreground">Avg saved per session</p>
|
||||
<p className="text-5xl font-semibold tracking-tight text-foreground">{usd(stats.saved_per_session)}</p>
|
||||
<p className="text-sm text-muted-foreground">across {stats.sessions.toLocaleString()} sessions</p>
|
||||
</div>
|
||||
</div>
|
||||
</Card>
|
||||
|
|
@ -233,9 +220,10 @@ interface BenchmarksBodyProps {
|
|||
error: unknown;
|
||||
data: AutoRouterBenchmarksResponse | undefined;
|
||||
selectedKey: string;
|
||||
autoRouters: readonly AutoRouterDeployment[];
|
||||
}
|
||||
|
||||
const BenchmarksBody: React.FC<BenchmarksBodyProps> = ({ isPending, error, data, selectedKey }) => {
|
||||
const BenchmarksBody: React.FC<BenchmarksBodyProps> = ({ isPending, error, data, selectedKey, autoRouters }) => {
|
||||
if (isPending) return <Message>Loading auto-router usage...</Message>;
|
||||
if (error instanceof ApiError && error.status === 403) {
|
||||
return <Message>Auto-router usage is visible to proxy admin roles only</Message>;
|
||||
|
|
@ -249,6 +237,8 @@ const BenchmarksBody: React.FC<BenchmarksBodyProps> = ({ isPending, error, data,
|
|||
<>
|
||||
<HeroCard view={view} />
|
||||
|
||||
<TierTurnsChart view={view} autoRouters={autoRouters} />
|
||||
|
||||
<div className="grid grid-cols-1 gap-4 sm:grid-cols-3">
|
||||
<Metric label="Avg turns per session" value={stats.avg_turns_per_session.toFixed(1)} />
|
||||
<Metric label="Avg session length" value={durationLabel(stats.avg_session_seconds)} />
|
||||
|
|
@ -282,6 +272,7 @@ const AutoRouterBenchmarksTab: React.FC<AutoRouterBenchmarksTabProps> = ({ acces
|
|||
const [range, setRange] = useState<BenchmarkWindow>("30d");
|
||||
const { data, isPending, error } = useAutoRouterBenchmarks(accessToken, range);
|
||||
const [selectedKey, setSelectedKey] = useState<string>(ALL_ROUTERS);
|
||||
const { data: autoRouters } = useAutoRouters();
|
||||
|
||||
const groups = data?.groups ?? [];
|
||||
const selectedLabel = data ? viewFor(data, selectedKey).label : "All auto-routers";
|
||||
|
|
@ -319,7 +310,13 @@ const AutoRouterBenchmarksTab: React.FC<AutoRouterBenchmarksTabProps> = ({ acces
|
|||
</div>
|
||||
</div>
|
||||
|
||||
<BenchmarksBody isPending={isPending} error={error} data={data} selectedKey={selectedKey} />
|
||||
<BenchmarksBody
|
||||
isPending={isPending}
|
||||
error={error}
|
||||
data={data}
|
||||
selectedKey={selectedKey}
|
||||
autoRouters={autoRouters ?? []}
|
||||
/>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,12 +1,23 @@
|
|||
import React from "react";
|
||||
import { fireEvent, render, waitFor } from "@testing-library/react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
|
||||
const mockUserDailyActivityCall = vi.fn();
|
||||
const { useAuthorizedMock, mockToolSpendResponse } = vi.hoisted(() => ({
|
||||
useAuthorizedMock: vi.fn(),
|
||||
mockToolSpendResponse: { by_tool: [], daily: [], start_date: null, end_date: null },
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: useAuthorizedMock,
|
||||
}));
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
userDailyActivityCall: (...args: unknown[]) => mockUserDailyActivityCall(...args),
|
||||
getToolSpend: vi.fn().mockResolvedValue({ by_tool: [], daily: [], start_date: null, end_date: null }),
|
||||
getToolSpend: vi.fn().mockResolvedValue(mockToolSpendResponse),
|
||||
getGeneralSettingsCall: vi.fn().mockResolvedValue([]),
|
||||
organizationListCall: vi.fn().mockResolvedValue([]),
|
||||
}));
|
||||
|
||||
vi.mock("@/components/shared/advanced_date_picker", () => ({
|
||||
|
|
@ -38,9 +49,13 @@ const singlePage = {
|
|||
describe("CostOptimizationView daily activity", () => {
|
||||
it("fetches daily activity once for the page and shares it with every tab that needs it", async () => {
|
||||
mockUserDailyActivityCall.mockResolvedValue(singlePage);
|
||||
useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" });
|
||||
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||
|
||||
const { getByRole, getByTestId } = render(
|
||||
<CostOptimizationView accessToken="test-token" userId="u1" userRole="proxy_admin" />,
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<CostOptimizationView accessToken="test-token" userId="u1" userRole="proxy_admin" />
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
await waitFor(() => expect(mockUserDailyActivityCall).toHaveBeenCalledTimes(1));
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
import React from "react";
|
||||
import { fireEvent, render } from "@testing-library/react";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
|
||||
const { useAuthorizedMock } = vi.hoisted(() => ({ useAuthorizedMock: vi.fn() }));
|
||||
|
||||
|
|
@ -7,6 +9,13 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
|||
default: useAuthorizedMock,
|
||||
}));
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
organizationListCall: vi.fn().mockResolvedValue([]),
|
||||
userDailyActivityCall: vi
|
||||
.fn()
|
||||
.mockResolvedValue({ results: [], metadata: { total_pages: 1, has_more: false, page: 1 } }),
|
||||
}));
|
||||
|
||||
vi.mock("./UsageTab", () => ({ __esModule: true, default: () => <div data-testid="usage-tab" /> }));
|
||||
vi.mock("./PromptCompressionTab", () => ({ __esModule: true, default: () => <div data-testid="compression-tab" /> }));
|
||||
vi.mock("./PromptCachingTab", () => ({ __esModule: true, default: () => <div data-testid="caching-tab" /> }));
|
||||
|
|
@ -19,7 +28,12 @@ import CostOptimizationView from "./CostOptimizationView";
|
|||
|
||||
const renderView = (userRole = "Admin") => {
|
||||
useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole });
|
||||
return render(<CostOptimizationView accessToken="test-token" userId="u1" userRole={userRole} />);
|
||||
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||
return render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<CostOptimizationView accessToken="test-token" userId="u1" userRole={userRole} />
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
};
|
||||
|
||||
describe("CostOptimizationView", () => {
|
||||
|
|
|
|||
|
|
@ -0,0 +1,163 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import React from "react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
import type { AutoRouterDeployment } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
|
||||
vi.mock("@/components/shared/charts", () => ({
|
||||
DonutChart: ({ label }: { label: string }) => <div data-testid="donut">{label}</div>,
|
||||
SEQUENTIAL_COLOR_RAMP: ["indigo", "blue"],
|
||||
chartColorValue: (color: string) => color,
|
||||
}));
|
||||
|
||||
import TierTurnsChart, { tierDisplayLabel } from "./TierTurnsChart";
|
||||
import type { AutoRouterBenchmarkGroup, BenchmarkView } from "./autoRouterBenchmarks";
|
||||
|
||||
const totalsOnly = {
|
||||
sessions: 3,
|
||||
turns: 9,
|
||||
avg_turns_per_session: 3,
|
||||
avg_session_seconds: 60,
|
||||
avg_tokens_per_session: 100,
|
||||
spend: 1,
|
||||
saved_spend: 1,
|
||||
baseline_spend: 2,
|
||||
saved_pct: 50,
|
||||
saved_per_session: 0.33,
|
||||
cache: {
|
||||
coverage_pct: 0,
|
||||
hit_rate_pct: 0,
|
||||
same_model: { turns: 0, hits: 0, hit_rate_pct: 0 },
|
||||
first_visit: { turns: 0, hits: 0, hit_rate_pct: 0 },
|
||||
return_to_tier: { turns: 0, hits: 0, hit_rate_pct: 0 },
|
||||
unordered_turns: 0,
|
||||
return_misses_expired: 0,
|
||||
return_misses_within_ttl: 0,
|
||||
return_misses_unknown: 0,
|
||||
ttl_5m_turns: 0,
|
||||
ttl_1h_turns: 0,
|
||||
},
|
||||
};
|
||||
|
||||
const groupView = (overrides: Partial<AutoRouterBenchmarkGroup> = {}): BenchmarkView => ({
|
||||
label: "claude-auto",
|
||||
stats: {
|
||||
...totalsOnly,
|
||||
router_name: "claude-auto",
|
||||
router_type: "complexity",
|
||||
tier_turns: { SIMPLE: 3, COMPLEX: 1 },
|
||||
...overrides,
|
||||
} as AutoRouterBenchmarkGroup,
|
||||
});
|
||||
|
||||
const deployment = (config: unknown): AutoRouterDeployment => ({
|
||||
model_name: "claude-auto",
|
||||
litellm_params: { model: "auto_router/claude-auto", complexity_router_config: config },
|
||||
});
|
||||
|
||||
describe("tierDisplayLabel", () => {
|
||||
it("prefers the admin's custom label for a canonical complexity tier", () => {
|
||||
expect(tierDisplayLabel("SIMPLE", { SIMPLE: "Cheap" })).toBe("Cheap");
|
||||
});
|
||||
|
||||
it("falls back to the canonical name when that tier has no custom label", () => {
|
||||
expect(tierDisplayLabel("COMPLEX", { SIMPLE: "Cheap" })).toBe("Complex");
|
||||
expect(tierDisplayLabel("REASONING", undefined)).toBe("Reasoning");
|
||||
});
|
||||
|
||||
it("shows a non-complexity tier verbatim, since no label map covers a quality router's tier", () => {
|
||||
expect(tierDisplayLabel("3", { SIMPLE: "Cheap" })).toBe("3");
|
||||
});
|
||||
});
|
||||
|
||||
describe("TierTurnsChart", () => {
|
||||
it("labels each slice with its tier and share of the tiered turns", () => {
|
||||
render(<TierTurnsChart view={groupView()} autoRouters={[deployment({ tier_labels: { SIMPLE: "Cheap" } })]} />);
|
||||
|
||||
expect(screen.getByText("Cheap 75%")).toBeInTheDocument();
|
||||
expect(screen.getByText("Complex 25%")).toBeInTheDocument();
|
||||
expect(screen.getByTestId("donut")).toHaveTextContent("4 total turns");
|
||||
});
|
||||
|
||||
it("reads tier_labels out of a config stored as a JSON string", () => {
|
||||
const stored = JSON.stringify({ tier_labels: { SIMPLE: "Cheap" } });
|
||||
render(<TierTurnsChart view={groupView()} autoRouters={[deployment(stored)]} />);
|
||||
|
||||
expect(screen.getByText("Cheap 75%")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("uses canonical names when the router is not in the deployment list", () => {
|
||||
render(<TierTurnsChart view={groupView()} autoRouters={[]} />);
|
||||
|
||||
expect(screen.getByText("Simple 75%")).toBeInTheDocument();
|
||||
expect(screen.getByText("Complex 25%")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("lists each tier's assigned models below its name and share", () => {
|
||||
render(
|
||||
<TierTurnsChart
|
||||
view={groupView()}
|
||||
autoRouters={[deployment({ tiers: { SIMPLE: ["gpt-4o-mini"], COMPLEX: ["gpt-4o", "claude-3-opus"] } })]}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("gpt-4o-mini")).toBeInTheDocument();
|
||||
expect(screen.getByText("gpt-4o, claude-3-opus")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("widens a bare string tier (pinned single model) into its one-model list", () => {
|
||||
render(<TierTurnsChart view={groupView()} autoRouters={[deployment({ tiers: { SIMPLE: "gpt-4o-mini" } })]} />);
|
||||
|
||||
expect(screen.getByText("gpt-4o-mini")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("omits the model line for a tier with no configured models", () => {
|
||||
render(<TierTurnsChart view={groupView()} autoRouters={[deployment({ tiers: { SIMPLE: [] } })]} />);
|
||||
|
||||
expect(screen.getByText("Simple 75%")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows no models for a quality router's numeric tier, which has no per-tier model list", () => {
|
||||
render(
|
||||
<TierTurnsChart
|
||||
view={groupView({ router_type: "quality", tier_turns: { "3": 3, "1": 1 } })}
|
||||
autoRouters={[deployment({ quality_router_config: { available_models: ["gpt-4o"] } })]}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("3 75%")).toBeInTheDocument();
|
||||
expect(screen.getByText("1 25%")).toBeInTheDocument();
|
||||
expect(screen.queryByText("gpt-4o")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("ignores a same-named deployment of a different router type", () => {
|
||||
const qualityDeployment = {
|
||||
model_name: "claude-auto",
|
||||
litellm_params: { model: "auto_router/claude-auto", quality_router_config: { available_models: ["gpt-4o"] } },
|
||||
};
|
||||
|
||||
render(
|
||||
<TierTurnsChart
|
||||
view={groupView()} // complexity router
|
||||
autoRouters={[qualityDeployment] as AutoRouterDeployment[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("Simple 75%")).toBeInTheDocument();
|
||||
expect(screen.queryByText("gpt-4o")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders nothing for the all-routers view, which carries no router identity", () => {
|
||||
const { container } = render(
|
||||
<TierTurnsChart view={{ label: "All auto-routers", stats: totalsOnly }} autoRouters={[]} />,
|
||||
);
|
||||
|
||||
expect(container).toBeEmptyDOMElement();
|
||||
});
|
||||
|
||||
it("renders nothing when the router recorded no tiers", () => {
|
||||
const { container } = render(<TierTurnsChart view={groupView({ tier_turns: {} })} autoRouters={[]} />);
|
||||
|
||||
expect(container).toBeEmptyDOMElement();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,148 @@
|
|||
"use client";
|
||||
|
||||
import React from "react";
|
||||
|
||||
import type { AutoRouterDeployment } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import { hydrateTierLabels } from "@/components/add_model/build_complexity_router_config";
|
||||
import {
|
||||
TIER_KEYS,
|
||||
effectiveTierLabel,
|
||||
type ComplexityTierLabels,
|
||||
type ComplexityTiers,
|
||||
} from "@/components/add_model/ComplexityRouterConfig";
|
||||
import { normalizeTierModels } from "@/components/add_model/complexity_router_tiers";
|
||||
import { chartColorValue, DonutChart, type ChartColor } from "@/components/shared/charts";
|
||||
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
|
||||
import { viewGroup, type BenchmarkView } from "./autoRouterBenchmarks";
|
||||
|
||||
const safeParse = (value: string): unknown => {
|
||||
try {
|
||||
return JSON.parse(value);
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
};
|
||||
|
||||
const asRecord = (value: unknown): Record<string, unknown> => {
|
||||
const parsed: unknown = typeof value === "string" ? safeParse(value) : value;
|
||||
return typeof parsed === "object" && parsed !== null && !Array.isArray(parsed)
|
||||
? (parsed as Record<string, unknown>)
|
||||
: {};
|
||||
};
|
||||
|
||||
const isComplexityTier = (tier: string): tier is keyof ComplexityTiers =>
|
||||
(TIER_KEYS as readonly string[]).includes(tier);
|
||||
|
||||
export const tierDisplayLabel = (tier: string, tierLabels: ComplexityTierLabels | undefined): string =>
|
||||
isComplexityTier(tier) ? effectiveTierLabel(tier, tierLabels) : tier;
|
||||
|
||||
const CONFIG_KEY_BY_ROUTER_TYPE: Record<string, keyof NonNullable<AutoRouterDeployment["litellm_params"]>> = {
|
||||
complexity: "complexity_router_config",
|
||||
quality: "quality_router_config",
|
||||
auto_router: "auto_router_config",
|
||||
adaptive: "adaptive_router_config",
|
||||
};
|
||||
|
||||
const deploymentFor = (
|
||||
routerName: string,
|
||||
routerType: string,
|
||||
autoRouters: readonly AutoRouterDeployment[],
|
||||
): AutoRouterDeployment | undefined => {
|
||||
const configKey = CONFIG_KEY_BY_ROUTER_TYPE[routerType];
|
||||
if (!configKey) return undefined;
|
||||
return autoRouters.find((d) => d.model_name === routerName && d.litellm_params?.[configKey]);
|
||||
};
|
||||
|
||||
const tierLabelsFor = (
|
||||
routerName: string,
|
||||
routerType: string,
|
||||
autoRouters: readonly AutoRouterDeployment[],
|
||||
): ComplexityTierLabels | undefined => {
|
||||
const deployment = deploymentFor(routerName, routerType, autoRouters);
|
||||
if (!deployment) return undefined;
|
||||
const config = asRecord(deployment.litellm_params?.complexity_router_config);
|
||||
return hydrateTierLabels(config.tier_labels);
|
||||
};
|
||||
|
||||
const tierModelsFor = (
|
||||
tier: string,
|
||||
routerName: string,
|
||||
routerType: string,
|
||||
autoRouters: readonly AutoRouterDeployment[],
|
||||
): string[] => {
|
||||
if (!isComplexityTier(tier)) return [];
|
||||
const deployment = deploymentFor(routerName, routerType, autoRouters);
|
||||
if (!deployment) return [];
|
||||
const config = asRecord(deployment.litellm_params?.complexity_router_config);
|
||||
const tiers = asRecord(config.tiers);
|
||||
return normalizeTierModels(tiers[tier]);
|
||||
};
|
||||
|
||||
interface TierTurnsChartProps {
|
||||
view: BenchmarkView;
|
||||
autoRouters: readonly AutoRouterDeployment[];
|
||||
}
|
||||
|
||||
const TIER_DONUT_COLORS: readonly ChartColor[] = ["#c7d2fe", "#1e293b", "#d4b483", "#87a878"];
|
||||
|
||||
const TierTurnsChart: React.FC<TierTurnsChartProps> = ({ view, autoRouters }) => {
|
||||
const group = viewGroup(view);
|
||||
const entries = Object.entries(group?.tier_turns ?? {}).filter(([, turns]) => turns > 0);
|
||||
if (!group || entries.length === 0) return null;
|
||||
|
||||
const tierLabels = tierLabelsFor(group.router_name, group.router_type, autoRouters);
|
||||
const total = entries.reduce((sum, [, turns]) => sum + turns, 0);
|
||||
const slices = entries.map(([tier, turns]) => ({
|
||||
tier: tierDisplayLabel(tier, tierLabels),
|
||||
turns,
|
||||
models: tierModelsFor(tier, group.router_name, group.router_type, autoRouters),
|
||||
}));
|
||||
const colors = slices.map((_, idx) => TIER_DONUT_COLORS[idx % TIER_DONUT_COLORS.length]);
|
||||
|
||||
return (
|
||||
<Card>
|
||||
<CardHeader>
|
||||
<CardTitle>Routing by tier</CardTitle>
|
||||
<p className="text-sm text-muted-foreground">
|
||||
Turns each tier served. Turns the classifier sent to the default model belong to no tier and are not counted
|
||||
here, so this can total less than the router's turns.
|
||||
</p>
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
<div className="grid grid-cols-1 items-center gap-6 lg:grid-cols-2">
|
||||
<DonutChart
|
||||
className="h-80"
|
||||
data={slices}
|
||||
index="tier"
|
||||
category="turns"
|
||||
colors={colors}
|
||||
valueFormatter={(value) => value.toLocaleString()}
|
||||
showLabel
|
||||
label={`${total.toLocaleString()} total turns`}
|
||||
/>
|
||||
<ul className="flex flex-col gap-6">
|
||||
{slices.map((slice, idx) => (
|
||||
<li key={slice.tier} className="flex items-start gap-2">
|
||||
<span
|
||||
className="mt-1.5 h-2 w-2 shrink-0 rounded-full ring-4 ring-white"
|
||||
style={{ backgroundColor: chartColorValue(colors[idx]) }}
|
||||
/>
|
||||
<div className="min-w-0">
|
||||
<p className="text-sm text-muted-foreground">
|
||||
{slice.tier} {Math.round((100 * slice.turns) / total).toLocaleString()}%
|
||||
</p>
|
||||
{slice.models.length > 0 && (
|
||||
<p className="text-xs break-words text-muted-foreground/70">{slice.models.join(", ")}</p>
|
||||
)}
|
||||
</div>
|
||||
</li>
|
||||
))}
|
||||
</ul>
|
||||
</div>
|
||||
</CardContent>
|
||||
</Card>
|
||||
);
|
||||
};
|
||||
|
||||
export default TierTurnsChart;
|
||||
|
|
@ -24,9 +24,12 @@ export const windowFor = (range: BenchmarkWindow, now: Date): { start_date: stri
|
|||
|
||||
export interface BenchmarkView {
|
||||
label: string;
|
||||
stats: AutoRouterBenchmarkTotals;
|
||||
stats: AutoRouterBenchmarkTotals | AutoRouterBenchmarkGroup;
|
||||
}
|
||||
|
||||
export const viewGroup = (view: BenchmarkView): AutoRouterBenchmarkGroup | null =>
|
||||
"router_name" in view.stats ? view.stats : null;
|
||||
|
||||
export const groupKey = (group: AutoRouterBenchmarkGroup): string => `${group.router_name} ${group.router_type}`;
|
||||
|
||||
export const groupLabel = (group: AutoRouterBenchmarkGroup, groups: readonly AutoRouterBenchmarkGroup[]): string => {
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -21403,7 +21403,7 @@ export interface components {
|
|||
spend: number;
|
||||
/**
|
||||
* Tier Turns
|
||||
* @description Turns per tier, keyed by the tier name the routing decision recorded at request time (never re-derived at read time, since the tier-to-model mapping is mutable config). Tier names are scoped to this group's router_type and are not comparable across types: a complexity router reports 'simple'/'medium'/'complex'/'reasoning', a quality router reports its numeric quality tier, and an adaptive router records no tier at all. Turns no tier served (the classifier fell back to default_model) are absent rather than pooled under a sentinel key, so the values may sum to less than turns
|
||||
* @description Turns per tier, keyed by the tier name the routing decision recorded at request time (never re-derived at read time, since the tier-to-model mapping is mutable config). Tier names are scoped to this group's router_type and are not comparable across types: a complexity router reports 'SIMPLE'/'MEDIUM'/'COMPLEX'/'REASONING', a quality router reports its numeric quality tier, and an adaptive router records no tier at all. Turns no tier served (the classifier fell back to default_model) are absent rather than pooled under a sentinel key, so the values may sum to less than turns
|
||||
*/
|
||||
tier_turns?: {
|
||||
[key: string]: number;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue