mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'litellm_internal_staging' into litellm_shadcn_logs_drawer_header_0813
This commit is contained in:
commit
1dc0ea3d11
70 changed files with 4273 additions and 1097 deletions
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 23914
|
||||
"limit": 22947
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2580
|
||||
"limit": 2579
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 323
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 7573
|
||||
"limit": 7312
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5719
|
||||
"limit": 5707
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15657
|
||||
"limit": 15642
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,22 +99,22 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44832
|
||||
"limit": 44776
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 39269
|
||||
"limit": 39237
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19988
|
||||
"limit": 19969
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30923
|
||||
"limit": 30881
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 118
|
||||
"limit": 117
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 699
|
||||
|
|
|
|||
|
|
@ -0,0 +1,49 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_ShadowEvalJob" (
|
||||
"id" TEXT NOT NULL,
|
||||
"api_key_id" TEXT NOT NULL,
|
||||
"router_name" TEXT NOT NULL,
|
||||
"judge_model" TEXT NOT NULL,
|
||||
"shadow_percentage" DOUBLE PRECISION NOT NULL,
|
||||
"max_turns" INTEGER NOT NULL,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"created_by" TEXT,
|
||||
"ends_at" TIMESTAMP(3) NOT NULL,
|
||||
"stopped_at" TIMESTAMP(3),
|
||||
|
||||
CONSTRAINT "LiteLLM_ShadowEvalJob_pkey" PRIMARY KEY ("id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_ShadowEvalAttempt" (
|
||||
"id" TEXT NOT NULL,
|
||||
"job_id" TEXT NOT NULL,
|
||||
"request_id" TEXT NOT NULL,
|
||||
"outcome" TEXT NOT NULL,
|
||||
"tier" TEXT,
|
||||
"real_model" TEXT,
|
||||
"shadow_model" TEXT,
|
||||
"confidence" DOUBLE PRECISION,
|
||||
"judge_cost" DOUBLE PRECISION NOT NULL DEFAULT 0,
|
||||
"error" TEXT,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_ShadowEvalAttempt_pkey" PRIMARY KEY ("id")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_ShadowEvalJob_api_key_id_idx" ON "LiteLLM_ShadowEvalJob"("api_key_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_ShadowEvalJob_created_at_idx" ON "LiteLLM_ShadowEvalJob"("created_at");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_ShadowEvalAttempt_job_id_idx" ON "LiteLLM_ShadowEvalAttempt"("job_id");
|
||||
|
||||
|
||||
-- One active job per key, enforced by the database rather than a read-then-create in the
|
||||
-- start endpoint, which races against a concurrent start on another pod. Partial indexes
|
||||
-- are not expressible in schema.prisma, so this lives here only. Active means not yet
|
||||
-- stopped; the start endpoint stamps stopped_at on expired jobs before creating.
|
||||
CREATE UNIQUE INDEX "LiteLLM_ShadowEvalJob_one_active_per_key"
|
||||
ON "LiteLLM_ShadowEvalJob"("api_key_id") WHERE "stopped_at" IS NULL;
|
||||
|
|
@ -1450,6 +1450,44 @@ model LiteLLM_AutoRouterSession {
|
|||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic.
|
||||
// A sampled slice of requests is duplicated through the router in a detached task and an
|
||||
// LLM judge compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// stopped_at; every count, status, and spend figure is derived from the append-only
|
||||
// attempt rows, so nothing can disagree across pods or stop races.
|
||||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
api_key_id String // hashed virtual key whose traffic is shadowed
|
||||
router_name String
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample budget: judge at most this many turns
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
ends_at DateTime
|
||||
stopped_at DateTime?
|
||||
|
||||
@@index([api_key_id])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
// One row per sampled pipeline: a blind verdict (real | shadow | tie) or an error.
|
||||
model LiteLLM_ShadowEvalAttempt {
|
||||
id String @id @default(cuid())
|
||||
job_id String
|
||||
request_id String // the judged real request
|
||||
outcome String // real | shadow | tie | error
|
||||
tier String? // router's tier for the prompt, when classified
|
||||
real_model String?
|
||||
shadow_model String?
|
||||
confidence Float?
|
||||
judge_cost Float @default(0)
|
||||
error String?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
@@index([job_id])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ A2A Streaming Events (in order):
|
|||
4. Status update (kind: "status-update") - Final status "completed" with final=true
|
||||
"""
|
||||
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from collections.abc import AsyncIterator, Callable, Coroutine, Mapping
|
||||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
|
|
@ -54,7 +54,7 @@ class A2ACompletionBridgeHandler:
|
|||
agent_extra_headers: Mapping[str, str] | None,
|
||||
*,
|
||||
stream: bool,
|
||||
) -> Mapping[str, Any]:
|
||||
) -> Mapping[str, object]:
|
||||
# Extract message from params
|
||||
message: Final = params.get("message", {})
|
||||
|
||||
|
|
@ -63,7 +63,7 @@ class A2ACompletionBridgeHandler:
|
|||
|
||||
# Get completion params
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
model: Final = litellm_params.get("model", "agent")
|
||||
model: Final[str] = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
# Skip prepending if model already starts with the provider prefix
|
||||
|
|
@ -109,13 +109,16 @@ class A2ACompletionBridgeHandler:
|
|||
return completion_params
|
||||
|
||||
@staticmethod
|
||||
async def _acompletion(completion_params: Mapping[str, Any]) -> ModelResponse | CustomStreamWrapper:
|
||||
return await litellm.acompletion(**completion_params)
|
||||
async def _acompletion(completion_params: Mapping[str, object]) -> ModelResponse | CustomStreamWrapper:
|
||||
acompletion_fn: Final[Callable[..., Coroutine[object, object, ModelResponse | CustomStreamWrapper]]] = vars(
|
||||
litellm
|
||||
)["acompletion"]
|
||||
return await acompletion_fn(**completion_params)
|
||||
|
||||
@staticmethod
|
||||
async def handle_non_streaming(
|
||||
request_id: str,
|
||||
params: dict[str, Any],
|
||||
params: dict[str, object],
|
||||
litellm_params: dict[str, Any],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -296,8 +299,8 @@ class A2ACompletionBridgeHandler:
|
|||
# Convenience functions that delegate to the class methods
|
||||
async def handle_a2a_completion(
|
||||
request_id: str,
|
||||
params: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
) -> dict[str, object]:
|
||||
|
|
@ -313,8 +316,8 @@ async def handle_a2a_completion(
|
|||
|
||||
async def handle_a2a_completion_streaming(
|
||||
request_id: str,
|
||||
params: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
api_base: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
) -> AsyncIterator[dict[str, object]]:
|
||||
|
|
|
|||
|
|
@ -12,7 +12,8 @@ Provides standalone functions with @client decorator for LiteLLM logging integra
|
|||
import asyncio
|
||||
import datetime
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Coroutine
|
||||
from collections.abc import AsyncIterator, Coroutine, Mapping
|
||||
from types import ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -38,12 +39,15 @@ if TYPE_CHECKING:
|
|||
SendMessageResponse,
|
||||
SendStreamingMessageRequest,
|
||||
SendStreamingMessageResponse,
|
||||
SendStreamingMessageSuccessResponse,
|
||||
Task,
|
||||
)
|
||||
from a2a.types.a2a_pb2 import SendMessageRequest as CoreSendMessageRequest
|
||||
from a2a.types.a2a_pb2 import StreamResponse as CoreStreamResponse
|
||||
|
||||
# Runtime imports — requires a2a-sdk>=1.1.0
|
||||
A2A_SDK_AVAILABLE = False
|
||||
_a2a_conversions: Any = None
|
||||
_a2a_conversions: ModuleType | None = None
|
||||
|
||||
try:
|
||||
from a2a.client import Client, ClientCallContext, ClientConfig, create_client
|
||||
|
|
@ -128,7 +132,7 @@ _A2A_COST_PARAM_KEYS: Final = ("cost_per_query", "input_cost_per_token", "output
|
|||
|
||||
def _set_litellm_params_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> None:
|
||||
"""
|
||||
Merge the agent's pricing params into model_call_details["litellm_params"]
|
||||
|
|
@ -150,7 +154,7 @@ def _set_litellm_params_on_logging_obj(
|
|||
logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params}
|
||||
|
||||
|
||||
def _get_a2a_model_info(a2a_client: Any, kwargs: dict[str, Any]) -> str:
|
||||
def _get_a2a_model_info(a2a_client: "A2AClientType", kwargs: dict[str, Any]) -> str:
|
||||
"""
|
||||
Extract agent info and set model/custom_llm_provider for cost tracking.
|
||||
|
||||
|
|
@ -179,7 +183,7 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: dict[str, Any]) -> str:
|
|||
return agent_name
|
||||
|
||||
|
||||
def _get_a2a_client_agent_card(a2a_client: Any) -> Optional["AgentCard"]:
|
||||
def _get_a2a_client_agent_card(a2a_client: "A2AClientType") -> Optional["AgentCard"]:
|
||||
agent_card = cast(Optional["AgentCard"], getattr(a2a_client, "_litellm_agent_card", None))
|
||||
if agent_card is not None:
|
||||
return agent_card
|
||||
|
|
@ -191,9 +195,9 @@ def _get_a2a_client_agent_card(a2a_client: Any) -> Optional["AgentCard"]:
|
|||
|
||||
async def _send_message_via_completion_bridge(
|
||||
request: "SendMessageRequest",
|
||||
custom_llm_provider: str,
|
||||
custom_llm_provider: object,
|
||||
api_base: str | None,
|
||||
litellm_params: dict[str, Any],
|
||||
litellm_params: dict[str, object],
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
|
|
@ -224,6 +228,20 @@ def _get_a2a_call_context(a2a_client: "A2AClientType") -> Optional["A2ACallConte
|
|||
return getattr(a2a_client, "_litellm_call_context", None)
|
||||
|
||||
|
||||
def _to_core_send_message_request(request: "SendMessageRequest") -> "CoreSendMessageRequest":
|
||||
from a2a.compat.v0_3 import conversions
|
||||
|
||||
return conversions.to_core_send_message_request(request)
|
||||
|
||||
|
||||
def _to_compat_stream_response(
|
||||
event: "CoreStreamResponse", request_id: str | int
|
||||
) -> "SendStreamingMessageSuccessResponse":
|
||||
from a2a.compat.v0_3 import conversions
|
||||
|
||||
return conversions.to_compat_stream_response(event, request_id=request_id)
|
||||
|
||||
|
||||
async def _send_message(a2a_client: "A2AClientType", request: "SendMessageRequest") -> "SendMessageResponse":
|
||||
"""Send a non-streaming message via a2a-sdk 1.x and return JSON-RPC response."""
|
||||
if _a2a_conversions is None:
|
||||
|
|
@ -231,17 +249,14 @@ async def _send_message(a2a_client: "A2AClientType", request: "SendMessageReques
|
|||
"The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk"
|
||||
)
|
||||
|
||||
pb_request: Final = _a2a_conversions.to_core_send_message_request(request)
|
||||
pb_request: Final = _to_core_send_message_request(request)
|
||||
last_event = None
|
||||
async for event in a2a_client.send_message(pb_request, context=_get_a2a_call_context(a2a_client)):
|
||||
last_event = event
|
||||
if last_event is None:
|
||||
raise RuntimeError("A2A send_message failed: no response received from agent.")
|
||||
|
||||
stream_compat: Final = _a2a_conversions.to_compat_stream_response(
|
||||
last_event,
|
||||
request_id=request.id,
|
||||
)
|
||||
stream_compat: Final = _to_compat_stream_response(last_event, request_id=request.id)
|
||||
result: Final = stream_compat.result
|
||||
if not isinstance(result, (Message, Task)):
|
||||
raise RuntimeError(
|
||||
|
|
@ -306,12 +321,9 @@ async def _stream_messages(
|
|||
"The 'a2a' package is required for A2A agent invocation. Install it with: pip install a2a-sdk"
|
||||
)
|
||||
|
||||
pb_request: Final = _a2a_conversions.to_core_send_message_request(request)
|
||||
pb_request: Final[CoreSendMessageRequest] = _a2a_conversions.to_core_send_message_request(request)
|
||||
async for event in a2a_client.send_message(pb_request, context=_get_a2a_call_context(a2a_client)):
|
||||
compat_chunk = _a2a_conversions.to_compat_stream_response(
|
||||
event,
|
||||
request_id=request.id,
|
||||
)
|
||||
compat_chunk = _to_compat_stream_response(event, request_id=request.id)
|
||||
yield SendStreamingMessageResponse(root=compat_chunk)
|
||||
|
||||
|
||||
|
|
@ -368,10 +380,10 @@ async def asend_message(
|
|||
a2a_client: Optional["A2AClientType"] = None,
|
||||
request: Optional["SendMessageRequest"] = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
agent_id: str | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
**kwargs: Any,
|
||||
**kwargs: object,
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
Async: Send a message to an A2A agent.
|
||||
|
|
@ -485,7 +497,7 @@ async def asend_message(
|
|||
response: Final = LiteLLMSendMessageResponse.from_a2a_response(a2a_response, request_id=str(request.id))
|
||||
|
||||
# Calculate token usage from request and response
|
||||
response_dict: Final = a2a_response.model_dump(mode="json", exclude_none=True)
|
||||
response_dict: Final[dict[str, object]] = a2a_response.model_dump(mode="json", exclude_none=True)
|
||||
(
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
|
|
@ -516,7 +528,7 @@ def send_message(
|
|||
a2a_client: "A2AClientType",
|
||||
request: "SendMessageRequest",
|
||||
**kwargs: Any,
|
||||
) -> LiteLLMSendMessageResponse | Coroutine[Any, Any, LiteLLMSendMessageResponse]:
|
||||
) -> LiteLLMSendMessageResponse | Coroutine[object, object, LiteLLMSendMessageResponse]:
|
||||
"""
|
||||
Sync: Send a message to an A2A agent.
|
||||
|
||||
|
|
@ -545,9 +557,9 @@ def _build_streaming_logging_obj(
|
|||
request: "SendStreamingMessageRequest",
|
||||
agent_name: str,
|
||||
agent_id: str | None,
|
||||
litellm_params: dict[str, Any] | None,
|
||||
metadata: dict[str, Any] | None,
|
||||
proxy_server_request: dict[str, Any] | None,
|
||||
litellm_params: dict[str, object] | None,
|
||||
metadata: dict[str, object] | None,
|
||||
proxy_server_request: dict[str, object] | None,
|
||||
) -> Logging:
|
||||
"""Build logging object for streaming A2A requests."""
|
||||
start_time: Final = datetime.datetime.now()
|
||||
|
|
@ -588,10 +600,10 @@ async def asend_message_streaming(
|
|||
a2a_client: Optional["A2AClientType"] = None,
|
||||
request: Optional["SendStreamingMessageRequest"] = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
agent_id: str | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
proxy_server_request: dict[str, Any] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
proxy_server_request: dict[str, object] | None = None,
|
||||
agent_extra_headers: dict[str, str] | None = None,
|
||||
**kwargs: object,
|
||||
) -> AsyncIterator[Any]:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
import json
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Final, cast
|
||||
from typing import Any, Final, TypedDict, cast
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.litellm_core_utils.json_validation_rule import normalize_tool_schema
|
||||
|
|
@ -28,6 +30,19 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
|
||||
class _GenAITextPart(TypedDict, total=False):
|
||||
text: ReadOnly[str]
|
||||
|
||||
|
||||
class _GenAISystemInstruction(TypedDict, total=False):
|
||||
parts: ReadOnly[list[_GenAITextPart]]
|
||||
|
||||
|
||||
class _GenAIPart(TypedDict, total=False):
|
||||
text: ReadOnly[str]
|
||||
functionCall: ReadOnly[dict[str, object]]
|
||||
|
||||
|
||||
class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
||||
"""
|
||||
Wrapper for streaming Google GenAI generate_content responses.
|
||||
|
|
@ -36,9 +51,9 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
|
||||
sent_first_chunk: bool = False
|
||||
# State tracking for accumulating partial tool calls
|
||||
accumulated_tool_calls: dict[str, dict[str, Any]]
|
||||
accumulated_tool_calls: dict[str, dict[str, str]]
|
||||
|
||||
def __init__(self, completion_stream: Any):
|
||||
def __init__(self, completion_stream: object):
|
||||
self.sent_first_chunk = False
|
||||
self.accumulated_tool_calls = {}
|
||||
self._returned_response = False
|
||||
|
|
@ -85,7 +100,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# After the stream is exhausted, check for any remaining accumulated tool calls
|
||||
if self.accumulated_tool_calls:
|
||||
try:
|
||||
parts: Final = []
|
||||
parts: Final[list[_GenAIPart]] = []
|
||||
for (
|
||||
tool_call_index,
|
||||
tool_call_data,
|
||||
|
|
@ -94,7 +109,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
# For tool calls with no arguments, accumulated_args will be "", which is not valid JSON.
|
||||
# We default to an empty JSON object in this case.
|
||||
parsed_args = json.loads(tool_call_data["arguments"] or "{}")
|
||||
function_call_part = {
|
||||
function_call_part: _GenAIPart = {
|
||||
"functionCall": {
|
||||
"name": tool_call_data["name"] or "undefined_tool_name",
|
||||
"args": parsed_args,
|
||||
|
|
@ -110,7 +125,7 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
tool_call_data["arguments"],
|
||||
)
|
||||
if parts:
|
||||
final_chunk: Final = {
|
||||
final_chunk: Final[dict[str, object]] = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
|
|
@ -273,9 +288,9 @@ class GoogleGenAIAdapter:
|
|||
|
||||
def _add_generic_litellm_params_to_request(
|
||||
self,
|
||||
completion_request_dict: dict[str, Any],
|
||||
completion_request_dict: dict[str, object],
|
||||
litellm_params: GenericLiteLLMParams | None = None,
|
||||
) -> dict:
|
||||
) -> dict[str, object]:
|
||||
"""Add generic litellm params to request. e.g add api_base, api_key, api_version, etc.
|
||||
|
||||
Args:
|
||||
|
|
@ -295,7 +310,7 @@ class GoogleGenAIAdapter:
|
|||
|
||||
def translate_completion_output_params_streaming(
|
||||
self,
|
||||
completion_stream: Any,
|
||||
completion_stream: object,
|
||||
) -> AsyncIterator[bytes] | None:
|
||||
"""Transform streaming completion output to Google GenAI format"""
|
||||
google_genai_wrapper: Final = GoogleGenAIStreamWrapper(completion_stream=completion_stream)
|
||||
|
|
@ -307,12 +322,12 @@ class GoogleGenAIAdapter:
|
|||
tools: list[dict[str, Any]],
|
||||
) -> list[ChatCompletionToolParam]:
|
||||
"""Transform Google GenAI tools to OpenAI tools format"""
|
||||
openai_tools: Final[list[dict[str, Any]]] = []
|
||||
openai_tools: Final[list[dict[str, object]]] = []
|
||||
|
||||
for tool in tools:
|
||||
if "functionDeclarations" in tool:
|
||||
for func_decl in tool["functionDeclarations"]:
|
||||
function_chunk: dict[str, Any] = {
|
||||
function_chunk: dict[str, object] = {
|
||||
"name": func_decl.get("name", ""),
|
||||
}
|
||||
|
||||
|
|
@ -321,7 +336,7 @@ class GoogleGenAIAdapter:
|
|||
if "parametersJsonSchema" in func_decl:
|
||||
function_chunk["parameters"] = func_decl["parametersJsonSchema"]
|
||||
|
||||
openai_tool = {"type": "function", "function": function_chunk}
|
||||
openai_tool: dict[str, object] = {"type": "function", "function": function_chunk}
|
||||
openai_tools.append(openai_tool)
|
||||
|
||||
# normalize the tool schemas
|
||||
|
|
@ -345,7 +360,7 @@ class GoogleGenAIAdapter:
|
|||
def _transform_contents_to_messages(
|
||||
self,
|
||||
contents: list[dict[str, Any]],
|
||||
system_instruction: dict[str, Any] | None = None,
|
||||
system_instruction: _GenAISystemInstruction | None = None,
|
||||
) -> list[AllMessageValues]:
|
||||
"""Transform Google GenAI contents to OpenAI messages format"""
|
||||
messages: Final[list[AllMessageValues]] = []
|
||||
|
|
@ -461,7 +476,7 @@ class GoogleGenAIAdapter:
|
|||
def translate_completion_to_generate_content(
|
||||
self,
|
||||
response: ModelResponse,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Transform litellm completion response to Google GenAI generate_content format
|
||||
|
||||
|
|
@ -490,7 +505,7 @@ class GoogleGenAIAdapter:
|
|||
parts = [{"text": message_content}] if message_content else []
|
||||
|
||||
# Create Google GenAI format response
|
||||
generate_content_response: Final[dict[str, Any]] = {
|
||||
generate_content_response: Final[dict[str, object]] = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
|
|
@ -524,7 +539,7 @@ class GoogleGenAIAdapter:
|
|||
self,
|
||||
response: ModelResponse | ModelResponseStream,
|
||||
wrapper: GoogleGenAIStreamWrapper,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Transform streaming litellm completion chunk to Google GenAI generate_content format
|
||||
|
||||
|
|
@ -560,7 +575,7 @@ class GoogleGenAIAdapter:
|
|||
return None
|
||||
|
||||
# Create Google GenAI streaming format response
|
||||
streaming_chunk: Final[dict[str, Any]] = {
|
||||
streaming_chunk: Final[dict[str, object]] = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
|
|
@ -597,9 +612,9 @@ class GoogleGenAIAdapter:
|
|||
def _transform_openai_message_to_google_genai_parts(
|
||||
self,
|
||||
message: Any,
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[_GenAIPart]:
|
||||
"""Transform OpenAI message to Google GenAI parts format"""
|
||||
parts: Final[list[dict[str, Any]]] = []
|
||||
parts: Final[list[_GenAIPart]] = []
|
||||
|
||||
# Add text content if present
|
||||
if hasattr(message, "content") and message.content:
|
||||
|
|
@ -614,7 +629,7 @@ class GoogleGenAIAdapter:
|
|||
except json.JSONDecodeError:
|
||||
args = {}
|
||||
|
||||
function_call_part = {
|
||||
function_call_part: _GenAIPart = {
|
||||
"functionCall": {
|
||||
"name": tool_call.function.name or "undefined_tool_name",
|
||||
"args": args,
|
||||
|
|
@ -626,14 +641,14 @@ class GoogleGenAIAdapter:
|
|||
|
||||
def _transform_openai_delta_to_google_genai_parts_with_accumulation(
|
||||
self, delta: Any, wrapper: GoogleGenAIStreamWrapper
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[_GenAIPart]:
|
||||
"""Transforms OpenAI delta to Google GenAI parts, accumulating streaming tool calls."""
|
||||
|
||||
# 1. Initialize wrapper state if it doesn't exist
|
||||
if not hasattr(wrapper, "accumulated_tool_calls"):
|
||||
wrapper.accumulated_tool_calls = {}
|
||||
|
||||
parts: Final[list[dict[str, Any]]] = []
|
||||
parts: Final[list[_GenAIPart]] = []
|
||||
|
||||
if hasattr(delta, "content") and delta.content:
|
||||
parts.append({"text": delta.content})
|
||||
|
|
@ -686,7 +701,7 @@ class GoogleGenAIAdapter:
|
|||
# The part will be created by a later chunk that brings the name.
|
||||
if accumulated_name:
|
||||
# If successful, create the part and clean up
|
||||
function_call_part = {"functionCall": {"name": accumulated_name, "args": parsed_args}}
|
||||
function_call_part: _GenAIPart = {"functionCall": {"name": accumulated_name, "args": parsed_args}}
|
||||
parts.append(function_call_part)
|
||||
|
||||
# Remove the completed tool call from the accumulator
|
||||
|
|
|
|||
|
|
@ -6,12 +6,13 @@ import random
|
|||
import time
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict
|
||||
|
||||
import httpx
|
||||
from typing_extensions import Never, ReadOnly
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
|
|
@ -48,7 +49,20 @@ _WEBHOOK_PATH_PROMPT_MODERATION: Final = "/v1/before_prompt/openai/v1"
|
|||
_WEBHOOK_PATH_LOGGING_BATCH: Final = "/v1/litellm/batch"
|
||||
_MAX_QUEUE_SIZE: Final = 10_000
|
||||
_DROP_WARNING_INTERVAL_SECONDS: Final = 60.0
|
||||
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
|
||||
_EMPTY_MAPPING: Final[Mapping[str, Never]] = MappingProxyType({})
|
||||
|
||||
|
||||
class _ServiceToolCall(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
|
||||
|
||||
class _ServiceMessage(TypedDict, total=False):
|
||||
content: ReadOnly[str]
|
||||
tool_calls: ReadOnly[Sequence[_ServiceToolCall]]
|
||||
|
||||
|
||||
class _ServiceChoice(TypedDict, total=False):
|
||||
message: ReadOnly[_ServiceMessage]
|
||||
|
||||
|
||||
class _MalformedToolBlockingResponseError(Exception):
|
||||
|
|
@ -143,7 +157,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
else {"Content-Type": "application/json"}
|
||||
)
|
||||
|
||||
self._periodic_flush_task: asyncio.Task[Any] | None = self._start_periodic_flush_task()
|
||||
self._periodic_flush_task: asyncio.Task[None] | None = self._start_periodic_flush_task()
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
|
|
@ -191,7 +205,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
params={"timeout": httpx.Timeout(5.0, connect=2.0)},
|
||||
)
|
||||
|
||||
def _start_periodic_flush_task(self) -> asyncio.Task[Any] | None:
|
||||
def _start_periodic_flush_task(self) -> asyncio.Task[None] | None:
|
||||
"""Start the periodic flush task only when an event loop is already running."""
|
||||
try:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
|
|
@ -212,7 +226,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
Closing them here would close the shared connection pool for every
|
||||
other logger instance; let LiteLLM manage their lifecycle instead.
|
||||
"""
|
||||
task: Final = getattr(self, "_periodic_flush_task", None)
|
||||
task: Final[asyncio.Task[None] | None] = getattr(self, "_periodic_flush_task", None)
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
|
||||
|
|
@ -253,7 +267,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
|
||||
@staticmethod
|
||||
async def _guarded(
|
||||
coro: Any,
|
||||
coro: Awaitable[GenericGuardrailAPIInputs],
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
label: str,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
|
|
@ -400,7 +414,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
request_data["_rubrik_logging_obj"] = logging_obj
|
||||
|
||||
@staticmethod
|
||||
def _normalize_tool_calls(tool_calls: Any) -> tuple[ChatCompletionMessageToolCall, ...]:
|
||||
def _normalize_tool_calls(tool_calls: Sequence[object]) -> tuple[ChatCompletionMessageToolCall, ...]:
|
||||
"""Convert tool_calls from inputs to ChatCompletionMessageToolCall objects."""
|
||||
return tuple(RubrikLogger._normalize_tool_call(tc) for tc in tool_calls)
|
||||
|
||||
|
|
@ -427,7 +441,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
raise TypeError(f"Cannot normalize tool_call of type {type(tc).__name__}: {tc!r}")
|
||||
|
||||
@staticmethod
|
||||
def _join_texts(texts: Any) -> str:
|
||||
def _join_texts(texts: Sequence[str] | None) -> str:
|
||||
"""Join response text segments into the single content string the
|
||||
webhook evaluates. Empty when there is no assistant text."""
|
||||
if not texts:
|
||||
|
|
@ -439,14 +453,14 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
tool_calls: Sequence[ChatCompletionMessageToolCall],
|
||||
content: str,
|
||||
request_id: str | None,
|
||||
) -> Mapping[str, Any]:
|
||||
) -> Mapping[str, object]:
|
||||
"""Build an OpenAI ChatCompletion-format dict (assistant text + tool
|
||||
calls) for the after_completion webhook.
|
||||
|
||||
``content`` is sent so the webhook can moderate the response text;
|
||||
``None`` when the assistant produced no text (tool-call-only response).
|
||||
"""
|
||||
message: Final[dict[str, Any]] = {
|
||||
message: Final[dict[str, object]] = {
|
||||
"role": "assistant",
|
||||
"content": content or None,
|
||||
}
|
||||
|
|
@ -467,7 +481,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
}
|
||||
|
||||
@staticmethod
|
||||
def _flatten_messages_for_moderation(messages: Any) -> tuple[Mapping[str, Any], ...]:
|
||||
def _flatten_messages_for_moderation(messages: Sequence[object] | None) -> tuple[Mapping[str, Any], ...]:
|
||||
"""Collapse each message's content to a plain string for the webhook.
|
||||
|
||||
litellm normalizes Anthropic ``/v1/messages`` requests to OpenAI shape,
|
||||
|
|
@ -506,8 +520,8 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
@staticmethod
|
||||
def _build_prompt_moderation_payload(
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: Mapping[str, Any],
|
||||
) -> Mapping[str, Any]:
|
||||
request_data: Mapping[str, object],
|
||||
) -> Mapping[str, object]:
|
||||
"""Build the bare OpenAI request the before_prompt webhook consumes.
|
||||
|
||||
Unlike the after_completion envelope, this endpoint takes a raw OpenAI
|
||||
|
|
@ -516,7 +530,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
``/v1/messages`` requests too. Optional fields are sent only when
|
||||
present so the payload stays clean.
|
||||
"""
|
||||
payload: Final[dict[str, Any]] = {
|
||||
payload: Final[dict[str, object]] = {
|
||||
"model": inputs.get("model") or request_data.get("model") or "",
|
||||
"messages": RubrikLogger._flatten_messages_for_moderation(inputs.get("structured_messages")),
|
||||
}
|
||||
|
|
@ -540,8 +554,8 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
@staticmethod
|
||||
def _extract_request_data(
|
||||
call_details: Mapping[str, Any],
|
||||
request_data: Mapping[str, Any] | None,
|
||||
) -> Mapping[str, Any]:
|
||||
request_data: Mapping[str, object] | None,
|
||||
) -> Mapping[str, object]:
|
||||
"""Extract original request data from model_call_details for the
|
||||
response moderation service envelope.
|
||||
|
||||
|
|
@ -576,7 +590,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
}
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_proxy_server_request(proxy_server_request: Any) -> Any:
|
||||
def _sanitize_proxy_server_request(proxy_server_request: object) -> object:
|
||||
"""Allowlist only routing fields (``url``, ``method``) when forwarding
|
||||
``proxy_server_request`` to an external webhook, dropping inbound
|
||||
``headers`` (Authorization, Cookie, x-api-key, ...) and the raw
|
||||
|
|
@ -586,17 +600,18 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
return {key: proxy_server_request[key] for key in ("url", "method") if key in proxy_server_request}
|
||||
|
||||
@staticmethod
|
||||
def _resolve_model(request_data: Mapping[str, Any], call_details: Mapping[str, Any]) -> str:
|
||||
def _resolve_model(request_data: Mapping[str, object], call_details: Mapping[str, str]) -> str:
|
||||
"""Get the model name for the ModifyResponseException."""
|
||||
response: Final = request_data.get("response")
|
||||
if response and hasattr(response, "model"):
|
||||
return response.model or "unknown"
|
||||
response_model: Final[str | None] = getattr(response, "model", None)
|
||||
return response_model or "unknown"
|
||||
return call_details.get("model", "unknown")
|
||||
|
||||
# -- Logging hooks ---------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _correlation_id(call_details: Mapping[str, Any], request_data: Mapping[str, Any] | None = None) -> str | None:
|
||||
def _correlation_id(call_details: Mapping[str, str], request_data: Mapping[str, str] | None = None) -> str | None:
|
||||
"""The id that joins a blocked request's two S3 logs by filename: the
|
||||
moderation (``_blocking``) log and the failure (response) log.
|
||||
|
||||
|
|
@ -610,7 +625,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
return call_details.get("litellm_call_id") or (request_data or _EMPTY_MAPPING).get("litellm_call_id")
|
||||
|
||||
@classmethod
|
||||
def _apply_correlation_id(cls, payload: dict[str, Any], source: Mapping[str, Any]) -> None:
|
||||
def _apply_correlation_id(cls, payload: dict[str, object], source: Mapping[str, str]) -> None:
|
||||
"""Pin ``payload["id"]`` to ``litellm_call_id`` in place so this log
|
||||
shares its S3 filename id with the moderation (``_blocking``) and
|
||||
failure logs for the same request -- for every provider.
|
||||
|
|
@ -630,7 +645,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
payload["id"] = correlated
|
||||
|
||||
@staticmethod
|
||||
def _prepend_system_prompt(payload: dict[str, Any], source: Mapping[str, Any]) -> None:
|
||||
def _prepend_system_prompt(payload: dict[str, object], source: Mapping[str, object]) -> None:
|
||||
"""Prepend ``source["system"]`` onto ``payload["messages"]``.
|
||||
|
||||
Builds a NEW messages list rather than mutating ``payload["messages"]``
|
||||
|
|
@ -658,7 +673,9 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
exc_info=True,
|
||||
)
|
||||
|
||||
async def _prepare_log_payload(self, kwargs: Mapping[str, Any], event_type: str) -> StandardLoggingPayload | None:
|
||||
async def _prepare_log_payload(
|
||||
self, kwargs: Mapping[str, object], event_type: str
|
||||
) -> StandardLoggingPayload | None:
|
||||
"""Shared logic for success logging (sampled)."""
|
||||
if random.random() > self.sampling_rate:
|
||||
verbose_logger.debug("Skipping Rubrik %s logging (sampling_rate=%s)", event_type, self.sampling_rate)
|
||||
|
|
@ -697,7 +714,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
self._dropped_since_warning = 0
|
||||
self._last_drop_warning_time = now
|
||||
|
||||
async def _enqueue_log_event(self, kwargs: Mapping[str, Any], event_type: str):
|
||||
async def _enqueue_log_event(self, kwargs: Mapping[str, object], event_type: str):
|
||||
try:
|
||||
payload: Final = await self._prepare_log_payload(kwargs, event_type)
|
||||
if payload is None:
|
||||
|
|
@ -862,7 +879,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
|
||||
base: Final = call_details.get("standard_logging_object")
|
||||
if base is not None:
|
||||
payload: dict = safe_deep_copy(base)
|
||||
payload: dict[str, object] = safe_deep_copy(base)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"Rubrik: standard_logging_object not yet on model_call_details "
|
||||
|
|
@ -908,7 +925,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
cls,
|
||||
call_details: Mapping[str, Any],
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
# Convert datetime to a Unix float so json.dumps can serialize it.
|
||||
# httpx's json= parameter uses stdlib json.dumps with no custom encoder.
|
||||
_raw_start: Final = call_details.get("start_time")
|
||||
|
|
@ -996,7 +1013,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
|
||||
# -- Webhook services ------------------------------------------------------
|
||||
|
||||
async def _post_json(self, endpoint: str, payload: Mapping[str, Any], service_name: str) -> Mapping[str, Any]:
|
||||
async def _post_json(self, endpoint: str, payload: Mapping[str, object], service_name: str) -> Mapping[str, Any]:
|
||||
"""POST ``payload`` to a Rubrik webhook and return its dict response.
|
||||
|
||||
Raises:
|
||||
|
|
@ -1010,7 +1027,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
headers=self._headers,
|
||||
)
|
||||
http_response.raise_for_status()
|
||||
result: Final = http_response.json()
|
||||
result: Final[object] = http_response.json()
|
||||
if not isinstance(result, dict):
|
||||
raise TypeError(
|
||||
f"{service_name} returned non-dict JSON "
|
||||
|
|
@ -1021,8 +1038,8 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
|
||||
async def _post_to_response_moderation_endpoint(
|
||||
self,
|
||||
response_data: Mapping[str, Any],
|
||||
request_data: Mapping[str, Any],
|
||||
response_data: Mapping[str, object],
|
||||
request_data: Mapping[str, object],
|
||||
) -> Mapping[str, Any]:
|
||||
"""Post the ``{request, response}`` envelope to the after_completion
|
||||
webhook and return its (possibly rewritten) response.
|
||||
|
|
@ -1039,7 +1056,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
"Response moderation service",
|
||||
)
|
||||
|
||||
async def _post_to_prompt_moderation_endpoint(self, payload: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
async def _post_to_prompt_moderation_endpoint(self, payload: Mapping[str, object]) -> Mapping[str, Any]:
|
||||
"""Post a bare OpenAI request to the before_prompt webhook.
|
||||
|
||||
Returns ``{}`` (passthrough) or a synthetic chat.completion (block).
|
||||
|
|
@ -1054,7 +1071,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
chat.completion whose ``choices[0].message.content`` is the refusal
|
||||
explanation.
|
||||
"""
|
||||
choices: Final = service_response.get("choices")
|
||||
choices: Final[Sequence[_ServiceChoice] | None] = service_response.get("choices")
|
||||
if not choices:
|
||||
return None
|
||||
message: Final = choices[0].get("message") or _EMPTY_MAPPING
|
||||
|
|
@ -1086,7 +1103,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
Expects service_response in OpenAI chat completion format:
|
||||
{"choices": [{"message": {"tool_calls": [...], "content": "..."}}]}
|
||||
"""
|
||||
choices: Final = service_response.get("choices") or ()
|
||||
choices: Final[Sequence[_ServiceChoice]] = service_response.get("choices") or ()
|
||||
if not choices:
|
||||
raise _MalformedToolBlockingResponseError("Response moderation service returned empty response")
|
||||
|
||||
|
|
|
|||
563
litellm/integrations/shadow_eval_logger.py
Normal file
563
litellm/integrations/shadow_eval_logger.py
Normal file
|
|
@ -0,0 +1,563 @@
|
|||
"""Shadow Eval Logger: samples a shadowed key's successful chat requests, duplicates each
|
||||
through the auto-router in a detached task, blind-judges real vs shadow, and appends one
|
||||
``LiteLLM_ShadowEvalAttempt`` row (verdict or error) as the feature's only hot-path write.
|
||||
Counts, status, and spend derive from those rows at read time, so nothing can disagree
|
||||
across pods or stop races; the hook reads active jobs through a short-TTL cache."""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import random
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
from litellm.litellm_core_utils.internal_call_metadata import sanitized_forwardable_call_metadata
|
||||
from litellm.litellm_core_utils.llm_judge import (
|
||||
default_router_provider,
|
||||
extract_text_from_content,
|
||||
judge_acompletion,
|
||||
parse_json_verdict,
|
||||
)
|
||||
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.router import Router
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
# A job starting, stopping, or hitting its turn budget propagates to sampling within one
|
||||
# TTL; the turn budget can overshoot by at most one TTL of in-flight samples per pod.
|
||||
_JOBS_CACHE_TTL_SECONDS: Final = 10
|
||||
|
||||
# Concurrent shadow+judge pipelines per pod: a traffic spike turns into skipped samples
|
||||
# rather than an unbounded task pileup.
|
||||
_MAX_CONCURRENT_SHADOW_TASKS: Final = 16
|
||||
|
||||
# Total character budget for the judge's user prompt, however long the conversation and
|
||||
# the two responses are, so the prompt can never overflow a judge model's context window.
|
||||
_MAX_JUDGE_RESPONSE_CHARS: Final = 8_000
|
||||
_MAX_JUDGE_PROMPT_CHARS: Final = 24_000
|
||||
|
||||
# The judge answers with a small JSON object; a tighter budget truncates the JSON
|
||||
# mid-object and the attempt is lost to an error row.
|
||||
JUDGE_MAX_OUTPUT_TOKENS: Final = 500
|
||||
|
||||
_MAX_ERROR_CHARS: Final = 500
|
||||
|
||||
_EMPTY_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
_SAMPLED_CALL_TYPES: Final = frozenset({"completion", "acompletion"})
|
||||
|
||||
PAIRWISE_JUDGE_SYSTEM_PROMPT: Final = """You are an impartial quality judge comparing two responses to the same conversation.
|
||||
|
||||
The responses are labeled A and B in random order. You do not know which system produced which.
|
||||
|
||||
Criteria: correctness, completeness, clarity, conciseness.
|
||||
|
||||
Return ONLY valid JSON in this exact format, no other text:
|
||||
{
|
||||
"preference": "A" | "B" | "tie",
|
||||
"confidence": <0.0 to 1.0>,
|
||||
"reasoning": "<one sentence>"
|
||||
}"""
|
||||
|
||||
|
||||
class PairwiseVerdict(BaseModel):
|
||||
"""The judge's blind A/B verdict, validated at the parse boundary."""
|
||||
|
||||
preference: str = "tie"
|
||||
confidence: float = 0.0
|
||||
|
||||
|
||||
def _sample_hits(request_id: str, job_id: str, percentage: float) -> bool:
|
||||
"""Deterministically decide whether a request falls in the shadowed slice: hash-based
|
||||
rather than random so retries sample the same way and pods agree without coordination."""
|
||||
digest: Final = hashlib.sha256(f"{job_id}:{request_id}".encode()).digest()
|
||||
bucket: Final = int.from_bytes(digest[:8], "big") / float(2**64)
|
||||
return bucket * 100.0 < percentage
|
||||
|
||||
|
||||
def _judge_call_cost(response: object) -> float:
|
||||
"""Price a judge call, treating an unmapped judge model as free rather than fatal."""
|
||||
import litellm
|
||||
|
||||
try:
|
||||
return litellm.completion_cost(completion_response=response) or 0.0
|
||||
except Exception: # noqa: BLE001 # unmapped judge model: the verdict still counts, cost stays 0
|
||||
return 0.0
|
||||
|
||||
|
||||
def _unmask_preference(raw_preference: str, real_is_a: bool) -> str:
|
||||
"""Map the judge's blind A/B/tie verdict back to real/shadow/tie."""
|
||||
normalized: Final = raw_preference.strip().lower()
|
||||
if normalized == "a":
|
||||
return "real" if real_is_a else "shadow"
|
||||
if normalized == "b":
|
||||
return "shadow" if real_is_a else "real"
|
||||
return "tie"
|
||||
|
||||
|
||||
def _judge_user_prompt(conversation: str, response_a: str, response_b: str) -> str:
|
||||
"""The judge prompt under one total character budget: each response is capped, and
|
||||
the conversation tail gets whatever budget the responses left over."""
|
||||
a: Final = response_a[:_MAX_JUDGE_RESPONSE_CHARS]
|
||||
b: Final = response_b[:_MAX_JUDGE_RESPONSE_CHARS]
|
||||
conversation_budget: Final = _MAX_JUDGE_PROMPT_CHARS - len(a) - len(b)
|
||||
return (
|
||||
f"Conversation:\n{conversation[-conversation_budget:]}\n\n"
|
||||
f"Response A:\n{a}\n\n"
|
||||
f"Response B:\n{b}\n\n"
|
||||
"Which response is better?"
|
||||
)
|
||||
|
||||
|
||||
async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool:
|
||||
"""Whether the shadowed key or its team is over budget, decided by the same owners
|
||||
the request path uses, so counter keys and thresholds can never drift from auth's.
|
||||
|
||||
Advisory and fail-open: real traffic on an over-budget key is already rejected at
|
||||
auth (so nothing reaches the success hook), and this gate only closes the race
|
||||
where the key crosses its budget while a request is in flight.
|
||||
"""
|
||||
try:
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_team_max_budget_check,
|
||||
_virtual_key_max_budget_check,
|
||||
get_team_object,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
auth: Final = metadata.get("user_api_key_auth")
|
||||
if not isinstance(auth, UserAPIKeyAuth):
|
||||
return False
|
||||
try:
|
||||
await _virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj)
|
||||
if auth.team_id:
|
||||
team: Final = await get_team_object(
|
||||
team_id=auth.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_cache_only=True,
|
||||
)
|
||||
await _team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj)
|
||||
except BudgetExceededError:
|
||||
return True
|
||||
except Exception as e: # noqa: BLE001 # advisory gate: a failed read must not block sampling
|
||||
verbose_logger.debug("shadow_eval: budget read failed: %s", e)
|
||||
return False
|
||||
|
||||
|
||||
def _request_was_routed_by(request_metadata: Mapping[str, object], router_name: str) -> bool:
|
||||
"""Duplicating a request the shadowed router already served compares the router to
|
||||
itself: guaranteed ties, judge spend for zero information."""
|
||||
decision: Final = request_metadata.get("routing_decision")
|
||||
if not isinstance(decision, Mapping):
|
||||
return False
|
||||
return decision.get("router_model_name") == router_name
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CallFailure:
|
||||
"""A shadow or judge call that produced no usable response. cost carries any judge
|
||||
spend the failed attempt still billed, so job-level judge_spend never undercounts."""
|
||||
|
||||
error: str
|
||||
cost: float = 0.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ShadowResponse:
|
||||
"""A successful shadow call, with what the attempt row records."""
|
||||
|
||||
text: str
|
||||
model: str
|
||||
tier: str | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _JudgeVerdict:
|
||||
"""A parsed judge verdict, unmasked back to real/shadow/tie."""
|
||||
|
||||
preference: str
|
||||
confidence: float
|
||||
cost: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActiveShadowEvalJob:
|
||||
"""One active job as the sampling path needs it: immutable config plus the attempt
|
||||
count as of the cache fill (the turn budget's staleness is bounded by the cache TTL)."""
|
||||
|
||||
id: str
|
||||
router_name: str
|
||||
shadow_percentage: float
|
||||
judge_model: str
|
||||
max_turns: int
|
||||
ends_at: datetime
|
||||
attempts: int
|
||||
|
||||
|
||||
def _as_utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value
|
||||
|
||||
|
||||
_jobs_cache: Final = InMemoryCache(max_size_in_memory=4, default_ttl=_JOBS_CACHE_TTL_SECONDS)
|
||||
_JOBS_CACHE_KEY: Final = "shadow_eval:active_jobs"
|
||||
|
||||
|
||||
class ShadowEvalLogger(CustomLogger):
|
||||
"""Fires blind pairwise shadow evaluations for keys with an active shadow-eval job."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
router_provider: Callable[[], "Router | None"] | None = None,
|
||||
prisma_provider: Callable[[], "PrismaClient | None"] | None = None,
|
||||
jobs_cache: InMemoryCache | None = None,
|
||||
) -> None:
|
||||
"""Providers are callables so the proxy's lazily-initialized globals are resolved
|
||||
at call time, not at logger construction."""
|
||||
self._router_provider = router_provider or default_router_provider
|
||||
self._prisma_provider = prisma_provider or _default_prisma_provider
|
||||
self._jobs_cache = jobs_cache or _jobs_cache
|
||||
self._inflight_shadow_tasks: int = 0
|
||||
# Starts per job since the last cache fill, never decremented within a
|
||||
# generation; the refill absorbs written rows and resets.
|
||||
self._job_starts: dict[str, int] = {} # mutable-ok: per-generation counter
|
||||
|
||||
async def _active_jobs(self) -> Mapping[str, ActiveShadowEvalJob]:
|
||||
"""Active jobs by api_key_id, cache-first. A DB fault returns empty without
|
||||
caching, so sampling pauses for that request and the next one retries."""
|
||||
cached: Final = await self._jobs_cache.async_get_cache(_JOBS_CACHE_KEY)
|
||||
if cached is not None:
|
||||
return cached # pyright: ignore[reportReturnType] # cache stores exactly this mapping shape
|
||||
prisma: Final = self._prisma_provider()
|
||||
if prisma is None:
|
||||
return _EMPTY_JOBS
|
||||
try:
|
||||
records: Final = await prisma.db.litellm_shadowevaljob.find_many(
|
||||
where={ # mutable-ok: Prisma filter
|
||||
"stopped_at": None,
|
||||
"ends_at": {"gt": datetime.now(timezone.utc)}, # mutable-ok: Prisma filter
|
||||
},
|
||||
)
|
||||
grouped: Final = (
|
||||
await prisma.db.litellm_shadowevalattempt.group_by(
|
||||
by=["job_id"],
|
||||
count=True,
|
||||
where={"job_id": {"in": [str(record.id) for record in records]}}, # mutable-ok: Prisma filter
|
||||
)
|
||||
if records
|
||||
else ()
|
||||
)
|
||||
attempt_counts: Final = {str(row["job_id"]): int(row["_count"]["_all"]) for row in grouped or []}
|
||||
jobs: Final = {
|
||||
str(record.api_key_id): ActiveShadowEvalJob(
|
||||
id=str(record.id),
|
||||
router_name=str(record.router_name),
|
||||
shadow_percentage=float(record.shadow_percentage),
|
||||
judge_model=str(record.judge_model),
|
||||
max_turns=int(record.max_turns),
|
||||
ends_at=_as_utc(record.ends_at),
|
||||
attempts=attempt_counts.get(str(record.id), 0),
|
||||
)
|
||||
for record in records or []
|
||||
}
|
||||
await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs)
|
||||
self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill
|
||||
return jobs
|
||||
except Exception as e: # noqa: BLE001 # a DB blip must never break request logging
|
||||
verbose_logger.debug("shadow_eval: active-job read failed: %s", e)
|
||||
return _EMPTY_JOBS
|
||||
|
||||
#### hook ####
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
try:
|
||||
payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") # pyright: ignore[reportAssignmentType] # untyped callback kwargs
|
||||
if payload is None:
|
||||
return
|
||||
raw_meta: Final = get_litellm_metadata_from_kwargs(dict(kwargs)) # mutable-ok: helper needs dict
|
||||
request_metadata: Final = raw_meta if isinstance(raw_meta, Mapping) else _EMPTY_METADATA
|
||||
if request_metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY):
|
||||
return # internal sub-call (our own shadow/judge, a classifier), not user traffic
|
||||
# redaction rewrites logged content before callbacks run, so this hook
|
||||
# only ever sees placeholders for a redacted request
|
||||
if should_redact_message_logging(dict(kwargs)): # mutable-ok: predicate takes a plain dict
|
||||
return
|
||||
metadata: Final = payload.get("metadata") or _EMPTY_METADATA
|
||||
api_key_hash: Final = metadata.get("user_api_key_hash")
|
||||
if not api_key_hash:
|
||||
return
|
||||
job: Final = (await self._active_jobs()).get(str(api_key_hash))
|
||||
if job is None:
|
||||
return
|
||||
if datetime.now(timezone.utc) >= job.ends_at:
|
||||
return
|
||||
if job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns:
|
||||
return
|
||||
request_id: Final = payload.get("id") or ""
|
||||
if not request_id:
|
||||
return
|
||||
if not _sample_hits(request_id, job.id, job.shadow_percentage):
|
||||
return
|
||||
if payload.get("call_type") not in _SAMPLED_CALL_TYPES:
|
||||
return # only known chat-shaped traffic is comparable; unknown or missing types fail closed
|
||||
if _request_was_routed_by(request_metadata, job.router_name):
|
||||
return
|
||||
if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS:
|
||||
return
|
||||
raw_messages: Final = kwargs.get("messages")
|
||||
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1
|
||||
self._inflight_shadow_tasks += 1
|
||||
task: Final = asyncio.create_task(
|
||||
self._run_shadow_eval(
|
||||
job=job,
|
||||
request_id=request_id,
|
||||
messages=tuple(m for m in raw_messages if isinstance(m, Mapping))
|
||||
if isinstance(raw_messages, Sequence)
|
||||
else (),
|
||||
response_obj=response_obj,
|
||||
real_model=payload.get("model") or "",
|
||||
model_parameters=MappingProxyType(
|
||||
dict(payload.get("model_parameters") or {}) # mutable-ok: frozen snapshot
|
||||
),
|
||||
parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot
|
||||
)
|
||||
)
|
||||
task.add_done_callback(self._release_shadow_slot)
|
||||
except Exception as e: # noqa: BLE001 # logging hooks must never fail the request
|
||||
verbose_logger.debug("shadow_eval: failed to schedule task: %s", e)
|
||||
|
||||
def _release_shadow_slot(self, _task: "asyncio.Task[None]") -> None:
|
||||
self._inflight_shadow_tasks -= 1
|
||||
|
||||
#### the detached pipeline: one attempt row per sampled request, verdict or error ####
|
||||
|
||||
async def _run_shadow_eval(
|
||||
self,
|
||||
job: ActiveShadowEvalJob,
|
||||
request_id: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
response_obj: object,
|
||||
real_model: str,
|
||||
model_parameters: Mapping[str, object],
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> None:
|
||||
"""Budget gate -> shadow call -> blind judge -> one attempt row. The prisma gate
|
||||
sits above the dispatch so no provider spend happens without a place to record
|
||||
the outcome, and the budget read lives here rather than in the success hook."""
|
||||
prisma: Final = self._prisma_provider()
|
||||
try:
|
||||
if prisma is None:
|
||||
return
|
||||
real_text: Final = self._extract_response_text(response_obj)
|
||||
if not real_text or not messages:
|
||||
return
|
||||
if await _key_or_team_is_over_budget(parent_metadata):
|
||||
return
|
||||
|
||||
shadow: Final = await self._call_router_shadow(job.router_name, messages, model_parameters, parent_metadata)
|
||||
if isinstance(shadow, _CallFailure):
|
||||
await self._record_attempt(prisma, job, request_id, outcome="error", error=shadow.error)
|
||||
return
|
||||
|
||||
verdict: Final = await self._call_judge(
|
||||
judge_model=job.judge_model,
|
||||
messages=messages,
|
||||
real_text=real_text,
|
||||
shadow_text=shadow.text,
|
||||
parent_metadata=parent_metadata,
|
||||
)
|
||||
if isinstance(verdict, _CallFailure):
|
||||
await self._record_attempt(
|
||||
prisma,
|
||||
job,
|
||||
request_id,
|
||||
outcome="error",
|
||||
error=verdict.error,
|
||||
shadow=shadow,
|
||||
judge_cost=verdict.cost,
|
||||
)
|
||||
return
|
||||
await self._record_attempt(
|
||||
prisma,
|
||||
job,
|
||||
request_id,
|
||||
outcome=verdict.preference,
|
||||
shadow=shadow,
|
||||
real_model=real_model,
|
||||
confidence=verdict.confidence,
|
||||
judge_cost=verdict.cost,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # detached task: record what happened, never raise
|
||||
verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e)
|
||||
await self._record_attempt(prisma, job, request_id, outcome="error", error=f"pipeline error: {e}")
|
||||
|
||||
@staticmethod
|
||||
async def _record_attempt(
|
||||
prisma: "PrismaClient | None",
|
||||
job: ActiveShadowEvalJob,
|
||||
request_id: str,
|
||||
*,
|
||||
outcome: str,
|
||||
shadow: _ShadowResponse | None = None,
|
||||
real_model: str = "",
|
||||
confidence: float | None = None,
|
||||
judge_cost: float = 0.0,
|
||||
error: str | None = None,
|
||||
) -> None:
|
||||
if prisma is None:
|
||||
return
|
||||
try:
|
||||
await prisma.db.litellm_shadowevalattempt.create(
|
||||
data={ # mutable-ok: Prisma payload
|
||||
"job_id": job.id,
|
||||
"request_id": request_id,
|
||||
"outcome": outcome,
|
||||
"tier": shadow.tier if shadow else None,
|
||||
"real_model": real_model or None,
|
||||
"shadow_model": shadow.model if shadow else None,
|
||||
"confidence": confidence,
|
||||
"judge_cost": judge_cost,
|
||||
"error": error[:_MAX_ERROR_CHARS] if error else None,
|
||||
}
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a lost row degrades sample size, nothing can disagree with it
|
||||
verbose_logger.debug("shadow_eval: attempt write failed for %s: %s", request_id, e)
|
||||
|
||||
async def _call_router_shadow(
|
||||
self,
|
||||
router_name: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
model_parameters: Mapping[str, object],
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> "_ShadowResponse | _CallFailure":
|
||||
"""Send the prompt through the auto-router being evaluated. The metadata carries
|
||||
the shadowed key's identity (spend attribution) and receives the router's routing
|
||||
decision write-back, read back for tier attribution."""
|
||||
router: Final = self._router_provider()
|
||||
if router is None:
|
||||
return _CallFailure("no router configured on this pod")
|
||||
shadow_metadata: Final[dict[str, object]] = ( # mutable-ok: router writes its routing decision back
|
||||
sanitized_forwardable_call_metadata(parent_metadata, SHADOW_EVAL_ROUTER_CALL_ORIGIN)
|
||||
)
|
||||
shadow_params: Final = { # mutable-ok: splatted as kwargs
|
||||
k: v for k, v in model_parameters.items() if k not in ("stream", "metadata")
|
||||
}
|
||||
try:
|
||||
response: Final = await router.acompletion(
|
||||
model=router_name,
|
||||
messages=messages, # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
|
||||
metadata=shadow_metadata,
|
||||
num_retries=0,
|
||||
fallbacks=[], # mutable-ok: SDK kwarg; a failed shadow is a recorded error, never a spend multiplier
|
||||
**shadow_params,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # provider errors become error rows, not crashes
|
||||
verbose_logger.debug("shadow_eval: router call failed: %s", e)
|
||||
return _CallFailure(f"shadow router call failed: {e}")
|
||||
text: Final = self._extract_response_text(response)
|
||||
if not text:
|
||||
return _CallFailure("shadow router returned an empty response")
|
||||
raw_decision: Final = shadow_metadata.get("routing_decision")
|
||||
routing_decision: Final = raw_decision if isinstance(raw_decision, Mapping) else _EMPTY_METADATA
|
||||
raw_tier: Final = routing_decision.get("tier_label") or routing_decision.get("tier")
|
||||
return _ShadowResponse(
|
||||
text=text,
|
||||
model=str(getattr(response, "model", None) or routing_decision.get("routed_model") or ""),
|
||||
tier=str(raw_tier) if raw_tier is not None else None,
|
||||
)
|
||||
|
||||
async def _call_judge(
|
||||
self,
|
||||
judge_model: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
real_text: str,
|
||||
shadow_text: str,
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> "_JudgeVerdict | _CallFailure":
|
||||
"""Blind pairwise judge with A/B labels randomized to cancel position bias."""
|
||||
real_is_a: Final = random.random() < 0.5
|
||||
response_a: Final = real_text if real_is_a else shadow_text
|
||||
response_b: Final = shadow_text if real_is_a else real_text
|
||||
|
||||
conversation: Final = "\n".join(
|
||||
f"{str(m.get('role', 'user')).upper()}: {extract_text_from_content(m.get('content'))}"
|
||||
for m in messages
|
||||
if m.get("content") is not None
|
||||
)
|
||||
judge_metadata: Final = sanitized_forwardable_call_metadata(parent_metadata, SHADOW_EVAL_JUDGE_CALL_ORIGIN)
|
||||
judge_messages: Final = [ # mutable-ok: SDK takes a list
|
||||
{"role": "system", "content": PAIRWISE_JUDGE_SYSTEM_PROMPT}, # mutable-ok: SDK message
|
||||
{
|
||||
"role": "user",
|
||||
"content": _judge_user_prompt(conversation, response_a, response_b),
|
||||
}, # mutable-ok: SDK message
|
||||
]
|
||||
try:
|
||||
response: Final = await judge_acompletion(
|
||||
self._router_provider(),
|
||||
judge_model,
|
||||
judge_messages, # pyright: ignore[reportArgumentType] # plain SDK message dicts
|
||||
temperature=0,
|
||||
max_tokens=JUDGE_MAX_OUTPUT_TOKENS,
|
||||
metadata=judge_metadata,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # judge outages become error rows, not crashes
|
||||
verbose_logger.debug("shadow_eval: judge call failed: %s", e)
|
||||
return _CallFailure(f"judge call failed: {e}")
|
||||
try:
|
||||
raw: Final = response["choices"][0]["message"]["content"] or ""
|
||||
verdict: Final = PairwiseVerdict.model_validate(parse_json_verdict(raw))
|
||||
except Exception as e: # noqa: BLE001 # malformed verdicts become error rows
|
||||
verbose_logger.debug("shadow_eval: unparseable judge verdict: %s", e)
|
||||
return _CallFailure(f"unparseable judge verdict: {e}", cost=_judge_call_cost(response))
|
||||
return _JudgeVerdict(
|
||||
preference=_unmask_preference(verdict.preference, real_is_a),
|
||||
confidence=max(0.0, min(1.0, verdict.confidence)),
|
||||
cost=_judge_call_cost(response),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_text(response_obj: object) -> str:
|
||||
"""Extract the assistant's text from a ModelResponse-shaped object or dict."""
|
||||
try:
|
||||
content: Final = (
|
||||
response_obj["choices"][0]["message"]["content"]
|
||||
if isinstance(response_obj, Mapping)
|
||||
else response_obj.choices[0].message.content # pyright: ignore[reportAttributeAccessIssue] # duck-typed ModelResponse
|
||||
)
|
||||
except (AttributeError, KeyError, IndexError, TypeError):
|
||||
return ""
|
||||
return extract_text_from_content(content)
|
||||
|
||||
|
||||
_EMPTY_JOBS: Final[Mapping[str, ActiveShadowEvalJob]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _default_prisma_provider() -> "PrismaClient | None":
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
except ImportError:
|
||||
return None
|
||||
return prisma_client
|
||||
|
|
@ -2,8 +2,8 @@
|
|||
Handler for transforming interactions API requests to litellm.responses requests.
|
||||
"""
|
||||
|
||||
from collections.abc import AsyncIterator, Coroutine, Iterator
|
||||
from typing import Any, Final, cast
|
||||
from collections.abc import AsyncIterator, Callable, Coroutine, Iterator
|
||||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm.interactions.litellm_responses_transformation.streaming_iterator import (
|
||||
|
|
@ -37,7 +37,7 @@ class LiteLLMResponsesInteractionsHandler:
|
|||
) -> (
|
||||
InteractionsAPIResponse
|
||||
| Iterator[InteractionsAPIStreamingResponse]
|
||||
| Coroutine[Any, Any, InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]]
|
||||
| Coroutine[object, object, InteractionsAPIResponse | AsyncIterator[InteractionsAPIStreamingResponse]]
|
||||
):
|
||||
"""
|
||||
Handle Interactions API request by calling litellm.responses().
|
||||
|
|
@ -55,13 +55,15 @@ class LiteLLMResponsesInteractionsHandler:
|
|||
InteractionsAPIResponse or streaming iterator
|
||||
"""
|
||||
# Transform interactions request to responses request
|
||||
responses_request = LiteLLMResponsesInteractionsConfig.transform_interactions_request_to_responses_request(
|
||||
model=model,
|
||||
input=input,
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=stream,
|
||||
**kwargs,
|
||||
responses_request: Final = (
|
||||
LiteLLMResponsesInteractionsConfig.transform_interactions_request_to_responses_request(
|
||||
model=model,
|
||||
input=input,
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=stream,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
if _is_async:
|
||||
|
|
@ -76,7 +78,10 @@ class LiteLLMResponsesInteractionsHandler:
|
|||
# Call litellm.responses()
|
||||
# Note: litellm.responses() returns Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]
|
||||
# but the type checker may see it as a coroutine in some contexts
|
||||
responses_response: Final = litellm.responses(
|
||||
responses_fn: Final[Callable[..., ResponsesAPIResponse | BaseResponsesAPIStreamingIterator]] = vars(litellm)[
|
||||
"responses"
|
||||
]
|
||||
responses_response: Final = responses_fn(
|
||||
**responses_request,
|
||||
)
|
||||
|
||||
|
|
@ -92,8 +97,7 @@ class LiteLLMResponsesInteractionsHandler:
|
|||
)
|
||||
|
||||
# At this point, responses_response must be ResponsesAPIResponse (not streaming)
|
||||
# Cast to satisfy type checker since we've already checked it's not a streaming iterator
|
||||
responses_api_response: Final = cast(ResponsesAPIResponse, responses_response)
|
||||
responses_api_response: Final = responses_response
|
||||
|
||||
# Transform responses response to interactions response
|
||||
return LiteLLMResponsesInteractionsConfig.transform_responses_response_to_interactions_response(
|
||||
|
|
@ -112,7 +116,10 @@ class LiteLLMResponsesInteractionsHandler:
|
|||
"""Async handler for interactions API requests."""
|
||||
# Call litellm.aresponses()
|
||||
# Note: litellm.aresponses() returns Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]
|
||||
responses_response: Final = await litellm.aresponses(
|
||||
aresponses_fn: Final[
|
||||
Callable[..., Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator]]
|
||||
] = vars(litellm)["aresponses"]
|
||||
responses_response: Final = await aresponses_fn(
|
||||
**responses_request,
|
||||
)
|
||||
|
||||
|
|
@ -128,8 +135,7 @@ class LiteLLMResponsesInteractionsHandler:
|
|||
)
|
||||
|
||||
# At this point, responses_response must be ResponsesAPIResponse (not streaming)
|
||||
# Cast to satisfy type checker since we've already checked it's not a streaming iterator
|
||||
responses_api_response: Final = cast(ResponsesAPIResponse, responses_response)
|
||||
responses_api_response: Final = responses_response
|
||||
|
||||
# Transform responses response to interactions response
|
||||
return LiteLLMResponsesInteractionsConfig.transform_responses_response_to_interactions_response(
|
||||
|
|
|
|||
94
litellm/litellm_core_utils/internal_call_metadata.py
Normal file
94
litellm/litellm_core_utils/internal_call_metadata.py
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
"""Metadata a request forwards to the internal LLM sub-calls it triggers.
|
||||
|
||||
Internal features (the auto-router's classifier and embeddings, shadow eval's shadow and
|
||||
judge calls) bill real provider spend that nobody typed a prompt for. That spend must land
|
||||
on the same key/team/org/user as the request that caused it, so the sub-call carries the
|
||||
caller's identity metadata, minus two things that must never be forwarded as-is:
|
||||
|
||||
* ``user_api_key_budget_reservation`` (and the reservation nested inside
|
||||
``user_api_key_auth``) belongs to the parent completion. If a sub-call's cost callback
|
||||
sees it, that callback finalizes the reservation and the parent's own callback then
|
||||
skips incrementing the key/team budget counters, losing the parent's spend.
|
||||
``user_api_key_auth`` itself is kept, sanitized, because model access-group filtering
|
||||
needs it.
|
||||
* The sub-call is stamped with ``INTERNAL_CALL_ORIGIN_METADATA_KEY`` so its spend log row
|
||||
records that it is not traffic the caller sent.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.types.utils import InternalCallOrigin
|
||||
|
||||
BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
|
||||
|
||||
_USER_API_KEY_AUTH_KEY: Final = "user_api_key_auth"
|
||||
|
||||
FORWARDABLE_IDENTITY_METADATA_KEYS: Final = frozenset(
|
||||
{
|
||||
"user_api_key",
|
||||
"user_api_key_hash",
|
||||
"user_api_key_alias",
|
||||
"user_api_key_team_id",
|
||||
"user_api_key_org_id",
|
||||
"user_api_key_user_id",
|
||||
"user_api_key_end_user_id",
|
||||
_USER_API_KEY_AUTH_KEY,
|
||||
}
|
||||
)
|
||||
"""The caller-identity subset a detached sub-call needs to be attributed and
|
||||
budget-checked like the request that spawned it. Everything else on the parent's metadata
|
||||
(routing decision, guardrail state, logging payload) describes the parent call and would
|
||||
be a lie on a sub-call that runs after it returned."""
|
||||
|
||||
|
||||
def sanitize_user_api_key_auth(auth: object) -> object:
|
||||
"""Copy of the auth object with its budget reservation removed; the cost callback
|
||||
falls back to reading the reservation from inside the auth object."""
|
||||
if isinstance(auth, dict):
|
||||
return {k: v for k, v in auth.items() if k != "budget_reservation"} # mutable-ok: SDK metadata value
|
||||
reservation: Final[object] = getattr(auth, "budget_reservation", None)
|
||||
model_copy: Final[object] = getattr(auth, "model_copy", None)
|
||||
if reservation is not None and callable(model_copy):
|
||||
return model_copy(update={"budget_reservation": None}) # mutable-ok: pydantic update payload
|
||||
return auth
|
||||
|
||||
|
||||
def _sanitized(parent_metadata: Mapping[str, object]) -> dict[str, object]: # mutable-ok: SDK metadata kwarg
|
||||
return { # mutable-ok: SDK metadata kwarg
|
||||
k: sanitize_user_api_key_auth(v) if k == _USER_API_KEY_AUTH_KEY else v
|
||||
for k, v in parent_metadata.items()
|
||||
if k not in BUDGET_RESERVATION_METADATA_KEYS
|
||||
}
|
||||
|
||||
|
||||
def forwarded_internal_call_metadata(
|
||||
parent_metadata: Mapping[str, object] | None,
|
||||
call_origin: InternalCallOrigin,
|
||||
) -> dict[str, object]: # mutable-ok: SDK metadata kwarg
|
||||
"""Parent metadata, minus its budget reservation, stamped with the sub-call's origin.
|
||||
|
||||
For sub-calls made inside the parent request (classifier, embeddings), where the
|
||||
parent's full context still describes the call being made.
|
||||
"""
|
||||
if not parent_metadata:
|
||||
return {} # mutable-ok: SDK metadata kwarg
|
||||
return _sanitized(parent_metadata) | { # mutable-ok: SDK metadata kwarg
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY: call_origin
|
||||
}
|
||||
|
||||
|
||||
def sanitized_forwardable_call_metadata(
|
||||
parent_metadata: Mapping[str, object],
|
||||
call_origin: InternalCallOrigin,
|
||||
) -> dict[str, object]: # mutable-ok: SDK metadata kwarg
|
||||
"""Just the caller's identity, stamped with the sub-call's origin.
|
||||
|
||||
For sub-calls detached from the parent request (shadow eval), which outlive it and
|
||||
must not inherit per-request state such as its routing decision or logging payload.
|
||||
"""
|
||||
identity: Final = {k: v for k, v in parent_metadata.items() if k in FORWARDABLE_IDENTITY_METADATA_KEYS}
|
||||
return _sanitized(identity) | {INTERNAL_CALL_ORIGIN_METADATA_KEY: call_origin} # mutable-ok: SDK metadata kwarg
|
||||
87
litellm/litellm_core_utils/llm_judge.py
Normal file
87
litellm/litellm_core_utils/llm_judge.py
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
"""Shared primitives for LLM-judge features (llm_as_a_judge guardrail, shadow eval)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import Router
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
JSON_FENCE_RE: Final = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.DOTALL | re.IGNORECASE)
|
||||
|
||||
|
||||
def default_router_provider() -> Router | None:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
return llm_router
|
||||
|
||||
|
||||
def parse_json_verdict(raw: str) -> dict[str, object]: # mutable-ok: plain parsed-JSON payload
|
||||
"""Parse a judge's JSON verdict, tolerating markdown fences and surrounding prose."""
|
||||
text = raw.strip() # rebind-ok: progressively narrowed to the JSON payload
|
||||
fenced: Final = JSON_FENCE_RE.search(text)
|
||||
if fenced is not None:
|
||||
text = fenced.group(1).strip() # rebind-ok: progressively narrowed to the JSON payload
|
||||
parsed: object
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
start: Final = text.find("{")
|
||||
end: Final = text.rfind("}")
|
||||
if start == -1 or end <= start:
|
||||
raise
|
||||
parsed = json.loads(text[start : end + 1])
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("judge response is not a JSON object")
|
||||
return {str(k): v for k, v in parsed.items()} # mutable-ok: plain parsed-JSON payload
|
||||
|
||||
|
||||
def extract_text_from_content(content: object) -> str:
|
||||
"""Return plain text from a message content field (str or multimodal list)."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
return " ".join(
|
||||
str(part.get("text", "")) for part in content if isinstance(part, dict) and part.get("type") == "text"
|
||||
)
|
||||
return ""
|
||||
|
||||
|
||||
def router_resolves_model(router: Router | None, model: str) -> bool:
|
||||
"""Whether the model name resolves through the proxy's router (configured deployment
|
||||
or model-group alias), the same check the judge dispatch itself makes, so start-time
|
||||
validation cannot accept a name the call path then fails on."""
|
||||
return router is not None and bool(model in router.model_group_alias or router.get_model_list(model_name=model))
|
||||
|
||||
|
||||
async def judge_acompletion(
|
||||
router: Router | None,
|
||||
judge_model: str,
|
||||
messages: list[AllMessageValues], # mutable-ok: the SDK acompletion signature takes a list
|
||||
**params: object,
|
||||
) -> ModelResponse:
|
||||
"""Dispatch a judge call through the proxy's router when the judge model is a
|
||||
configured deployment (DB-stored credentials work), through the SDK for
|
||||
provider-qualified public names. The router path never retries or falls back:
|
||||
a failed judge call is the caller's counted failure, not a spend multiplier.
|
||||
Sampling preferences are advisory: models that removed sampling params (e.g.
|
||||
claude-sonnet-5) drop them instead of rejecting the judge call."""
|
||||
if router_resolves_model(router, judge_model):
|
||||
return await router.acompletion( # pyright: ignore[reportOptionalMemberAccess] # router_resolves_model implies router is not None
|
||||
model=judge_model,
|
||||
messages=messages,
|
||||
num_retries=0,
|
||||
fallbacks=[],
|
||||
drop_params=True,
|
||||
**params,
|
||||
)
|
||||
return await litellm.acompletion(model=judge_model, messages=messages, num_retries=0, drop_params=True, **params)
|
||||
|
|
@ -13,8 +13,10 @@ Mirrors Anthropic's native ``compact_20260112`` for non-Anthropic providers:
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NotRequired, Optional, TypedDict, Union, cast
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -29,9 +31,8 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.router import Router
|
||||
from litellm.types.llms.anthropic import (
|
||||
AllAnthropicPassThroughMessageValues,
|
||||
AllAnthropicToolsValues,
|
||||
AnthopicMessagesAssistantMessageParam,
|
||||
AnthropicMessagesUserMessageParam,
|
||||
)
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
|
@ -534,7 +535,7 @@ def _augment_system_with_summary(
|
|||
return [{"type": "text", "text": prefix.rstrip()}, *system]
|
||||
|
||||
|
||||
def _resolve_trigger_tokens(edit_spec: dict[str, object]) -> tuple[int, list[str]]:
|
||||
def _resolve_trigger_tokens(edit_spec: Mapping[str, object]) -> tuple[int, list[str]]:
|
||||
"""Validate and resolve ``trigger.value``.
|
||||
|
||||
Raises ``AnthropicContextManagementError`` if the explicitly-supplied value
|
||||
|
|
@ -568,7 +569,7 @@ def _resolve_trigger_tokens(edit_spec: dict[str, object]) -> tuple[int, list[str
|
|||
return value, warnings
|
||||
|
||||
|
||||
def _build_summary_prompt(edit_spec: dict[str, object], tools: list[dict[str, object]] | None) -> str:
|
||||
def _build_summary_prompt(edit_spec: Mapping[str, object], tools: Sequence[Mapping[str, object]] | None) -> str:
|
||||
custom: Final = edit_spec.get("instructions")
|
||||
if isinstance(custom, str) and custom.strip():
|
||||
return custom
|
||||
|
|
@ -623,7 +624,7 @@ def _count_effective_tokens(
|
|||
try:
|
||||
openai_shape = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=cast(
|
||||
"list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]",
|
||||
"list[AllAnthropicPassThroughMessageValues]",
|
||||
messages_without_compaction,
|
||||
)
|
||||
)
|
||||
|
|
@ -736,7 +737,7 @@ def _extract_summary_text(raw: str | None) -> str | None:
|
|||
|
||||
def _system_to_openai_message(
|
||||
system: str | list[dict[str, Any]] | None,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> dict[str, object] | None:
|
||||
"""Translate Anthropic-shaped ``system`` to an OpenAI system message.
|
||||
|
||||
Accepts a bare string or a list of Anthropic content blocks; returns
|
||||
|
|
@ -773,7 +774,7 @@ def _build_summary_messages(
|
|||
try:
|
||||
openai_messages = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
|
||||
messages=cast(
|
||||
"list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]",
|
||||
"list[AllAnthropicPassThroughMessageValues]",
|
||||
stripped,
|
||||
)
|
||||
)
|
||||
|
|
@ -809,7 +810,7 @@ def _is_user_message(msg: object) -> bool:
|
|||
return isinstance(msg, dict) and msg.get("role") == "user"
|
||||
|
||||
|
||||
def _append_text_to_content(content: Any, extra_text: str) -> Any:
|
||||
def _append_text_to_content(content: object, extra_text: str) -> object:
|
||||
"""Append ``extra_text`` to an OpenAI-shape message ``content`` field.
|
||||
|
||||
Handles the two common shapes: ``str`` and ``list`` of content parts.
|
||||
|
|
@ -820,10 +821,29 @@ def _append_text_to_content(content: Any, extra_text: str) -> Any:
|
|||
if isinstance(content, str):
|
||||
return f"{content}\n\n{extra_text}"
|
||||
if isinstance(content, list):
|
||||
return [*content, {"type": "text", "text": extra_text}]
|
||||
appended: Final[list[object]] = [*content, {"type": "text", "text": extra_text}]
|
||||
return appended
|
||||
return [content, {"type": "text", "text": extra_text}]
|
||||
|
||||
|
||||
class _SummaryCallUserKwarg(TypedDict, total=False):
|
||||
user: ReadOnly[object]
|
||||
|
||||
|
||||
class _SummaryCallRegionKwarg(TypedDict, total=False):
|
||||
allowed_model_region: ReadOnly[str]
|
||||
|
||||
|
||||
class _SummaryCallKwargs(TypedDict):
|
||||
model: ReadOnly[str]
|
||||
messages: ReadOnly[list[dict[str, object]]]
|
||||
max_tokens: ReadOnly[int]
|
||||
timeout: ReadOnly[float]
|
||||
litellm_metadata: ReadOnly[Mapping[str, object]]
|
||||
user: NotRequired[ReadOnly[object]]
|
||||
allowed_model_region: NotRequired[ReadOnly[str]]
|
||||
|
||||
|
||||
async def _call_summary_model(
|
||||
*,
|
||||
summary_model: str,
|
||||
|
|
@ -860,22 +880,24 @@ async def _call_summary_model(
|
|||
# the parent ``/v1/messages`` request. On timeout the caller catches the
|
||||
# exception and surfaces ``applied_edits[0].error = "summary_call_failed"``,
|
||||
# forwarding the request without compaction rather than hanging.
|
||||
call_kwargs: Final[dict[str, Any]] = {
|
||||
"model": summary_model,
|
||||
"messages": summary_messages,
|
||||
"max_tokens": max_tokens,
|
||||
"timeout": COMPACT_SUMMARY_TIMEOUT_SECONDS,
|
||||
"litellm_metadata": metadata,
|
||||
}
|
||||
# The end-user id must also travel as the top-level ``user`` kwarg: legacy
|
||||
# limiter hooks and prometheus end-user tracking read it from there rather
|
||||
# than from ``litellm_metadata``, so without it the summary tokens would not
|
||||
# debit the caller's end-user counters.
|
||||
end_user_id: Final = metadata.get("user_api_key_end_user_id")
|
||||
if end_user_id:
|
||||
call_kwargs["user"] = end_user_id
|
||||
if allowed_model_region is not None:
|
||||
call_kwargs["allowed_model_region"] = allowed_model_region
|
||||
call_kwargs: Final[_SummaryCallKwargs] = {
|
||||
"model": summary_model,
|
||||
"messages": summary_messages,
|
||||
"max_tokens": max_tokens,
|
||||
"timeout": COMPACT_SUMMARY_TIMEOUT_SECONDS,
|
||||
"litellm_metadata": metadata,
|
||||
**(_SummaryCallUserKwarg(user=end_user_id) if end_user_id else _SummaryCallUserKwarg()),
|
||||
**(
|
||||
_SummaryCallRegionKwarg(allowed_model_region=allowed_model_region)
|
||||
if allowed_model_region is not None
|
||||
else _SummaryCallRegionKwarg()
|
||||
),
|
||||
}
|
||||
if llm_router is not None and hasattr(llm_router, "acompletion"):
|
||||
return await llm_router.acompletion(**call_kwargs)
|
||||
return await litellm.acompletion(**call_kwargs)
|
||||
|
|
|
|||
|
|
@ -2,11 +2,12 @@ import asyncio
|
|||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Any, Final, Literal, NamedTuple, cast
|
||||
|
||||
import httpx
|
||||
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -23,6 +24,22 @@ from litellm.utils import _add_path_to_api_base
|
|||
azure_ad_cache: Final = DualCache()
|
||||
|
||||
|
||||
class _AzureAdTokenJson(TypedDict, total=False):
|
||||
access_token: ReadOnly[str]
|
||||
expires_in: ReadOnly[int]
|
||||
|
||||
|
||||
class _AzureV1ClientParams(TypedDict, total=False, extra_items=object):
|
||||
base_url: ReadOnly[str]
|
||||
|
||||
|
||||
class _AzureGatewayClientParams(TypedDict, total=False, extra_items=object):
|
||||
api_version: ReadOnly[str]
|
||||
base_url: ReadOnly[str]
|
||||
max_retries: ReadOnly[int]
|
||||
timeout: ReadOnly[float | httpx.Timeout]
|
||||
|
||||
|
||||
class AzureOpenAIError(BaseLLMException):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -220,7 +237,7 @@ def get_azure_ad_token_from_oidc(
|
|||
message=req_token.text,
|
||||
)
|
||||
|
||||
azure_ad_token_json: Final = req_token.json()
|
||||
azure_ad_token_json: Final[_AzureAdTokenJson] = req_token.json()
|
||||
azure_ad_token_access_token = azure_ad_token_json.get("access_token", None)
|
||||
azure_ad_token_expires_in: Final = azure_ad_token_json.get("expires_in", None)
|
||||
|
||||
|
|
@ -486,7 +503,7 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
|
||||
v1_api_key = _async_v1_api_key
|
||||
|
||||
v1_params: Final[dict[str, Any]] = {
|
||||
v1_params: Final[_AzureV1ClientParams] = {
|
||||
"api_key": v1_api_key,
|
||||
"base_url": f"{api_base}/openai/v1/",
|
||||
}
|
||||
|
|
@ -643,7 +660,7 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
api_base += "/"
|
||||
api_base += f"{model}"
|
||||
|
||||
azure_client_params: Final[dict[str, Any]] = {
|
||||
azure_client_params: Final[_AzureGatewayClientParams] = {
|
||||
"api_version": api_version,
|
||||
"base_url": f"{api_base}",
|
||||
"http_client": litellm.client_session,
|
||||
|
|
@ -702,7 +719,7 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
@staticmethod
|
||||
def _get_base_azure_url(
|
||||
api_base: str | None,
|
||||
litellm_params: GenericLiteLLMParams | dict[str, Any] | None,
|
||||
litellm_params: GenericLiteLLMParams | Mapping[str, object] | None,
|
||||
route: Literal["/openai/responses", "/openai/vector_stores"] | str,
|
||||
default_api_version: str | Literal["latest", "preview"] | None = None,
|
||||
) -> str:
|
||||
|
|
@ -757,7 +774,9 @@ class BaseAzureLLM(BaseOpenAILLM):
|
|||
return False
|
||||
return api_version in {"preview", "latest", "v1"}
|
||||
|
||||
def _resolve_env_var(self, litellm_params: dict[str, Any], param_key: str, env_var_key: str) -> str | None:
|
||||
def _resolve_env_var(
|
||||
self, litellm_params: Mapping[str, str | None], param_key: str, env_var_key: str
|
||||
) -> str | None:
|
||||
"""Resolve the environment variable for a given parameter key.
|
||||
|
||||
The logic here is different from `params.get(key, os.getenv(env_var))` because
|
||||
|
|
|
|||
|
|
@ -2,17 +2,18 @@ import base64
|
|||
import json
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Iterable, Mapping, MutableMapping
|
||||
from collections.abc import Iterable, Mapping, MutableMapping, Sequence
|
||||
from functools import cache
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, TypeAlias, TypedDict
|
||||
from urllib.parse import unquote
|
||||
|
||||
import httpx
|
||||
from httpx import Headers, Response
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -63,10 +64,39 @@ from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resol
|
|||
S3_SIGNED_GET_HEADERS_PARAM: Final = "_s3_signed_get_headers"
|
||||
|
||||
|
||||
def _frozen_mapping(items: Iterable[tuple[str, Any]]) -> Mapping[str, Any]:
|
||||
def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]:
|
||||
return MappingProxyType(dict(items))
|
||||
|
||||
|
||||
_EmbeddingBatchInput: TypeAlias = (
|
||||
str | int | float | Sequence[str] | Sequence[int] | Sequence[Sequence[int]] | Mapping[str, object]
|
||||
)
|
||||
|
||||
|
||||
class _OpenAIBatchRecordBody(TypedDict, total=False):
|
||||
model: ReadOnly[str]
|
||||
prompt: ReadOnly[str | Sequence[str] | Sequence[int] | Sequence[Sequence[int]]]
|
||||
input: ReadOnly[_EmbeddingBatchInput]
|
||||
metadata: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _OpenAIBatchRecord(TypedDict, total=False):
|
||||
custom_id: ReadOnly[str]
|
||||
url: ReadOnly[str]
|
||||
body: ReadOnly[_OpenAIBatchRecordBody]
|
||||
|
||||
|
||||
class _BedrockBatchRecord(TypedDict):
|
||||
recordId: ReadOnly[str]
|
||||
modelInput: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _S3UploadResponse(TypedDict, total=False):
|
||||
Key: ReadOnly[str]
|
||||
Bucket: ReadOnly[str]
|
||||
ContentLength: ReadOnly[int]
|
||||
|
||||
|
||||
# JSONL batch records are untyped json, so the `/v1/responses` fields are
|
||||
# validated into their concrete Responses API types before being handed to the
|
||||
# Responses-to-Chat bridge. Both adapters drop keys the Responses API doesn't
|
||||
|
|
@ -231,7 +261,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
def _get_s3_object_name_from_batch_jsonl(
|
||||
self,
|
||||
openai_jsonl_content: list[dict[str, Any]],
|
||||
openai_jsonl_content: Sequence[_OpenAIBatchRecord],
|
||||
) -> str:
|
||||
"""
|
||||
Gets a unique S3 object name for the Bedrock batch processing job
|
||||
|
|
@ -341,7 +371,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
OPENAI_RESPONSES_URL = "/v1/responses"
|
||||
|
||||
@staticmethod
|
||||
def _classify_batch_record(openai_jsonl_record: Mapping[str, Any]) -> BedrockBatchRecordKind:
|
||||
def _classify_batch_record(openai_jsonl_record: _OpenAIBatchRecord) -> BedrockBatchRecordKind:
|
||||
"""
|
||||
Decide which OpenAI endpoint shape an OpenAI batch JSONL line carries.
|
||||
|
||||
|
|
@ -484,7 +514,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
return value if isinstance(value, str) and value else None
|
||||
|
||||
@staticmethod
|
||||
def _coerce_embedding_input_to_string(raw_input: Any, model: str = "") -> str:
|
||||
def _coerce_embedding_input_to_string(raw_input: _EmbeddingBatchInput | None, model: str = "") -> str:
|
||||
"""
|
||||
Normalize an OpenAI /v1/embeddings `input` field into the single
|
||||
string that Bedrock Titan v2 InvokeModel expects in `inputText`.
|
||||
|
|
@ -541,8 +571,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
def _map_openai_embedding_to_bedrock_params(
|
||||
self,
|
||||
openai_request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
openai_request_body: _OpenAIBatchRecordBody,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Transform an OpenAI /v1/embeddings request body into the
|
||||
Bedrock InvokeModel `modelInput` for embedding models that AWS
|
||||
|
|
@ -588,7 +618,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
return dict(titan_config._transform_request(input=input_text, inference_params=inference_params))
|
||||
|
||||
@staticmethod
|
||||
def _transform_text_completion_body_to_chat_body(openai_request_body: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
def _transform_text_completion_body_to_chat_body(
|
||||
openai_request_body: _OpenAIBatchRecordBody,
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
Rewrite an OpenAI `/v1/completions` batch body as a Chat Completions body.
|
||||
|
||||
|
|
@ -610,7 +642,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _transform_responses_body_to_chat_body(openai_request_body: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
def _transform_responses_body_to_chat_body(openai_request_body: _OpenAIBatchRecordBody) -> Mapping[str, object]:
|
||||
"""
|
||||
Rewrite an OpenAI `/v1/responses` batch body as a Chat Completions body.
|
||||
|
||||
|
|
@ -631,23 +663,25 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
"Batch record for /v1/responses is missing required `input` field: "
|
||||
f"model={openai_request_body.get('model', '')}"
|
||||
)
|
||||
chat_body: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
|
||||
model=openai_request_body.get("model", ""),
|
||||
input=_responses_input_adapter().validate_python(responses_input),
|
||||
responses_api_request=_responses_request_adapter().validate_python(
|
||||
_frozen_mapping(
|
||||
(key, value) for key, value in openai_request_body.items() if key not in ("model", "input")
|
||||
)
|
||||
),
|
||||
metadata=openai_request_body.get("metadata"),
|
||||
chat_body: Final[Mapping[str, object]] = (
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
|
||||
model=openai_request_body.get("model", ""),
|
||||
input=_responses_input_adapter().validate_python(responses_input),
|
||||
responses_api_request=_responses_request_adapter().validate_python(
|
||||
_frozen_mapping(
|
||||
(key, value) for key, value in openai_request_body.items() if key not in ("model", "input")
|
||||
)
|
||||
),
|
||||
metadata=openai_request_body.get("metadata"),
|
||||
)
|
||||
)
|
||||
return _frozen_mapping((key, value) for key, value in chat_body.items() if key != "tools" or value)
|
||||
|
||||
@staticmethod
|
||||
def _transform_batch_body_to_chat_body(
|
||||
openai_request_body: Mapping[str, Any],
|
||||
openai_request_body: _OpenAIBatchRecordBody,
|
||||
record_kind: BedrockBatchRecordKind,
|
||||
) -> Mapping[str, Any]:
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
Normalize a non-embedding batch body to the Chat Completions shape the
|
||||
per-provider Bedrock transformations expect.
|
||||
|
|
@ -666,7 +700,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
self,
|
||||
openai_request_body: Mapping[str, Any],
|
||||
provider: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Transform OpenAI request body to Bedrock-compatible modelInput
|
||||
parameters using existing transformation logic.
|
||||
|
|
@ -677,7 +711,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
_model: Final = openai_request_body.get("model", "")
|
||||
_model: Final[str] = openai_request_body.get("model", "")
|
||||
messages: Final = openai_request_body.get("messages", [])
|
||||
optional_params: Final = {k: v for k, v in openai_request_body.items() if k not in ["model", "messages"]}
|
||||
|
||||
|
|
@ -733,8 +767,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
}
|
||||
|
||||
def _transform_openai_jsonl_content_to_bedrock_jsonl_content(
|
||||
self, openai_jsonl_content: list[dict[str, Any]]
|
||||
) -> list[dict[str, Any]]:
|
||||
self, openai_jsonl_content: Sequence[_OpenAIBatchRecord]
|
||||
) -> list[_BedrockBatchRecord]:
|
||||
"""
|
||||
Transforms OpenAI JSONL content to Bedrock batch format
|
||||
|
||||
|
|
@ -1026,7 +1060,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
response_headers: Final = raw_response.headers
|
||||
# Extract S3 object information from the response
|
||||
# S3 PUT object returns ETag and other metadata in headers
|
||||
content_length: Final = response_headers.get("Content-Length", "0")
|
||||
content_length: Final[str] = response_headers.get("Content-Length", "0")
|
||||
|
||||
# Use the actual upload URL that was used for the S3 upload
|
||||
upload_url: Final = litellm_params.get("upload_url")
|
||||
|
|
@ -1224,7 +1258,9 @@ class BedrockJsonlFilesTransformation:
|
|||
object_name: Final = self._get_s3_object_name(openai_jsonl_content=openai_jsonl_content)
|
||||
return bedrock_jsonl_string, object_name
|
||||
|
||||
def _transform_openai_jsonl_content_to_bedrock_jsonl_content(self, openai_jsonl_content: list[dict[str, Any]]):
|
||||
def _transform_openai_jsonl_content_to_bedrock_jsonl_content(
|
||||
self, openai_jsonl_content: Sequence[_OpenAIBatchRecord]
|
||||
):
|
||||
"""
|
||||
Delegate to the main BedrockFilesConfig transformation method
|
||||
"""
|
||||
|
|
@ -1233,7 +1269,7 @@ class BedrockJsonlFilesTransformation:
|
|||
|
||||
def _get_s3_object_name(
|
||||
self,
|
||||
openai_jsonl_content: list[dict[str, Any]],
|
||||
openai_jsonl_content: Sequence[_OpenAIBatchRecord],
|
||||
) -> str:
|
||||
"""
|
||||
Gets a unique S3 object name for the Bedrock batch processing job
|
||||
|
|
@ -1285,7 +1321,7 @@ class BedrockJsonlFilesTransformation:
|
|||
return content
|
||||
|
||||
def transform_s3_bucket_response_to_openai_file_object(
|
||||
self, create_file_data: CreateFileRequest, s3_upload_response: dict[str, Any]
|
||||
self, create_file_data: CreateFileRequest, s3_upload_response: _S3UploadResponse
|
||||
) -> OpenAIFileObject:
|
||||
"""
|
||||
Transforms S3 Bucket upload file response to OpenAI FileObject
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.constants import RUNWAYML_DEFAULT_API_VERSION
|
||||
|
|
@ -31,6 +33,29 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class _RunwayTaskResponse(TypedDict, total=False):
|
||||
id: ReadOnly[str]
|
||||
status: ReadOnly[str]
|
||||
createdAt: ReadOnly[str]
|
||||
completedAt: ReadOnly[str]
|
||||
output: ReadOnly[Sequence[str] | str]
|
||||
failureCode: ReadOnly[str]
|
||||
failure: ReadOnly[str]
|
||||
progress: ReadOnly[int]
|
||||
|
||||
|
||||
class _VideoObjectData(TypedDict, extra_items=object):
|
||||
id: ReadOnly[str]
|
||||
object: ReadOnly[Literal["video"]]
|
||||
status: ReadOnly[str]
|
||||
created_at: ReadOnly[int]
|
||||
|
||||
|
||||
def _parse_runway_task_response(raw_response: httpx.Response) -> _RunwayTaskResponse:
|
||||
response_data: Final[_RunwayTaskResponse] = raw_response.json()
|
||||
return response_data
|
||||
|
||||
|
||||
class RunwayMLVideoConfig(BaseVideoConfig):
|
||||
"""
|
||||
Configuration class for RunwayML video generation.
|
||||
|
|
@ -78,7 +103,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
- size -> ratio (convert "WIDTHxHEIGHT" to "WIDTH:HEIGHT")
|
||||
- seconds -> duration (convert to integer)
|
||||
"""
|
||||
mapped_params: Final[dict[str, Any]] = {}
|
||||
mapped_params: Final[dict[str, object]] = {}
|
||||
|
||||
# Handle input_reference parameter - map to promptImage
|
||||
if "input_reference" in video_create_optional_params:
|
||||
|
|
@ -180,7 +205,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
}
|
||||
"""
|
||||
# Build the request data
|
||||
request_data: Final[dict[str, Any]] = {
|
||||
request_data: Final[dict[str, object]] = {
|
||||
"model": model,
|
||||
"promptText": prompt,
|
||||
}
|
||||
|
|
@ -189,7 +214,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
request_data.update(video_create_optional_request_params)
|
||||
|
||||
# RunwayML uses JSON body, no files multipart
|
||||
files_list: Final[list[tuple[str, Any]]] = []
|
||||
files_list: Final[RequestFiles] = []
|
||||
|
||||
# Append the specific endpoint for video generation
|
||||
full_api_base: Final = f"{api_base}/image_to_video"
|
||||
|
|
@ -216,10 +241,10 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
|
||||
We map this to OpenAI VideoObject format.
|
||||
"""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_runway_task_response(raw_response)
|
||||
|
||||
# Map RunwayML task response to VideoObject format
|
||||
video_data: Final[dict[str, Any]] = {
|
||||
video_data: Final[_VideoObjectData] = {
|
||||
"id": response_data.get("id", ""),
|
||||
"object": "video",
|
||||
"status": self._map_runway_status(response_data.get("status", "pending")),
|
||||
|
|
@ -326,7 +351,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
# Get task status to retrieve video URL
|
||||
url: Final = f"{api_base}/tasks/{encoded_video_id}"
|
||||
|
||||
params: Final[dict[str, Any]] = {}
|
||||
params: Final[dict[str, str]] = {}
|
||||
|
||||
return url, params
|
||||
|
||||
|
|
@ -421,7 +446,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
Transform the video remix request for RunwayML API.
|
||||
|
|
@ -448,7 +473,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
order: str | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_query: Mapping[str, object] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
Transform the video list request for RunwayML API.
|
||||
|
|
@ -484,7 +509,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
# Construct the URL for task cancellation
|
||||
url: Final = f"{api_base}/tasks/{encoded_video_id}/cancel"
|
||||
|
||||
data: Final[dict[str, Any]] = {}
|
||||
data: Final[dict[str, str]] = {}
|
||||
|
||||
return url, data
|
||||
|
||||
|
|
@ -494,7 +519,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> VideoObject:
|
||||
"""Transform the RunwayML video delete/cancel response."""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_runway_task_response(raw_response)
|
||||
|
||||
video_obj: Final = VideoObject(
|
||||
id=response_data.get("id", ""),
|
||||
|
|
@ -524,7 +549,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
url: Final = f"{api_base}/tasks/{encoded_video_id}"
|
||||
|
||||
# Empty dict for GET request (no body)
|
||||
data: Final[dict[str, Any]] = {}
|
||||
data: Final[dict[str, str]] = {}
|
||||
|
||||
return url, data
|
||||
|
||||
|
|
@ -537,10 +562,10 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
"""
|
||||
Transform the RunwayML video status retrieve response.
|
||||
"""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_runway_task_response(raw_response)
|
||||
|
||||
# Map RunwayML task response to VideoObject format
|
||||
video_data: Final[dict[str, Any]] = {
|
||||
video_data: Final[_VideoObjectData] = {
|
||||
"id": response_data.get("id", ""),
|
||||
"object": "video",
|
||||
"status": self._map_runway_status(response_data.get("status", "pending")),
|
||||
|
|
@ -572,7 +597,7 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
|
||||
return video_obj
|
||||
|
||||
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
|
||||
def transform_video_create_character_request(self, name, video: object, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video create character is not supported for RunwayML")
|
||||
|
||||
def transform_video_create_character_response(self, raw_response, logging_obj):
|
||||
|
|
|
|||
|
|
@ -5,12 +5,13 @@ import json
|
|||
import os
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Iterator
|
||||
from typing import Any, Final
|
||||
from collections.abc import Callable, Iterable, Iterator, Mapping
|
||||
from typing import Any, Final, TypedDict
|
||||
|
||||
import httpx
|
||||
from httpx import Headers, Response
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -50,6 +51,7 @@ from litellm.types.llms.openai import (
|
|||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
OpenAIFileObject,
|
||||
OpenAIFilesPurpose,
|
||||
PathLike,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai import GcsBucketResponse
|
||||
|
|
@ -62,6 +64,46 @@ _GCP_LABEL_VALUE_MAX_LEN: Final = 63
|
|||
_CUSTOM_ID_RAW_LABEL_PREFIX: Final = "b32_"
|
||||
|
||||
|
||||
class _GcsObjectMetadataJson(TypedDict, total=False):
|
||||
purpose: ReadOnly[OpenAIFilesPurpose]
|
||||
|
||||
|
||||
class _GcsObjectJson(TypedDict, total=False):
|
||||
id: ReadOnly[str]
|
||||
name: ReadOnly[str]
|
||||
size: ReadOnly[str]
|
||||
timeCreated: ReadOnly[str]
|
||||
metadata: ReadOnly[_GcsObjectMetadataJson]
|
||||
|
||||
|
||||
class _VertexBatchRowRequest(TypedDict, total=False):
|
||||
labels: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _VertexBatchRow(TypedDict, total=False):
|
||||
request: ReadOnly[_VertexBatchRowRequest]
|
||||
status: ReadOnly[str]
|
||||
processed_time: ReadOnly[str]
|
||||
|
||||
|
||||
class _OpenAIBatchOutputError(TypedDict):
|
||||
code: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
||||
|
||||
class _OpenAIBatchOutputResponse(TypedDict):
|
||||
status_code: ReadOnly[int]
|
||||
request_id: ReadOnly[str]
|
||||
body: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _OpenAIBatchOutputRow(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
custom_id: ReadOnly[str]
|
||||
response: ReadOnly[_OpenAIBatchOutputResponse | None]
|
||||
error: ReadOnly[_OpenAIBatchOutputError | None]
|
||||
|
||||
|
||||
def _sanitize_gcp_label_value(value: str) -> str:
|
||||
"""
|
||||
Sanitize a string to meet GCP label value constraints.
|
||||
|
|
@ -106,7 +148,7 @@ def _decode_gcp_label_value_chunks(values: list[str]) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: Any) -> None:
|
||||
def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: object) -> None:
|
||||
"""
|
||||
Store OpenAI batch custom_id for Vertex batch correlation.
|
||||
|
||||
|
|
@ -122,7 +164,7 @@ def _set_litellm_batch_custom_id_labels(labels: dict[str, str], custom_id: Any)
|
|||
labels[f"litellm_custom_id_raw_{index}"] = raw_label_chunk
|
||||
|
||||
|
||||
def _get_litellm_batch_custom_id_from_labels(labels: dict[str, Any]) -> str:
|
||||
def _get_litellm_batch_custom_id_from_labels(labels: Mapping[str, object]) -> str:
|
||||
"""Prefer encoded custom_id when present (see _set_litellm_batch_custom_id_labels)."""
|
||||
raw: Final = labels.get("litellm_custom_id_raw")
|
||||
if raw:
|
||||
|
|
@ -186,7 +228,7 @@ def _iter_openai_jsonl_lines(openai_file_content: FileTypes) -> Iterator[str]:
|
|||
``str.splitlines()`` + ``line.strip()`` for ``\\n`` / ``\\r\\n`` delimited
|
||||
JSONL.
|
||||
"""
|
||||
content: Any = openai_file_content
|
||||
content: FileTypes | str = openai_file_content
|
||||
if isinstance(content, tuple):
|
||||
content = content[1]
|
||||
|
||||
|
|
@ -246,6 +288,11 @@ def _iter_openai_jsonl_entries(
|
|||
yield json.loads(line)
|
||||
|
||||
|
||||
def _parse_vertex_batch_output_row(line: str) -> _VertexBatchRow:
|
||||
row: Final[_VertexBatchRow] = json.loads(line)
|
||||
return row
|
||||
|
||||
|
||||
class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
|
||||
"""Streams an OpenAI batch JSONL upload as Vertex-wrapped JSONL one row at a
|
||||
time, so the transformed payload is never held in full.
|
||||
|
|
@ -463,7 +510,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
"""
|
||||
Transform VertexAI File upload response into OpenAI-style FileObject
|
||||
"""
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final[GcsBucketResponse] = raw_response.json()
|
||||
|
||||
try:
|
||||
response_object: Final = GcsBucketResponse(**response_json)
|
||||
|
|
@ -523,7 +570,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
) -> OpenAIFileObject:
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final[_GcsObjectJson] = raw_response.json()
|
||||
gcs_id = response_json.get("id", "")
|
||||
gcs_id = "/".join(gcs_id.split("/")[:-1]) if gcs_id else ""
|
||||
return OpenAIFileObject(
|
||||
|
|
@ -682,7 +729,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
# discriminating fields. Anything else (e.g. a binary file whose
|
||||
# first line is not valid UTF-8/JSON) raises and falls through to the
|
||||
# passthrough below, leaving the content untouched.
|
||||
first_row: Final = json.loads(first_line)
|
||||
first_row: Final = _parse_vertex_batch_output_row(first_line)
|
||||
is_vertex_batch_output: Final = (
|
||||
"request" in first_row
|
||||
and "response" in first_row
|
||||
|
|
@ -723,7 +770,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
for line in itertools.chain([first_line], lines):
|
||||
try:
|
||||
openai_output = self._transform_single_vertex_batch_output_to_openai(
|
||||
vertex_output=json.loads(line),
|
||||
vertex_output=_parse_vertex_batch_output_row(line),
|
||||
vertex_gemini_config=vertex_gemini_config,
|
||||
logging_obj=batch_transform_logging_obj,
|
||||
mock_httpx_response=mock_httpx_response,
|
||||
|
|
@ -742,18 +789,18 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
|
||||
def _transform_single_vertex_batch_output_to_openai(
|
||||
self,
|
||||
vertex_output: dict[str, Any],
|
||||
vertex_output: _VertexBatchRow,
|
||||
vertex_gemini_config: VertexGeminiConfig,
|
||||
logging_obj: Logging,
|
||||
mock_httpx_response: httpx.Response,
|
||||
) -> dict[str, Any]:
|
||||
) -> _OpenAIBatchOutputRow:
|
||||
"""
|
||||
Transform a single Vertex AI batch output line to OpenAI format.
|
||||
Uses the existing VertexGeminiConfig transformation for the response.
|
||||
"""
|
||||
# Extract custom_id from request labels (prefer raw for OpenAI round-trip)
|
||||
request_data: Final = vertex_output.get("request", {})
|
||||
labels: Final = request_data.get("labels", {}) or {}
|
||||
labels: Final[Mapping[str, object]] = request_data.get("labels", {}) or {}
|
||||
custom_id: Final = _get_litellm_batch_custom_id_from_labels(labels)
|
||||
|
||||
# Check if there's an error
|
||||
|
|
|
|||
|
|
@ -7,10 +7,12 @@ Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-refer
|
|||
|
||||
import base64
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
|
|
@ -40,11 +42,37 @@ else:
|
|||
BaseLLMException = Any
|
||||
|
||||
|
||||
class _VeoVideo(TypedDict, total=False):
|
||||
gcsUri: ReadOnly[str]
|
||||
bytesBase64Encoded: ReadOnly[str]
|
||||
mimeType: ReadOnly[str]
|
||||
|
||||
|
||||
class _VeoOperationResponse(TypedDict, total=False):
|
||||
videos: ReadOnly[Sequence[_VeoVideo]]
|
||||
|
||||
|
||||
class _VeoOperationMetadata(TypedDict, total=False):
|
||||
createTime: ReadOnly[str]
|
||||
|
||||
|
||||
class _VeoOperation(TypedDict, total=False):
|
||||
name: ReadOnly[str]
|
||||
done: ReadOnly[bool]
|
||||
metadata: ReadOnly[_VeoOperationMetadata]
|
||||
response: ReadOnly[_VeoOperationResponse]
|
||||
|
||||
|
||||
def _parse_veo_operation(raw_response: httpx.Response) -> _VeoOperation:
|
||||
operation: Final[_VeoOperation] = raw_response.json()
|
||||
return operation
|
||||
|
||||
|
||||
def _build_vertex_video_usage_from_request_data(
|
||||
request_data: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, float | str]:
|
||||
"""Build usage metadata (duration, resolution) for video cost calculation."""
|
||||
usage_data: Final[dict[str, Any]] = {}
|
||||
usage_data: Final[dict[str, float | str]] = {}
|
||||
if not request_data:
|
||||
return usage_data
|
||||
|
||||
|
|
@ -125,7 +153,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
video_create_optional_params: VideoCreateOptionalRequestParams,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Map OpenAI-style parameters to Veo format.
|
||||
|
||||
|
|
@ -135,7 +163,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
- size → aspectRatio (e.g., "1280x720" → "16:9")
|
||||
- seconds → durationSeconds (defaults to 4 seconds if not provided)
|
||||
"""
|
||||
mapped_params: Final[dict[str, Any]] = {}
|
||||
mapped_params: Final[dict[str, object]] = {}
|
||||
|
||||
# Map input_reference to image (will be processed in transform_video_create_request)
|
||||
if "input_reference" in video_create_optional_params:
|
||||
|
|
@ -289,7 +317,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
}
|
||||
"""
|
||||
# Build instance with prompt
|
||||
instance_dict: Final[dict[str, Any]] = {"prompt": prompt}
|
||||
instance_dict: Final[dict[str, object]] = {"prompt": prompt}
|
||||
params_copy: Final = video_create_optional_request_params.copy()
|
||||
|
||||
# Check if user wants to provide full instance dict
|
||||
|
|
@ -324,13 +352,13 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
# {"parameters": {"parameters": {...}}} ← wrong
|
||||
# {"parameters": {...}} ← correct
|
||||
nested_params: Final = params_copy.pop("parameters", None)
|
||||
vertex_params: Final[dict[str, Any]] = {}
|
||||
vertex_params: Final[dict[str, object]] = {}
|
||||
if isinstance(nested_params, dict):
|
||||
vertex_params.update(nested_params)
|
||||
vertex_params.update(params_copy)
|
||||
|
||||
# Build request data directly (TypedDict doesn't have model_dump)
|
||||
request_data: Final[dict[str, Any]] = {"instances": [instance_dict]}
|
||||
request_data: Final[dict[str, object]] = {"instances": [instance_dict]}
|
||||
|
||||
# Only add parameters if there are any
|
||||
if vertex_params:
|
||||
|
|
@ -363,7 +391,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
- status: "processing"
|
||||
- usage: includes duration_seconds and optional video_resolution for cost calculation
|
||||
"""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_veo_operation(raw_response)
|
||||
|
||||
operation_name: Final = response_data.get("name")
|
||||
if not operation_name:
|
||||
|
|
@ -441,7 +469,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
}
|
||||
}
|
||||
"""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_veo_operation(raw_response)
|
||||
|
||||
operation_name: Final = response_data.get("name", "")
|
||||
is_done: Final = response_data.get("done", False)
|
||||
|
|
@ -513,7 +541,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
|
||||
Extracts the base64 encoded video from the response and decodes it to bytes.
|
||||
"""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_veo_operation(raw_response)
|
||||
|
||||
if not response_data.get("done", False):
|
||||
raise ValueError(
|
||||
|
|
@ -548,7 +576,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
Video remix is not supported by Veo API.
|
||||
|
|
@ -574,7 +602,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
order: str | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
Video list is not supported by Veo API.
|
||||
|
|
@ -615,7 +643,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
"""Video delete is not supported."""
|
||||
raise NotImplementedError("Video delete is not supported by Vertex AI Veo.")
|
||||
|
||||
def transform_video_create_character_request(self, name, video, api_base, litellm_params, headers):
|
||||
def transform_video_create_character_request(self, name, video: object, api_base, litellm_params, headers):
|
||||
raise NotImplementedError("video create character is not supported for Vertex AI")
|
||||
|
||||
def transform_video_create_character_response(self, raw_response, logging_obj):
|
||||
|
|
@ -649,7 +677,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
prefetched_source_data: dict[str, Any] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
|
|
@ -667,12 +695,13 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
if not prefetched_source_data.get("done", False):
|
||||
raise ValueError("Source video generation is not complete yet. Check the video status before editing.")
|
||||
|
||||
videos: Final = prefetched_source_data.get("response", {}).get("videos", [])
|
||||
source_response: Final[_VeoOperationResponse] = prefetched_source_data.get("response", {})
|
||||
videos: Final = source_response.get("videos", [])
|
||||
if not videos:
|
||||
raise ValueError("No videos found in the completed operation. Cannot edit.")
|
||||
|
||||
source_video: Final = videos[0]
|
||||
video_input: Final[dict[str, Any]] = {}
|
||||
video_input: Final[dict[str, str]] = {}
|
||||
if "gcsUri" in source_video:
|
||||
video_input["gcsUri"] = source_video["gcsUri"]
|
||||
elif "bytesBase64Encoded" in source_video:
|
||||
|
|
@ -684,13 +713,13 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
operation_name: Final = extract_original_video_id(video_id)
|
||||
model: Final = self.extract_model_from_operation_name(operation_name) or ""
|
||||
|
||||
instance_dict: Final[dict[str, Any]] = {"prompt": prompt, "video": video_input}
|
||||
request_data: Final[dict[str, Any]] = {"instances": [instance_dict]}
|
||||
instance_dict: Final[dict[str, object]] = {"prompt": prompt, "video": video_input}
|
||||
request_data: Final[dict[str, object]] = {"instances": [instance_dict]}
|
||||
|
||||
if extra_body:
|
||||
extra_body_copy: Final = dict(extra_body)
|
||||
nested_params: Final = extra_body_copy.pop("parameters", None)
|
||||
vertex_params: Final[dict[str, Any]] = {}
|
||||
vertex_params: Final[dict[str, object]] = {}
|
||||
if isinstance(nested_params, dict):
|
||||
vertex_params.update(nested_params)
|
||||
vertex_params.update(extra_body_copy)
|
||||
|
|
@ -716,7 +745,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
|
|||
usage includes duration_seconds and optional video_resolution from the
|
||||
edit request parameters for cost calculation.
|
||||
"""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final = _parse_veo_operation(raw_response)
|
||||
|
||||
operation_name: Final = response_data.get("name")
|
||||
if not operation_name:
|
||||
|
|
|
|||
|
|
@ -8,11 +8,33 @@ import json
|
|||
import os
|
||||
import re
|
||||
from collections.abc import Iterable, Iterator, Mapping, MutableMapping, MutableSequence
|
||||
from typing import Any, Final
|
||||
from collections.abc import Set as AbstractSet
|
||||
from typing import Any, Final, Protocol
|
||||
from urllib.parse import quote
|
||||
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
class _McpServerLike(Protocol):
|
||||
@property
|
||||
def server_id(self) -> str: ...
|
||||
@property
|
||||
def server_name(self) -> str | None: ...
|
||||
@property
|
||||
def alias(self) -> str | None: ...
|
||||
@property
|
||||
def short_prefix(self) -> str | None: ...
|
||||
|
||||
|
||||
class McpServerPayloadLike(Protocol):
|
||||
alias: str | None
|
||||
|
||||
@property
|
||||
def server_name(self) -> str | None: ...
|
||||
@property
|
||||
def tool_name_to_display_name(self) -> Mapping[str, str] | None: ...
|
||||
|
||||
|
||||
# Constants
|
||||
#
|
||||
# NOTE: The environment-backed values below are read once, when this module is
|
||||
|
|
@ -102,7 +124,7 @@ def compute_short_server_prefix(server_id: str, attempt: int = 0) -> str:
|
|||
# at the end so the first emitted char comes from the high-order
|
||||
# bits of the digest (which is the position we constrain to be
|
||||
# alphabetic).
|
||||
chars: Final = []
|
||||
chars: Final[list[str]] = []
|
||||
for position in range(SHORT_MCP_TOOL_PREFIX_LENGTH):
|
||||
is_first_char = position == SHORT_MCP_TOOL_PREFIX_LENGTH - 1
|
||||
alphabet = _BASE52_ALPHA_ALPHABET if is_first_char else _BASE62_ALPHABET
|
||||
|
|
@ -176,34 +198,34 @@ def lookup_mcp_server_auth_in_headers(
|
|||
MCP_TOOL_ALLOWLIST_ENFORCED_KEY: Final = "tool_allowlist_enforced"
|
||||
|
||||
|
||||
def _parse_mcp_info_dict(mcp_info: Any) -> dict[str, Any] | None:
|
||||
def _parse_mcp_info_dict(mcp_info: object) -> Mapping[str, object] | None:
|
||||
if mcp_info is None:
|
||||
return None
|
||||
if isinstance(mcp_info, dict):
|
||||
return mcp_info
|
||||
if isinstance(mcp_info, str):
|
||||
try:
|
||||
parsed: Final = json.loads(mcp_info)
|
||||
parsed: Final[object] = json.loads(mcp_info)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
return parsed if isinstance(parsed, dict) else None
|
||||
return None
|
||||
|
||||
|
||||
def is_server_tool_allowlist_enforced(mcp_server: Any) -> bool:
|
||||
def is_server_tool_allowlist_enforced(mcp_server: object) -> bool:
|
||||
mcp_info: Final = _parse_mcp_info_dict(getattr(mcp_server, "mcp_info", None))
|
||||
if not mcp_info:
|
||||
return False
|
||||
return bool(mcp_info.get(MCP_TOOL_ALLOWLIST_ENFORCED_KEY))
|
||||
|
||||
|
||||
def server_applies_tool_allowlist(mcp_server: Any) -> bool:
|
||||
def server_applies_tool_allowlist(mcp_server: object) -> bool:
|
||||
"""Whether server-level allowed_tools whitelist filtering is active."""
|
||||
allowed_tools: Final = getattr(mcp_server, "allowed_tools", None) or []
|
||||
allowed_tools: Final[object] = getattr(mcp_server, "allowed_tools", None) or []
|
||||
return is_server_tool_allowlist_enforced(mcp_server) or bool(allowed_tools)
|
||||
|
||||
|
||||
def validate_and_normalize_mcp_server_payload(payload: Any) -> None:
|
||||
def validate_and_normalize_mcp_server_payload(payload: McpServerPayloadLike) -> None:
|
||||
"""
|
||||
Validate and normalize MCP server payload fields (server_name, alias, and
|
||||
tool_name_to_display_name).
|
||||
|
|
@ -233,8 +255,8 @@ def validate_and_normalize_mcp_server_payload(payload: Any) -> None:
|
|||
validate_tool_display_names(payload.tool_name_to_display_name)
|
||||
|
||||
# Alias normalization and defaulting
|
||||
alias = getattr(payload, "alias", None)
|
||||
server_name: Final = getattr(payload, "server_name", None)
|
||||
alias: str | None = getattr(payload, "alias", None)
|
||||
server_name: Final[str | None] = getattr(payload, "server_name", None)
|
||||
|
||||
if not alias and server_name:
|
||||
alias = normalize_server_name(server_name)
|
||||
|
|
@ -257,7 +279,7 @@ def add_server_prefix_to_name(name: str, server_name: str) -> str:
|
|||
)
|
||||
|
||||
|
||||
def get_server_prefix(server: Any) -> str:
|
||||
def get_server_prefix(server: object) -> str:
|
||||
"""Return the prefix for a server.
|
||||
|
||||
When the short-prefix mode is enabled (``LITELLM_USE_SHORT_MCP_TOOL_PREFIX``)
|
||||
|
|
@ -270,23 +292,26 @@ def get_server_prefix(server: Any) -> str:
|
|||
alias if present, else server_name, else server_id.
|
||||
"""
|
||||
if is_short_mcp_tool_prefix_enabled():
|
||||
cached: Final = getattr(server, "short_prefix", None)
|
||||
cached: Final[str | None] = getattr(server, "short_prefix", None)
|
||||
if cached:
|
||||
return cached
|
||||
server_id: Final = getattr(server, "server_id", None)
|
||||
server_id: Final[str | None] = getattr(server, "server_id", None)
|
||||
if server_id:
|
||||
return compute_short_server_prefix(server_id)
|
||||
|
||||
if hasattr(server, "alias") and server.alias:
|
||||
return server.alias
|
||||
if hasattr(server, "server_name") and server.server_name:
|
||||
return server.server_name
|
||||
alias: Final[str | None] = getattr(server, "alias", None)
|
||||
if alias:
|
||||
return alias
|
||||
server_name: Final[str | None] = getattr(server, "server_name", None)
|
||||
if server_name:
|
||||
return server_name
|
||||
if hasattr(server, "server_id"):
|
||||
return server.server_id
|
||||
fallback_server_id: Final[str] = getattr(server, "server_id", "")
|
||||
return fallback_server_id
|
||||
return ""
|
||||
|
||||
|
||||
def iter_known_server_prefixes(server: Any) -> Iterator[str]:
|
||||
def iter_known_server_prefixes(server: _McpServerLike) -> Iterator[str]:
|
||||
"""Yield every prefix form that may appear in tool names for ``server``.
|
||||
|
||||
Always includes the *current* prefix returned by ``get_server_prefix``.
|
||||
|
|
@ -304,7 +329,7 @@ def iter_known_server_prefixes(server: Any) -> Iterator[str]:
|
|||
yield from _emit(get_server_prefix(server))
|
||||
yield from _emit(getattr(server, "short_prefix", None))
|
||||
|
||||
server_id: Final = getattr(server, "server_id", None)
|
||||
server_id: Final[str | None] = getattr(server, "server_id", None)
|
||||
if server_id:
|
||||
try:
|
||||
yield from _emit(compute_short_server_prefix(server_id))
|
||||
|
|
@ -397,7 +422,7 @@ def match_known_server_prefix(name: str, known_prefixes: Iterable[str]) -> tuple
|
|||
return None
|
||||
|
||||
|
||||
def strip_known_server_prefix(name: str, server: Any | None) -> str:
|
||||
def strip_known_server_prefix(name: str, server: _McpServerLike | None) -> str:
|
||||
"""Strip ``server``'s registered prefix from a prefixed tool/resource name.
|
||||
|
||||
Unlike :func:`split_server_prefix_from_name`, which guesses the boundary at
|
||||
|
|
@ -420,7 +445,7 @@ def strip_known_server_prefix(name: str, server: Any | None) -> str:
|
|||
|
||||
def is_tool_name_prefixed(
|
||||
tool_name: str,
|
||||
known_server_prefixes: set | None = None,
|
||||
known_server_prefixes: AbstractSet[str] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if tool name has a known MCP server prefix.
|
||||
|
|
@ -640,7 +665,7 @@ def parse_admin_env_vars(
|
|||
if raw is None:
|
||||
continue
|
||||
if hasattr(raw, "model_dump"):
|
||||
entry = raw.model_dump()
|
||||
entry: Mapping[str, object] = raw.model_dump()
|
||||
elif isinstance(raw, dict):
|
||||
entry = raw
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from itertools import groupby
|
|||
from typing import TYPE_CHECKING, Final, NamedTuple
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -180,12 +181,17 @@ def build_autorouter_turn_transaction(
|
|||
|
||||
The routing_decision record is what says a request was auto-routed at all, so a
|
||||
request without one (including the auto-router's own classifier sub-calls) never
|
||||
reaches the rollup. Failed requests served nothing and are excluded. Cache facts
|
||||
are derived from the payload's own usage record through the savings owner, never
|
||||
handed in beside it.
|
||||
reaches the rollup. Internal sub-calls that DO carry one (a shadow eval's duplicate
|
||||
of a request through the router) are excluded by their internal_call_origin stamp:
|
||||
they are not traffic a user sent, so counting them would manufacture sessions and
|
||||
savings in the adoption metrics. Failed requests served nothing and are excluded.
|
||||
Cache facts are derived from the payload's own usage record through the savings
|
||||
owner, never handed in beside it.
|
||||
"""
|
||||
if payload.get("status") != "success":
|
||||
return None
|
||||
if metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY):
|
||||
return None
|
||||
routing_decision: Final = metadata.get("routing_decision")
|
||||
if not isinstance(routing_decision, Mapping) or not routing_decision:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.caching import RedisCache
|
|||
from litellm.constants import (
|
||||
DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME,
|
||||
DB_SPEND_UPDATE_JOB_NAME,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -1794,6 +1795,7 @@ class DBSpendUpdateWriter:
|
|||
if call_type:
|
||||
endpoint = ROUTE_ENDPOINT_MAPPING.get(call_type, None)
|
||||
|
||||
is_internal_call: Final = bool(_metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY))
|
||||
cache_read_input_tokens: Final = extract_cache_read_tokens(usage_obj)
|
||||
compression_saved_tokens: Final = extract_compression_saved_tokens(_metadata)
|
||||
savings_spend: Final = compute_savings_spend(
|
||||
|
|
@ -1818,15 +1820,20 @@ class DBSpendUpdateWriter:
|
|||
prompt_tokens=payload["prompt_tokens"],
|
||||
completion_tokens=payload["completion_tokens"],
|
||||
spend=payload["spend"],
|
||||
api_requests=1,
|
||||
successful_requests=1 if request_status == "success" else 0,
|
||||
failed_requests=1 if request_status != "success" else 0,
|
||||
# Internal sub-calls (auto-router classifier, shadow eval's shadow and
|
||||
# judge) bill real spend and tokens to the key, but they are not
|
||||
# requests the caller made: counting them inflates request-volume
|
||||
# readers, and an auto-router savings figure computed on a shadow
|
||||
# duplicate credits savings for traffic no user sent.
|
||||
api_requests=0 if is_internal_call else 1,
|
||||
successful_requests=1 if not is_internal_call and request_status == "success" else 0,
|
||||
failed_requests=1 if not is_internal_call and request_status != "success" else 0,
|
||||
cache_read_input_tokens=cache_read_input_tokens,
|
||||
cache_creation_input_tokens=extract_cache_creation_tokens(usage_obj),
|
||||
compression_saved_tokens=compression_saved_tokens,
|
||||
compression_savings_spend=savings_spend.compression,
|
||||
prompt_caching_savings_spend=savings_spend.prompt_caching,
|
||||
autorouter_savings_spend=savings_spend.autorouter,
|
||||
autorouter_savings_spend=0.0 if is_internal_call else savings_spend.autorouter,
|
||||
)
|
||||
return daily_transaction
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import requests
|
|||
from fastapi import HTTPException
|
||||
from httpx import HTTPStatusError
|
||||
from requests.auth import HTTPBasicAuth
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
|
|
@ -55,6 +56,26 @@ class _HiddenlayerResponse(TypedDict, total=False):
|
|||
modified_data: Mapping[str, _HiddenlayerModifiedSide]
|
||||
|
||||
|
||||
class _LoggedCallMetadata(TypedDict, total=False):
|
||||
headers: ReadOnly[Mapping[str, str]]
|
||||
|
||||
|
||||
class _LoggedCallLitellmParams(TypedDict, total=False):
|
||||
metadata: ReadOnly[_LoggedCallMetadata]
|
||||
|
||||
|
||||
class _HiddenlayerOutputMessage(TypedDict, total=False):
|
||||
content: ReadOnly[str | Sequence[Mapping[str, str]]]
|
||||
|
||||
|
||||
class _HiddenlayerChoiceMessage(TypedDict, total=False):
|
||||
content: ReadOnly[str]
|
||||
|
||||
|
||||
class _HiddenlayerChoice(TypedDict, total=False):
|
||||
message: ReadOnly[_HiddenlayerChoiceMessage]
|
||||
|
||||
|
||||
def is_saas(host: str) -> bool:
|
||||
"""Checks whether the connection is to the SaaS platform"""
|
||||
|
||||
|
|
@ -155,7 +176,10 @@ class HiddenlayerGuardrail(CustomGuardrail):
|
|||
# from the logger object on the response from the model.
|
||||
headers = request_data.get("proxy_server_request", {}).get("headers", {})
|
||||
if not headers and logging_obj and logging_obj.model_call_details:
|
||||
headers = logging_obj.model_call_details.get("litellm_params", {}).get("metadata", {}).get("headers", {})
|
||||
logged_litellm_params: Final[_LoggedCallLitellmParams] = logging_obj.model_call_details.get(
|
||||
"litellm_params", {}
|
||||
)
|
||||
headers = logged_litellm_params.get("metadata", {}).get("headers", {})
|
||||
|
||||
hl_request_metadata["requester_id"] = headers.get("hl-requester-id") or "LiteLLM"
|
||||
project_id: Final = headers.get("hl-project-id")
|
||||
|
|
@ -408,7 +432,8 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
|
|||
if input_type == "request":
|
||||
inputs["structured_messages"] = output
|
||||
|
||||
for message in output.get("messages", []):
|
||||
modified_messages: Final[Sequence[_HiddenlayerOutputMessage]] = output.get("messages", [])
|
||||
for message in modified_messages:
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, list):
|
||||
text_parts = [
|
||||
|
|
@ -422,7 +447,8 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
|
|||
inputs["texts"] = new_texts
|
||||
|
||||
elif input_type == "response" and inputs.get("texts"):
|
||||
inputs["texts"] = [output.get("choices", [{}])[-1].get("message", {}).get("content", "")]
|
||||
redacted_choices: Final[Sequence[_HiddenlayerChoice]] = output.get("choices", [{}])
|
||||
inputs["texts"] = [redacted_choices[-1].get("message", {}).get("content", "")]
|
||||
elif input_type == "response" and inputs.get("tool_calls"):
|
||||
inputs["tool_calls"] = output
|
||||
|
||||
|
|
|
|||
|
|
@ -1,16 +1,20 @@
|
|||
"""LLM-as-a-Judge guardrail: uses an LLM to score responses against weighted criteria."""
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.llm_judge import (
|
||||
default_router_provider,
|
||||
extract_text_from_content,
|
||||
judge_acompletion,
|
||||
parse_json_verdict,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus
|
||||
|
||||
|
|
@ -32,50 +36,9 @@ Return ONLY valid JSON in this exact format:
|
|||
|
||||
_VALID_ON_FAILURE: Final = frozenset({"block", "log"})
|
||||
|
||||
|
||||
def _default_router_provider() -> "Router | None":
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
return llm_router
|
||||
|
||||
|
||||
_JSON_FENCE_RE: Final = re.compile(r"```(?:json)?\s*(.*?)\s*```", re.DOTALL | re.IGNORECASE)
|
||||
|
||||
|
||||
def _parse_judge_verdict(raw: str) -> dict[str, Any]:
|
||||
"""Parse the judge's JSON verdict, tolerating markdown fences and surrounding prose."""
|
||||
text = raw.strip()
|
||||
fenced: Final = _JSON_FENCE_RE.search(text)
|
||||
if fenced is not None:
|
||||
text = fenced.group(1).strip()
|
||||
parsed: object
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
start: Final = text.find("{")
|
||||
end: Final = text.rfind("}")
|
||||
if start == -1 or end <= start:
|
||||
raise
|
||||
parsed = json.loads(text[start : end + 1])
|
||||
if not isinstance(parsed, dict):
|
||||
raise ValueError("judge response is not a JSON object")
|
||||
return cast(dict[str, Any], parsed) # cast-ok: narrowed to dict by the isinstance guard above
|
||||
|
||||
|
||||
def _extract_text_from_content(content: Any) -> str:
|
||||
"""Return plain text from a message content field (str or multimodal list)."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts: Final = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
parts.append(part.get("text", ""))
|
||||
return " ".join(parts)
|
||||
return ""
|
||||
_default_router_provider: Final = default_router_provider
|
||||
_parse_judge_verdict: Final = parse_json_verdict
|
||||
_extract_text_from_content: Final = extract_text_from_content
|
||||
|
||||
|
||||
def _get_litellm_param(
|
||||
|
|
@ -168,25 +131,13 @@ class LLMAsAJudgeGuardrail(CustomGuardrail):
|
|||
"content": _build_judge_prompt(self.criteria, messages, response_text),
|
||||
},
|
||||
]
|
||||
router: Final = self._router_provider()
|
||||
if router is not None and (
|
||||
self.judge_model in router.model_group_alias or router.get_model_list(model_name=self.judge_model)
|
||||
):
|
||||
response = await router.acompletion(
|
||||
model=self.judge_model,
|
||||
messages=judge_messages,
|
||||
response_format={"type": "json_object"},
|
||||
temperature=0,
|
||||
num_retries=0,
|
||||
fallbacks=[],
|
||||
)
|
||||
else:
|
||||
response = await litellm.acompletion(
|
||||
model=self.judge_model,
|
||||
messages=judge_messages,
|
||||
response_format={"type": "json_object"},
|
||||
temperature=0,
|
||||
)
|
||||
response: Final = await judge_acompletion(
|
||||
self._router_provider(),
|
||||
self.judge_model,
|
||||
judge_messages,
|
||||
response_format={"type": "json_object"},
|
||||
temperature=0,
|
||||
)
|
||||
raw: Final = response.choices[0].message.content or "{}"
|
||||
return _parse_judge_verdict(raw)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,9 +2,10 @@
|
|||
|
||||
import importlib
|
||||
import os
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from datetime import datetime, timezone
|
||||
from itertools import chain, count
|
||||
from typing import Any, Final, Literal, Optional, cast
|
||||
from typing import Any, Final, Literal, Optional, Protocol, cast
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
|
@ -59,6 +60,13 @@ from .guardrail_initializers import (
|
|||
initialize_tool_permission,
|
||||
)
|
||||
|
||||
|
||||
class _GuardrailRowLike(Protocol):
|
||||
@property
|
||||
def guardrail_id(self) -> str: ...
|
||||
def __iter__(self) -> Iterator[tuple[str, object]]: ...
|
||||
|
||||
|
||||
guardrail_initializer_registry: Final = {
|
||||
SupportedGuardrailIntegrations.BEDROCK.value: initialize_bedrock,
|
||||
SupportedGuardrailIntegrations.LAKERA.value: initialize_lakera,
|
||||
|
|
@ -125,7 +133,9 @@ def get_guardrail_initializer_from_hooks():
|
|||
|
||||
# Check for guardrail_initializer_registry dictionary
|
||||
if hasattr(module, "guardrail_initializer_registry"):
|
||||
registry = getattr(module, "guardrail_initializer_registry")
|
||||
registry: Mapping[str, Callable[..., CustomGuardrail]] | None = getattr(
|
||||
module, "guardrail_initializer_registry", None
|
||||
)
|
||||
if isinstance(registry, dict):
|
||||
discovered_initializers.update(registry)
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -135,7 +145,7 @@ def get_guardrail_initializer_from_hooks():
|
|||
# Check for standalone initialize_guardrail function (fallback for directory-based guardrails)
|
||||
elif hasattr(module, "initialize_guardrail"):
|
||||
# For directories with just initialize_guardrail, use the directory name as the key
|
||||
initialize_fn = getattr(module, "initialize_guardrail")
|
||||
initialize_fn: Callable[..., CustomGuardrail] | None = getattr(module, "initialize_guardrail", None)
|
||||
discovered_initializers[item] = initialize_fn
|
||||
verbose_proxy_logger.debug("Found initialize_guardrail function in %s", module_path)
|
||||
|
||||
|
|
@ -206,7 +216,9 @@ def get_guardrail_class_from_hooks():
|
|||
|
||||
# Check for guardrail_initializer_registry dictionary
|
||||
if hasattr(module, "guardrail_class_registry"):
|
||||
registry = getattr(module, "guardrail_class_registry")
|
||||
registry: Mapping[str, type[CustomGuardrail]] | None = getattr(
|
||||
module, "guardrail_class_registry", None
|
||||
)
|
||||
if isinstance(registry, dict):
|
||||
discovered_classes.update(registry)
|
||||
|
||||
|
|
@ -275,7 +287,7 @@ class GuardrailRegistry:
|
|||
guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {}))
|
||||
|
||||
# Create guardrail in DB
|
||||
created_guardrail: Final = await GuardrailsRepository(prisma_client).table.create(
|
||||
created_guardrail: Final[_GuardrailRowLike] = await GuardrailsRepository(prisma_client).table.create(
|
||||
data={
|
||||
"guardrail_name": guardrail_name,
|
||||
"litellm_params": litellm_params,
|
||||
|
|
@ -321,7 +333,7 @@ class GuardrailRegistry:
|
|||
guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {}))
|
||||
|
||||
# Update in DB
|
||||
updated_guardrail: Final = await GuardrailsRepository(prisma_client).table.update(
|
||||
updated_guardrail: Final[_GuardrailRowLike] = await GuardrailsRepository(prisma_client).table.update(
|
||||
where={"guardrail_id": guardrail_id},
|
||||
data={
|
||||
"guardrail_name": guardrail_name,
|
||||
|
|
@ -482,7 +494,7 @@ class InMemoryGuardrailHandler:
|
|||
custom_guardrail_callback = initializer(litellm_params, guardrail)
|
||||
elif isinstance(guardrail_type, str) and "." in guardrail_type:
|
||||
custom_guardrail_callback = self.initialize_custom_guardrail(
|
||||
guardrail=cast(dict, guardrail),
|
||||
guardrail=guardrail,
|
||||
guardrail_type=guardrail_type,
|
||||
litellm_params=litellm_params,
|
||||
config_file_path=config_file_path,
|
||||
|
|
@ -512,7 +524,7 @@ class InMemoryGuardrailHandler:
|
|||
"skip_tool_message_in_guardrail are enabled together, which excludes every message from "
|
||||
"scanning, so no request content would ever be scanned. Remove one of the two."
|
||||
)
|
||||
configured_run_in_parallel: Final = getattr(litellm_params, "run_in_parallel", None)
|
||||
configured_run_in_parallel: Final[bool | None] = getattr(litellm_params, "run_in_parallel", None)
|
||||
if configured_run_in_parallel is not None:
|
||||
custom_guardrail_callback.run_in_parallel = bool(configured_run_in_parallel)
|
||||
|
||||
|
|
@ -532,7 +544,7 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
def initialize_custom_guardrail(
|
||||
self,
|
||||
guardrail: dict,
|
||||
guardrail: Guardrail,
|
||||
guardrail_type: str,
|
||||
litellm_params: LitellmParams,
|
||||
config_file_path: str | None = None,
|
||||
|
|
@ -550,7 +562,9 @@ class InMemoryGuardrailHandler:
|
|||
guardrail_type,
|
||||
)
|
||||
|
||||
_guardrail_class: Final = get_instance_fn(guardrail_type, config_file_path=config_file_path)
|
||||
_guardrail_class: Final[Callable[..., CustomGuardrail]] = get_instance_fn(
|
||||
guardrail_type, config_file_path=config_file_path
|
||||
)
|
||||
|
||||
mode: Final = litellm_params.mode
|
||||
if mode is None:
|
||||
|
|
@ -683,8 +697,8 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
@staticmethod
|
||||
def _normalize_litellm_params_for_comparison(
|
||||
params: Any | None,
|
||||
) -> dict[str, Any] | None:
|
||||
params: LitellmParams | Mapping[str, object] | None,
|
||||
) -> Mapping[str, object] | None:
|
||||
"""
|
||||
Render litellm_params to a canonical dict so an in-memory LitellmParams and
|
||||
the raw dict loaded from the DB compare equal when they describe the same
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ from typing import (
|
|||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE
|
||||
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
|
|
@ -2991,6 +2991,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
rate_limit_type: Literal["output", "input", "total"],
|
||||
) -> list[RedisPipelineIncrementOperation]:
|
||||
"""Build Redis pipeline increment ops for TPM / parallel-request counters."""
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
get_model_group_from_litellm_kwargs,
|
||||
)
|
||||
|
|
@ -2998,6 +2999,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
# Get metadata from standard_logging_object - this correctly handles both
|
||||
# 'metadata' and 'litellm_metadata' fields from litellm_params
|
||||
standard_logging_object: Final = kwargs.get("standard_logging_object") or {}
|
||||
request_metadata: Final = get_litellm_metadata_from_kwargs(kwargs)
|
||||
if request_metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY):
|
||||
# Internal sub-calls bill spend to the caller but are not the caller's
|
||||
# traffic; charging them here would let background evals eat TPM headroom.
|
||||
return []
|
||||
standard_logging_metadata: Final = standard_logging_object.get("metadata") or {}
|
||||
|
||||
model_group: Final = get_model_group_from_litellm_kwargs(kwargs)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from pydantic import BaseModel, TypeAdapter
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
from litellm.litellm_core_utils.llm_judge import router_resolves_model
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -39,11 +40,16 @@ from litellm.types.management_endpoints.auto_router_endpoints import (
|
|||
AutoRouterRoutingTestRequest,
|
||||
AutoRouterRoutingTestResponse,
|
||||
RequestComplexityRouterConfig,
|
||||
ShadowEvalJobResponse,
|
||||
ShadowEvalResult,
|
||||
ShadowEvalSlice,
|
||||
StartShadowEvalRequest,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.router import Router
|
||||
else:
|
||||
try:
|
||||
|
|
@ -388,14 +394,7 @@ async def get_auto_router_benchmarks(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if user_api_key_dict.user_role not in (
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Only proxy admin roles can view auto-router benchmarks across the deployment",
|
||||
)
|
||||
_require_admin_viewer(user_api_key_dict, "view auto-router benchmarks across the deployment")
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
|
|
@ -430,3 +429,335 @@ async def get_auto_router_benchmarks(
|
|||
totals=_benchmark_totals(_summed_agg_row(rows)),
|
||||
groups=groups,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shadow eval: pre-adoption evaluation of an auto-router against live traffic.
|
||||
# The job row is immutable config plus stopped_at; status, counts, spend, and errors
|
||||
# are derived from the append-only attempt rows, so reads here are aggregations
|
||||
# bounded by each job's max_turns through the attempt table's job_id index.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin_viewer(user_api_key_dict: UserAPIKeyAuth, action: str) -> None:
|
||||
if user_api_key_dict.user_role not in (
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
):
|
||||
raise HTTPException(status_code=403, detail=f"Only proxy admin roles can {action}")
|
||||
|
||||
|
||||
def _require_admin_writer(user_api_key_dict: UserAPIKeyAuth, action: str) -> None:
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(status_code=403, detail=f"Only a proxy admin can {action}")
|
||||
|
||||
|
||||
def _is_configured_pre_routing_strategy(llm_router: "Router", router_name: str) -> bool:
|
||||
return any(
|
||||
router_name in registry
|
||||
for registry in (
|
||||
llm_router.auto_routers,
|
||||
llm_router.complexity_routers,
|
||||
llm_router.adaptive_routers,
|
||||
llm_router.quality_routers,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _validate_judge_model(llm_router: "Router | None", judge_model: str) -> None:
|
||||
"""Reject a judge model the dispatch path cannot resolve, at start rather than as a
|
||||
silently growing error count once the job is already sampling and billing."""
|
||||
if llm_router is not None and _is_configured_pre_routing_strategy(llm_router, judge_model):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"judge_model '{judge_model}' is an auto-router; the judge must be a plain model",
|
||||
)
|
||||
if router_resolves_model(llm_router, judge_model):
|
||||
return
|
||||
import litellm
|
||||
|
||||
try:
|
||||
litellm.get_llm_provider(model=judge_model)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"judge_model '{judge_model}' is neither a model configured on this proxy nor a "
|
||||
"provider-qualified public model name (e.g. 'anthropic/claude-sonnet-5')"
|
||||
),
|
||||
) from e
|
||||
|
||||
|
||||
def _is_unique_violation(error: Exception) -> bool:
|
||||
"""Whether a Prisma create failed on a unique index. One active job per key lives in
|
||||
a partial unique index (raw SQL in the migration; schema.prisma cannot express partial
|
||||
indexes), so the read-then-create check above it is advisory: two concurrent starts
|
||||
pass the read, and the loser must surface as the same 409 rather than a 500."""
|
||||
try:
|
||||
from prisma.errors import UniqueViolationError
|
||||
except ImportError:
|
||||
return "unique constraint" in str(error).lower() or "P2002" in str(error)
|
||||
return isinstance(error, UniqueViolationError)
|
||||
|
||||
|
||||
class _AttemptAggRow(BaseModel):
|
||||
grp: str
|
||||
turn_count: int
|
||||
real_wins: int
|
||||
shadow_wins: int
|
||||
ties: int
|
||||
avg_confidence: float | None
|
||||
|
||||
|
||||
_ATTEMPT_AGG_ROWS: Final = TypeAdapter(list[_AttemptAggRow])
|
||||
|
||||
_ATTEMPT_AGG_SELECT: Final = """
|
||||
COUNT(*)::int AS turn_count,
|
||||
COUNT(*) FILTER (WHERE outcome = 'real')::int AS real_wins,
|
||||
COUNT(*) FILTER (WHERE outcome = 'shadow')::int AS shadow_wins,
|
||||
COUNT(*) FILTER (WHERE outcome = 'tie')::int AS ties,
|
||||
AVG(confidence)::float AS avg_confidence
|
||||
FROM "LiteLLM_ShadowEvalAttempt"
|
||||
WHERE job_id = $1 AND outcome != 'error'
|
||||
GROUP BY 1
|
||||
"""
|
||||
|
||||
_ATTEMPT_AGG_BY_TIER_SQL: Final = "SELECT COALESCE(tier, 'UNCLASSIFIED') AS grp," + _ATTEMPT_AGG_SELECT
|
||||
_ATTEMPT_AGG_BY_MODEL_SQL: Final = "SELECT COALESCE(real_model, 'unknown') AS grp," + _ATTEMPT_AGG_SELECT
|
||||
|
||||
_SWEEP_FINISHED_JOBS_SQL: Final = """
|
||||
UPDATE "LiteLLM_ShadowEvalJob" j SET stopped_at = NOW()
|
||||
WHERE j.api_key_id = $1 AND j.stopped_at IS NULL
|
||||
AND (
|
||||
j.ends_at <= NOW()
|
||||
OR (SELECT COUNT(*) FROM "LiteLLM_ShadowEvalAttempt" a WHERE a.job_id = j.id) >= j.max_turns
|
||||
)
|
||||
"""
|
||||
|
||||
_ATTEMPT_TOTALS_SQL: Final = """
|
||||
SELECT
|
||||
COUNT(*) FILTER (WHERE outcome != 'error')::int AS judged_count,
|
||||
COUNT(*) FILTER (WHERE outcome = 'error')::int AS error_count,
|
||||
COALESCE(SUM(judge_cost), 0)::float AS judge_spend
|
||||
FROM "LiteLLM_ShadowEvalAttempt"
|
||||
WHERE job_id = $1
|
||||
"""
|
||||
|
||||
|
||||
class _AttemptTotalsRow(BaseModel):
|
||||
judged_count: int
|
||||
error_count: int
|
||||
judge_spend: float
|
||||
|
||||
|
||||
_ATTEMPT_TOTALS_ROWS: Final = TypeAdapter(list[_AttemptTotalsRow])
|
||||
|
||||
|
||||
def _pct_of(numerator: int, denominator: int) -> float:
|
||||
return _pct(numerator, denominator)
|
||||
|
||||
|
||||
def _slices(rows: Sequence[_AttemptAggRow]) -> tuple[ShadowEvalSlice, ...]:
|
||||
return tuple(
|
||||
ShadowEvalSlice(
|
||||
group=row.grp,
|
||||
turn_count=row.turn_count,
|
||||
real_win_rate_pct=_pct_of(row.real_wins, row.turn_count),
|
||||
shadow_win_rate_pct=_pct_of(row.shadow_wins, row.turn_count),
|
||||
tie_rate_pct=_pct_of(row.ties, row.turn_count),
|
||||
avg_judge_confidence=round(row.avg_confidence or 0.0, 3),
|
||||
)
|
||||
for row in sorted(rows, key=lambda r: r.turn_count, reverse=True)
|
||||
)
|
||||
|
||||
|
||||
async def _shadow_eval_results(prisma_client: "PrismaClient", job_id: str) -> ShadowEvalResult | None:
|
||||
"""Both stratifications of one job's verdicts. Tier answers "where does the router do
|
||||
well"; current-model answers "which of the models this key uses today would the router
|
||||
beat". Reads are bounded by the job's own attempts (<= max_turns) via the job_id index."""
|
||||
by_tier: Final = _ATTEMPT_AGG_ROWS.validate_python(
|
||||
await prisma_client.db.query_raw(_ATTEMPT_AGG_BY_TIER_SQL, job_id) or ()
|
||||
)
|
||||
if not by_tier:
|
||||
return None
|
||||
by_model: Final = _ATTEMPT_AGG_ROWS.validate_python(
|
||||
await prisma_client.db.query_raw(_ATTEMPT_AGG_BY_MODEL_SQL, job_id) or ()
|
||||
)
|
||||
total_turns: Final = sum(r.turn_count for r in by_tier)
|
||||
return ShadowEvalResult(
|
||||
by_tier=_slices(by_tier),
|
||||
by_current_model=_slices(by_model),
|
||||
overall_shadow_win_rate_pct=_pct_of(sum(r.shadow_wins for r in by_tier), total_turns),
|
||||
overall_tie_rate_pct=_pct_of(sum(r.ties for r in by_tier), total_turns),
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/auto_router/shadow_eval/start",
|
||||
tags=("auto router",),
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=ShadowEvalJobResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
async def start_shadow_eval(
|
||||
data: StartShadowEvalRequest,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> ShadowEvalJobResponse:
|
||||
"""
|
||||
Start a pre-adoption shadow eval: duplicate a sampled slice of a key's live traffic
|
||||
through an auto-router, judge real vs. shadow responses blind, and stratify win rates
|
||||
by the router's tier classification and by the incumbent model.
|
||||
|
||||
Shadow responses are never served to users. The job samples until it has judged
|
||||
max_turns turns, reaches the end of its window, or is stopped; sampling changes
|
||||
propagate to pods within about 10 seconds. Shadow and judge calls bill to the
|
||||
shadowed key but are excluded from request counts and auto-router adoption metrics.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import llm_router, prisma_client
|
||||
|
||||
_require_admin_writer(user_api_key_dict, "start a shadow eval")
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
if llm_router is None or not _is_configured_pre_routing_strategy(llm_router, data.router_name):
|
||||
raise HTTPException(status_code=400, detail=f"'{data.router_name}' is not a configured auto-router")
|
||||
_validate_judge_model(llm_router, data.judge_model)
|
||||
key_row: Final = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": data.api_key_id} # mutable-ok: Prisma filter
|
||||
)
|
||||
if key_row is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"api_key_id '{data.api_key_id}' is not a key on this proxy; pass the key's token hash, "
|
||||
"the value the key list and key info endpoints report"
|
||||
),
|
||||
)
|
||||
|
||||
# A job that expired or exhausted its turn budget stopped sampling on its own, but
|
||||
# still holds the one-active-per-key partial unique index until stamped; free it so
|
||||
# a new eval can start.
|
||||
await prisma_client.db.execute_raw(_SWEEP_FINISHED_JOBS_SQL, data.api_key_id)
|
||||
active: Final = await prisma_client.db.litellm_shadowevaljob.find_first(
|
||||
where={"api_key_id": data.api_key_id, "stopped_at": None}, # mutable-ok: Prisma filter
|
||||
)
|
||||
if active is not None:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"Key already has an active shadow eval job ({active.id}). Stop it first.",
|
||||
)
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
try:
|
||||
job: Final = await prisma_client.db.litellm_shadowevaljob.create(
|
||||
data={ # mutable-ok: Prisma payload
|
||||
"api_key_id": data.api_key_id,
|
||||
"router_name": data.router_name,
|
||||
"judge_model": data.judge_model,
|
||||
"shadow_percentage": data.shadow_percentage,
|
||||
"max_turns": data.max_turns,
|
||||
"created_by": user_api_key_dict.user_id,
|
||||
"ends_at": now + timedelta(days=data.duration_days),
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
if not _is_unique_violation(e):
|
||||
raise
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail="Key already has an active shadow eval job (started concurrently). Stop it first.",
|
||||
) from e
|
||||
return ShadowEvalJobResponse.model_validate(job, from_attributes=True)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/auto_router/shadow_eval",
|
||||
tags=("auto router",),
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=list[ShadowEvalJobResponse],
|
||||
)
|
||||
async def list_shadow_eval_jobs(
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
api_key_id: Annotated[str | None, Query(description="Filter to jobs shadowing this key")] = None,
|
||||
limit: Annotated[int, Query(ge=1, le=200, description="Newest jobs to return")] = 50,
|
||||
) -> tuple[ShadowEvalJobResponse, ...]:
|
||||
"""List shadow eval jobs, newest first. Counts and results ride the detail endpoint only."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
_require_admin_viewer(user_api_key_dict, "view shadow evals")
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
records: Final = await prisma_client.db.litellm_shadowevaljob.find_many(
|
||||
where={"api_key_id": api_key_id} if api_key_id else {}, # mutable-ok: Prisma filter
|
||||
order={"created_at": "desc"}, # mutable-ok: Prisma order
|
||||
take=limit,
|
||||
)
|
||||
return tuple(ShadowEvalJobResponse.model_validate(record, from_attributes=True) for record in records or ())
|
||||
|
||||
|
||||
@router.get(
|
||||
"/auto_router/shadow_eval/{job_id}",
|
||||
tags=("auto router",),
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=ShadowEvalJobResponse,
|
||||
)
|
||||
async def get_shadow_eval_job(
|
||||
job_id: str,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> ShadowEvalJobResponse:
|
||||
"""One job with derived counts, judge spend, latest error, and stratified results."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
_require_admin_viewer(user_api_key_dict, "view shadow evals")
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
record: Final = await prisma_client.db.litellm_shadowevaljob.find_unique(
|
||||
where={"id": job_id} # mutable-ok: Prisma filter
|
||||
)
|
||||
if record is None:
|
||||
raise HTTPException(status_code=404, detail=f"No shadow eval job {job_id}")
|
||||
totals: Final = _ATTEMPT_TOTALS_ROWS.validate_python(
|
||||
await prisma_client.db.query_raw(_ATTEMPT_TOTALS_SQL, job_id) or ()
|
||||
)
|
||||
latest_error: Final = await prisma_client.db.litellm_shadowevalattempt.find_first(
|
||||
where={"job_id": job_id, "outcome": "error"}, # mutable-ok: Prisma filter
|
||||
order={"created_at": "desc"}, # mutable-ok: Prisma order
|
||||
)
|
||||
return ShadowEvalJobResponse.model_validate(record, from_attributes=True).model_copy(
|
||||
update={ # mutable-ok: pydantic update payload
|
||||
"judged_count": totals[0].judged_count if totals else 0,
|
||||
"error_count": totals[0].error_count if totals else 0,
|
||||
"judge_spend": round(totals[0].judge_spend, 6) if totals else 0.0,
|
||||
"last_error": latest_error.error if latest_error else None,
|
||||
"results": await _shadow_eval_results(prisma_client, job_id),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/auto_router/shadow_eval/{job_id}/stop",
|
||||
tags=("auto router",),
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=ShadowEvalJobResponse,
|
||||
)
|
||||
async def stop_shadow_eval_job(
|
||||
job_id: str,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> ShadowEvalJobResponse:
|
||||
"""Stop an active shadow eval job. Attempts are kept; sampling halts within ~10s."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
_require_admin_writer(user_api_key_dict, "stop a shadow eval")
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
record: Final = await prisma_client.db.litellm_shadowevaljob.find_unique(
|
||||
where={"id": job_id} # mutable-ok: Prisma filter
|
||||
)
|
||||
if record is None:
|
||||
raise HTTPException(status_code=404, detail=f"No shadow eval job {job_id}")
|
||||
current: Final = ShadowEvalJobResponse.model_validate(record, from_attributes=True)
|
||||
if current.status != "running":
|
||||
raise HTTPException(status_code=400, detail=f"Job {job_id} is already {current.status}")
|
||||
updated: Final = await prisma_client.db.litellm_shadowevaljob.update(
|
||||
where={"id": job_id}, # mutable-ok: Prisma filter
|
||||
data={"stopped_at": datetime.now(timezone.utc)}, # mutable-ok: Prisma payload
|
||||
)
|
||||
return ShadowEvalJobResponse.model_validate(updated, from_attributes=True)
|
||||
|
|
|
|||
|
|
@ -10,12 +10,19 @@ All /customer management endpoints
|
|||
"""
|
||||
|
||||
#### END-USER/CUSTOMER MANAGEMENT ####
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Final
|
||||
from typing import TYPE_CHECKING, Final, Protocol, TypeVar, overload
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import LiteLLM_BudgetTable as PrismaBudgetRow
|
||||
from prisma.models import LiteLLM_EndUserTable as PrismaEndUserRow
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -41,6 +48,54 @@ from litellm.types.proxy.management_endpoints.customer_endpoints import (
|
|||
UnblockUsersResponse,
|
||||
)
|
||||
|
||||
_RowT_co: Final = TypeVar("_RowT_co", covariant=True)
|
||||
_STR_OBJECT_DICT: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _TableOps(Protocol[_RowT_co]):
|
||||
async def find_first(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> _RowT_co | None: ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> Sequence[_RowT_co]: ...
|
||||
|
||||
async def create(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> _RowT_co: ...
|
||||
|
||||
async def update(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
data: Mapping[str, object],
|
||||
include: Mapping[str, bool] | None = None,
|
||||
) -> _RowT_co | None: ...
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
data: Mapping[str, Mapping[str, object]],
|
||||
) -> _RowT_co: ...
|
||||
|
||||
async def delete_many(self, where: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
@overload
|
||||
def _typed_table(repo: EndUserRepository) -> "_TableOps[PrismaEndUserRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: BudgetRepository) -> "_TableOps[PrismaBudgetRow]": ...
|
||||
def _typed_table(repo: EndUserRepository | BudgetRepository) -> object:
|
||||
return repo.table
|
||||
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -89,7 +144,7 @@ async def block_user(data: BlockUsers):
|
|||
records: Final = []
|
||||
if prisma_client is not None:
|
||||
for id in data.user_ids:
|
||||
record = await EndUserRepository(prisma_client).table.upsert(
|
||||
record = await _typed_table(EndUserRepository(prisma_client)).upsert(
|
||||
where={"user_id": id},
|
||||
data={
|
||||
"create": {"user_id": id, "blocked": True},
|
||||
|
|
@ -184,7 +239,7 @@ def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None:
|
|||
budget_kv_pairs[field_name] = value
|
||||
|
||||
if budget_kv_pairs:
|
||||
budget_request: Final = BudgetNewRequest(**budget_kv_pairs)
|
||||
budget_request: Final = BudgetNewRequest.model_validate(budget_kv_pairs)
|
||||
validate_budget_duration(budget_request.budget_duration)
|
||||
if budget_request.budget_reset_at is None and budget_request.budget_duration is not None:
|
||||
budget_request.budget_reset_at = datetime.utcnow() + timedelta(
|
||||
|
|
@ -195,10 +250,10 @@ def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None:
|
|||
|
||||
|
||||
async def _handle_customer_object_permission_update(
|
||||
non_default_values: dict,
|
||||
non_default_values: dict[str, object],
|
||||
end_user_table_data_typed: LiteLLM_EndUserTable | None,
|
||||
update_end_user_table_data: dict,
|
||||
prisma_client,
|
||||
update_end_user_table_data: dict[str, object],
|
||||
prisma_client: "PrismaClient",
|
||||
) -> None:
|
||||
"""
|
||||
Handle object permission updates for customer endpoints.
|
||||
|
|
@ -344,13 +399,13 @@ async def new_end_user(
|
|||
},
|
||||
)
|
||||
|
||||
new_end_user_obj: dict = {}
|
||||
new_end_user_obj: dict[str, object] = {}
|
||||
|
||||
## CREATE BUDGET ## if set
|
||||
_new_budget: Final = new_budget_request(data)
|
||||
if _new_budget is not None:
|
||||
try:
|
||||
budget_record: Final = await BudgetRepository(prisma_client).table.create(
|
||||
budget_record: Final = await _typed_table(BudgetRepository(prisma_client)).create(
|
||||
data={
|
||||
**_new_budget.model_dump(exclude_unset=True),
|
||||
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
|
|
@ -364,16 +419,18 @@ async def new_end_user(
|
|||
elif data.budget_id is not None:
|
||||
new_end_user_obj["budget_id"] = data.budget_id
|
||||
|
||||
_user_data: Final = data.dict(exclude_none=True)
|
||||
_user_data: Final = _STR_OBJECT_DICT.validate_python(data.dict(exclude_none=True))
|
||||
|
||||
for k, v in _user_data.items():
|
||||
if k not in BudgetNewRequest.model_fields:
|
||||
new_end_user_obj[k] = v
|
||||
|
||||
## Handle Object Permission - MCP Servers, Vector Stores etc.
|
||||
new_end_user_obj = await _set_object_permission(
|
||||
data_json=new_end_user_obj,
|
||||
prisma_client=prisma_client,
|
||||
new_end_user_obj = _STR_OBJECT_DICT.validate_python(
|
||||
await _set_object_permission(
|
||||
data_json=new_end_user_obj,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
|
||||
# Ensure object_permission is not in the data being sent to create
|
||||
|
|
@ -386,7 +443,7 @@ async def new_end_user(
|
|||
new_end_user_obj.pop("object_permission", None)
|
||||
|
||||
## WRITE TO DB ##
|
||||
end_user_record: Final = await EndUserRepository(prisma_client).table.create(
|
||||
end_user_record: Final = await _typed_table(EndUserRepository(prisma_client)).create(
|
||||
data=new_end_user_obj,
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
|
@ -442,7 +499,7 @@ async def end_user_info(
|
|||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
user_info: Final = await EndUserRepository(prisma_client).table.find_first(
|
||||
user_info: Final = await _typed_table(EndUserRepository(prisma_client)).find_first(
|
||||
where={"user_id": end_user_id},
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
)
|
||||
|
|
@ -535,13 +592,13 @@ async def update_end_user(
|
|||
from litellm.proxy.proxy_server import litellm_proxy_admin_name, prisma_client
|
||||
|
||||
try:
|
||||
data_json: Final[dict] = data.json()
|
||||
data_json: Final = _STR_OBJECT_DICT.validate_python(data.json())
|
||||
# get the row from db
|
||||
if prisma_client is None:
|
||||
raise Exception("Not connected to DB!")
|
||||
|
||||
# get non default values for key
|
||||
non_default_values: Final = {}
|
||||
non_default_values: Final = dict[str, object]()
|
||||
for k, v in data_json.items():
|
||||
if v is not None and v not in (
|
||||
[],
|
||||
|
|
@ -551,7 +608,7 @@ async def update_end_user(
|
|||
non_default_values[k] = v
|
||||
|
||||
## Get end user table data ##
|
||||
end_user_table_data: Final = await EndUserRepository(prisma_client).table.find_first(
|
||||
end_user_table_data: Final = await _typed_table(EndUserRepository(prisma_client)).find_first(
|
||||
where={"user_id": data.user_id}, include={"litellm_budget_table": True}
|
||||
)
|
||||
|
||||
|
|
@ -563,14 +620,14 @@ async def update_end_user(
|
|||
param="user_id",
|
||||
)
|
||||
|
||||
end_user_table_data_typed: Final = LiteLLM_EndUserTable(**end_user_table_data.model_dump())
|
||||
end_user_table_data_typed: Final = LiteLLM_EndUserTable.model_validate(end_user_table_data.model_dump())
|
||||
|
||||
## Get budget table data ##
|
||||
end_user_budget_table: Final = end_user_table_data_typed.litellm_budget_table
|
||||
|
||||
## Get all params for budget table ##
|
||||
budget_table_data: Final = {}
|
||||
update_end_user_table_data: Final = {}
|
||||
budget_table_data: Final = dict[str, object]()
|
||||
update_end_user_table_data: Final = dict[str, object]()
|
||||
for k, v in non_default_values.items():
|
||||
# budget_id is for linking to existing budget, not for creating new budget
|
||||
if k == "budget_id":
|
||||
|
|
@ -593,7 +650,7 @@ async def update_end_user(
|
|||
if budget_table_data:
|
||||
if end_user_budget_table is None:
|
||||
## Create new budget ##
|
||||
budget_table_data_record = await BudgetRepository(prisma_client).table.create(
|
||||
budget_table_data_record = await _typed_table(BudgetRepository(prisma_client)).create(
|
||||
data={
|
||||
**budget_table_data,
|
||||
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
|
||||
|
|
@ -605,7 +662,7 @@ async def update_end_user(
|
|||
update_end_user_table_data["budget_id"] = budget_table_data_record.budget_id
|
||||
else:
|
||||
## Update existing budget ##
|
||||
budget_table_data_record = await BudgetRepository(prisma_client).table.update(
|
||||
budget_table_data_record = await _typed_table(BudgetRepository(prisma_client)).update(
|
||||
where={"budget_id": end_user_budget_table.budget_id},
|
||||
data=budget_table_data,
|
||||
)
|
||||
|
|
@ -625,7 +682,7 @@ async def update_end_user(
|
|||
if data.user_id is not None and len(data.user_id) > 0:
|
||||
update_end_user_table_data["user_id"] = data.user_id
|
||||
verbose_proxy_logger.debug("In update customer, user_id condition block.")
|
||||
response: Final = await EndUserRepository(prisma_client).table.update(
|
||||
response: Final = await _typed_table(EndUserRepository(prisma_client)).update(
|
||||
where={"user_id": data.user_id},
|
||||
data=update_end_user_table_data,
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
|
|
@ -688,7 +745,7 @@ async def delete_end_user(
|
|||
verbose_proxy_logger.debug("/customer/delete: Received data = %s", data)
|
||||
if data.user_ids is not None and isinstance(data.user_ids, list) and len(data.user_ids) > 0:
|
||||
# First check if all users exist
|
||||
existing_users: Final = await EndUserRepository(prisma_client).table.find_many(
|
||||
existing_users: Final = await _typed_table(EndUserRepository(prisma_client)).find_many(
|
||||
where={"user_id": {"in": data.user_ids}}
|
||||
)
|
||||
existing_user_ids: Final = {user.user_id for user in existing_users}
|
||||
|
|
@ -703,7 +760,7 @@ async def delete_end_user(
|
|||
)
|
||||
|
||||
# All users exist, proceed with deletion
|
||||
response: Final = await EndUserRepository(prisma_client).table.delete_many(
|
||||
response: Final = await _typed_table(EndUserRepository(prisma_client)).delete_many(
|
||||
where={"user_id": {"in": data.user_ids}}
|
||||
)
|
||||
verbose_proxy_logger.debug("received response from updating prisma client. response=%s", response)
|
||||
|
|
@ -764,7 +821,7 @@ async def list_end_user(
|
|||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
response: Final = await EndUserRepository(prisma_client).table.find_many(
|
||||
response: Final = await _typed_table(EndUserRepository(prisma_client)).find_many(
|
||||
include={"litellm_budget_table": True, "object_permission": True}
|
||||
)
|
||||
|
||||
|
|
@ -827,11 +884,10 @@ async def get_customer_daily_activity(
|
|||
exclude_end_user_ids_list = exclude_end_user_ids.split(",") if exclude_end_user_ids else None
|
||||
|
||||
# Fetch organization aliases for metadata
|
||||
where_condition: Final = {}
|
||||
where_condition: Final = dict[str, object]()
|
||||
if end_user_ids_list:
|
||||
where_condition["user_id"] = {"in": list(end_user_ids_list)}
|
||||
end_user_aliases: Final = await EndUserRepository(prisma_client).table.find_many(where=where_condition)
|
||||
end_user_alias_metadata: Final = {e.user_id: {"alias": e.alias} for e in end_user_aliases}
|
||||
end_user_aliases: Final = await _typed_table(EndUserRepository(prisma_client)).find_many(where=where_condition)
|
||||
|
||||
# Query daily activity for organizations
|
||||
return await get_daily_activity(
|
||||
|
|
@ -839,7 +895,7 @@ async def get_customer_daily_activity(
|
|||
table_name="litellm_dailyenduserspend",
|
||||
entity_id_field="end_user_id",
|
||||
entity_id=end_user_ids_list,
|
||||
entity_metadata_field=end_user_alias_metadata,
|
||||
entity_metadata_field={e.user_id: {"alias": e.alias} for e in end_user_aliases},
|
||||
exclude_entity_ids=exclude_end_user_ids_list,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.constants import LITELLM_PROXY_ADMIN_NAME
|
|||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
LITELLM_MCP_SERVER_DESCRIPTION,
|
||||
LITELLM_MCP_SERVER_NAME,
|
||||
McpServerPayloadLike,
|
||||
build_env_var_setup_url,
|
||||
collect_env_var_references,
|
||||
get_server_prefix,
|
||||
|
|
@ -196,7 +197,7 @@ if MCP_AVAILABLE:
|
|||
server: MCPServer
|
||||
expires_at: datetime
|
||||
|
||||
def _validate_mcp_server_name_fields(payload: Any) -> None:
|
||||
def _validate_mcp_server_name_fields(payload: McpServerPayloadLike) -> None:
|
||||
candidates: Final[list[tuple[str, str | None]]] = []
|
||||
|
||||
server_name: Final = getattr(payload, "server_name", None)
|
||||
|
|
@ -223,7 +224,7 @@ if MCP_AVAILABLE:
|
|||
detail={"error": error_messages_text},
|
||||
)
|
||||
|
||||
def validate_and_normalize_mcp_server_payload(payload: Any) -> None:
|
||||
def validate_and_normalize_mcp_server_payload(payload: McpServerPayloadLike) -> None:
|
||||
_base_validate_and_normalize_mcp_server_payload(payload)
|
||||
_validate_mcp_server_name_fields(payload)
|
||||
|
||||
|
|
|
|||
|
|
@ -10,13 +10,21 @@ POST /v1/tool/policy - Update the input_policy / output_policy for a
|
|||
"""
|
||||
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final
|
||||
from typing import TYPE_CHECKING, Annotated, Final, Protocol, TypeAlias, TypeVar, overload
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from pydantic import BaseModel, Field, TypeAdapter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import LiteLLM_DailyToolSpend as PrismaDailyToolSpendRow
|
||||
from prisma.models import LiteLLM_ObjectPermissionTable as PrismaObjectPermissionRow
|
||||
from prisma.models import LiteLLM_SpendLogs as PrismaSpendLogRow
|
||||
from prisma.models import LiteLLM_SpendLogToolIndex as PrismaSpendLogToolIndexRow
|
||||
from prisma.models import LiteLLM_TeamTable as PrismaTeamRow
|
||||
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationTokenRow
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -49,6 +57,72 @@ from litellm.types.tool_management import (
|
|||
ToolUsageLogsResponse,
|
||||
)
|
||||
|
||||
_RowT_co: Final = TypeVar("_RowT_co", covariant=True)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _TableOps(Protocol[_RowT_co]):
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = None,
|
||||
order: Mapping[str, object] | Sequence[Mapping[str, object]] | None = None,
|
||||
skip: int | None = None,
|
||||
take: int | None = None,
|
||||
) -> Sequence[_RowT_co]: ...
|
||||
|
||||
async def find_unique(self, where: Mapping[str, object]) -> _RowT_co | None: ...
|
||||
|
||||
async def count(self, where: Mapping[str, object] | None = None) -> int: ...
|
||||
|
||||
async def create(self, data: Mapping[str, object]) -> _RowT_co: ...
|
||||
|
||||
async def update_many(
|
||||
self,
|
||||
where: Mapping[str, object],
|
||||
data: Mapping[str, object],
|
||||
) -> int: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> _RowT_co | None: ...
|
||||
|
||||
async def group_by(
|
||||
self,
|
||||
by: Sequence[str],
|
||||
sum: Mapping[str, bool] | None = None,
|
||||
where: Mapping[str, object] | None = None,
|
||||
order: Mapping[str, object] | None = None,
|
||||
take: int | None = None,
|
||||
) -> Sequence[Mapping[str, object]]: ...
|
||||
|
||||
class _SpendLogRow(Protocol):
|
||||
@property
|
||||
def messages(self) -> object: ...
|
||||
@property
|
||||
def proxy_server_request(self) -> str | Mapping[str, object] | None: ...
|
||||
|
||||
|
||||
@overload
|
||||
def _typed_table(repo: DailyToolSpendRepository) -> "_TableOps[PrismaDailyToolSpendRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: SpendLogToolIndexRepository) -> "_TableOps[PrismaSpendLogToolIndexRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: SpendLogsRepository) -> "_TableOps[PrismaSpendLogRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: VerificationTokenRepository) -> "_TableOps[PrismaVerificationTokenRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: TeamRepository) -> "_TableOps[PrismaTeamRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: ObjectPermissionRepository) -> "_TableOps[PrismaObjectPermissionRow]": ...
|
||||
def _typed_table(
|
||||
repo: DailyToolSpendRepository
|
||||
| SpendLogToolIndexRepository
|
||||
| SpendLogsRepository
|
||||
| VerificationTokenRepository
|
||||
| TeamRepository
|
||||
| ObjectPermissionRepository,
|
||||
) -> object:
|
||||
return repo.table
|
||||
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
TOOL_POLICY_OPTIONS: Final = ToolPolicyOptionsResponse(
|
||||
|
|
@ -201,7 +275,7 @@ async def get_tool_spend(
|
|||
end_str: Final = end_day.strftime("%Y-%m-%d")
|
||||
date_window: Final = {"date": {"gte": start_str, "lte": end_str}}
|
||||
|
||||
table: Final = DailyToolSpendRepository(prisma_client).table
|
||||
table: Final = _typed_table(DailyToolSpendRepository(prisma_client))
|
||||
top_tools: Final = _TOP_TOOL_ROWS.validate_python(
|
||||
await table.group_by(
|
||||
by=["tool_name"],
|
||||
|
|
@ -222,7 +296,7 @@ async def get_tool_spend(
|
|||
for row in top_tools
|
||||
]
|
||||
|
||||
daily_rows: Final = (
|
||||
daily_rows: Final[Sequence[PrismaDailyToolSpendRow]] = (
|
||||
await table.find_many(
|
||||
where={**date_window, "tool_name": {"in": [row.tool_name for row in top_tools]}},
|
||||
order=[{"date": "asc"}, {"spend": "desc"}],
|
||||
|
|
@ -270,36 +344,43 @@ async def get_tool_detail(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> str | None:
|
||||
_ParsedJson: TypeAlias = dict[str, object] | list[object] | str | int | float | bool | None
|
||||
_PARSED_JSON: Final[TypeAdapter[_ParsedJson]] = TypeAdapter(_ParsedJson)
|
||||
_STR_OBJECT_DICT: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _input_snippet_for_tool_log(sl: "_SpendLogRow | None", max_len: int = 200) -> str | None:
|
||||
"""Short snippet from messages or proxy_server_request for tool usage log row."""
|
||||
if sl is None:
|
||||
return None
|
||||
messages: Final = getattr(sl, "messages", None)
|
||||
messages: Final = sl.messages
|
||||
if messages is not None:
|
||||
s = _snippet_str(messages, max_len)
|
||||
if s:
|
||||
return s
|
||||
psr = getattr(sl, "proxy_server_request", None)
|
||||
psr = sl.proxy_server_request
|
||||
if not psr:
|
||||
return None
|
||||
if isinstance(psr, str):
|
||||
import json
|
||||
|
||||
try:
|
||||
psr = json.loads(psr)
|
||||
psr = _PARSED_JSON.validate_python(json.loads(psr))
|
||||
except Exception:
|
||||
return _snippet_str(psr, max_len)
|
||||
if isinstance(psr, dict):
|
||||
msgs = psr.get("messages")
|
||||
if msgs is None and isinstance(psr.get("body"), dict):
|
||||
msgs = psr["body"].get("messages")
|
||||
if msgs is None:
|
||||
body: Final = psr.get("body")
|
||||
if isinstance(body, dict):
|
||||
msgs = _STR_OBJECT_DICT.validate_python(body).get("messages")
|
||||
s = _snippet_str(msgs, max_len)
|
||||
if s:
|
||||
return s
|
||||
return _snippet_str(psr, max_len)
|
||||
|
||||
|
||||
def _snippet_str(text: Any, max_len: int = 200) -> str | None:
|
||||
def _snippet_str(text: object, max_len: int = 200) -> str | None:
|
||||
if text is None:
|
||||
return None
|
||||
if isinstance(text, str):
|
||||
|
|
@ -344,7 +425,7 @@ async def get_tool_usage_logs(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
where: Final[dict] = {"tool_name": tool_name}
|
||||
where: Final[dict[str, object]] = {"tool_name": tool_name}
|
||||
if start_date or end_date:
|
||||
start_time_filter: datetime | None = None
|
||||
end_time_filter: datetime | None = None
|
||||
|
|
@ -363,14 +444,14 @@ async def get_tool_usage_logs(
|
|||
except ValueError:
|
||||
pass
|
||||
if start_time_filter is not None or end_time_filter is not None:
|
||||
where["start_time"] = {}
|
||||
if start_time_filter is not None:
|
||||
where["start_time"]["gte"] = start_time_filter
|
||||
if end_time_filter is not None:
|
||||
where["start_time"]["lte"] = end_time_filter
|
||||
where["start_time"] = {
|
||||
key: value
|
||||
for key, value in (("gte", start_time_filter), ("lte", end_time_filter))
|
||||
if value is not None
|
||||
}
|
||||
|
||||
total: Final = await SpendLogToolIndexRepository(prisma_client).table.count(where=where)
|
||||
index_rows: Final = await SpendLogToolIndexRepository(prisma_client).table.find_many(
|
||||
total: Final = await _typed_table(SpendLogToolIndexRepository(prisma_client)).count(where=where)
|
||||
index_rows: Final = await _typed_table(SpendLogToolIndexRepository(prisma_client)).find_many(
|
||||
where=where,
|
||||
order={"start_time": "desc"},
|
||||
skip=(page - 1) * page_size,
|
||||
|
|
@ -380,7 +461,9 @@ async def get_tool_usage_logs(
|
|||
if not request_ids:
|
||||
return ToolUsageLogsResponse(logs=[], total=total, page=page, page_size=page_size)
|
||||
|
||||
spend_logs = await SpendLogsRepository(prisma_client).table.find_many(where={"request_id": {"in": request_ids}})
|
||||
spend_logs = await _typed_table(SpendLogsRepository(prisma_client)).find_many(
|
||||
where={"request_id": {"in": request_ids}}
|
||||
)
|
||||
log_by_id: Final = {s.request_id: s for s in spend_logs}
|
||||
|
||||
logs_out: Final[list[ToolUsageLogEntry]] = []
|
||||
|
|
@ -449,24 +532,24 @@ async def _resolve_key_hash_to_object_permission_id(
|
|||
hashed: Final = key_hash if "sk-" not in (key_hash or "") else hash_token(key_hash)
|
||||
if not hashed:
|
||||
return None
|
||||
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed})
|
||||
row = await _typed_table(VerificationTokenRepository(prisma_client)).find_unique(where={"token": hashed})
|
||||
if row is None:
|
||||
return None
|
||||
op_id: Final = getattr(row, "object_permission_id", None)
|
||||
op_id: Final = row.object_permission_id
|
||||
if op_id:
|
||||
return op_id
|
||||
new_id: Final = str(uuid.uuid4())
|
||||
await ObjectPermissionRepository(prisma_client).table.create(
|
||||
await _typed_table(ObjectPermissionRepository(prisma_client)).create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
)
|
||||
updated_count: Final = await VerificationTokenRepository(prisma_client).table.update_many(
|
||||
updated_count: Final = await _typed_table(VerificationTokenRepository(prisma_client)).update_many(
|
||||
where={"token": hashed, "object_permission_id": None},
|
||||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
await ObjectPermissionRepository(prisma_client).table.delete(where={"object_permission_id": new_id})
|
||||
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed})
|
||||
return getattr(row, "object_permission_id", None) if row else None
|
||||
await _typed_table(ObjectPermissionRepository(prisma_client)).delete(where={"object_permission_id": new_id})
|
||||
row = await _typed_table(VerificationTokenRepository(prisma_client)).find_unique(where={"token": hashed})
|
||||
return row.object_permission_id if row else None
|
||||
return new_id
|
||||
|
||||
|
||||
|
|
@ -478,24 +561,24 @@ async def _resolve_team_id_to_object_permission_id(
|
|||
if not team_id or not team_id.strip():
|
||||
return None
|
||||
team_id_clean: Final = team_id.strip()
|
||||
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id_clean})
|
||||
row = await _typed_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id_clean})
|
||||
if row is None:
|
||||
return None
|
||||
op_id: Final = getattr(row, "object_permission_id", None)
|
||||
op_id: Final = row.object_permission_id
|
||||
if op_id:
|
||||
return op_id
|
||||
new_id: Final = str(uuid.uuid4())
|
||||
await ObjectPermissionRepository(prisma_client).table.create(
|
||||
await _typed_table(ObjectPermissionRepository(prisma_client)).create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
)
|
||||
updated_count: Final = await TeamRepository(prisma_client).table.update_many(
|
||||
updated_count: Final = await _typed_table(TeamRepository(prisma_client)).update_many(
|
||||
where={"team_id": team_id_clean, "object_permission_id": None},
|
||||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
await ObjectPermissionRepository(prisma_client).table.delete(where={"object_permission_id": new_id})
|
||||
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id_clean})
|
||||
return getattr(row, "object_permission_id", None) if row else None
|
||||
await _typed_table(ObjectPermissionRepository(prisma_client)).delete(where={"object_permission_id": new_id})
|
||||
row = await _typed_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id_clean})
|
||||
return row.object_permission_id if row else None
|
||||
return new_id
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4,11 +4,11 @@ usage/spend data by querying the aggregated daily activity endpoints.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence
|
||||
from datetime import date
|
||||
from typing import Any, Final, Literal, cast
|
||||
from typing import Any, Final, Literal, Protocol, cast, overload
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -73,9 +73,36 @@ class SSEErrorEvent(TypedDict):
|
|||
SSEEvent = SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent
|
||||
|
||||
|
||||
class _EntityEntry(TypedDict, total=False):
|
||||
metrics: ReadOnly[Mapping[str, float]]
|
||||
metadata: ReadOnly[Mapping[str, str]]
|
||||
|
||||
|
||||
class _DayDump(TypedDict, total=False):
|
||||
breakdown: ReadOnly[Mapping[str, Mapping[str, _EntityEntry]]]
|
||||
|
||||
|
||||
class _UsageDump(Protocol):
|
||||
@overload
|
||||
def get(self, key: Literal["metadata"], default: Mapping[str, float], /) -> Mapping[str, float]: ...
|
||||
@overload
|
||||
def get(self, key: Literal["results"], default: Sequence[_DayDump], /) -> Sequence[_DayDump]: ...
|
||||
|
||||
|
||||
class _ToolFunctionDef(TypedDict):
|
||||
name: ReadOnly[str]
|
||||
description: ReadOnly[str]
|
||||
parameters: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _ToolDef(TypedDict):
|
||||
type: ReadOnly[str]
|
||||
function: ReadOnly[_ToolFunctionDef]
|
||||
|
||||
|
||||
class ToolHandler(TypedDict):
|
||||
fetch: Callable[..., Any]
|
||||
summarise: Callable[[dict[str, Any]], str]
|
||||
fetch: Callable[..., Awaitable[_UsageDump]]
|
||||
summarise: Callable[[_UsageDump], str]
|
||||
label: str
|
||||
|
||||
|
||||
|
|
@ -88,7 +115,7 @@ _DATE_PARAMS: Final = {
|
|||
"end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"},
|
||||
}
|
||||
|
||||
_TOOL_USAGE: Final = {
|
||||
_TOOL_USAGE: Final[_ToolDef] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_usage_data",
|
||||
|
|
@ -111,7 +138,7 @@ _TOOL_USAGE: Final = {
|
|||
},
|
||||
}
|
||||
|
||||
_TOOL_TEAM: Final = {
|
||||
_TOOL_TEAM: Final[_ToolDef] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_team_usage_data",
|
||||
|
|
@ -133,7 +160,7 @@ _TOOL_TEAM: Final = {
|
|||
},
|
||||
}
|
||||
|
||||
_TOOL_TAG: Final = {
|
||||
_TOOL_TAG: Final[_ToolDef] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_tag_usage_data",
|
||||
|
|
@ -159,7 +186,7 @@ TOOLS_BASE: Final = [_TOOL_USAGE]
|
|||
TOOLS_ADMIN: Final = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG]
|
||||
|
||||
|
||||
def get_tools_for_role(is_admin: bool) -> list[dict[str, Any]]:
|
||||
def get_tools_for_role(is_admin: bool) -> list[_ToolDef]:
|
||||
"""Return the tool list appropriate for the user's role."""
|
||||
return TOOLS_ADMIN if is_admin else TOOLS_BASE
|
||||
|
||||
|
|
@ -254,7 +281,7 @@ async def _query_activity(
|
|||
)
|
||||
|
||||
|
||||
async def _fetch_usage_data(start_date: str, end_date: str, user_id: str | None = None) -> dict[str, Any]:
|
||||
async def _fetch_usage_data(start_date: str, end_date: str, user_id: str | None = None) -> _UsageDump:
|
||||
resp: Final = await _query_activity(
|
||||
TABLE_DAILY_USER_SPEND,
|
||||
ENTITY_FIELD_USER,
|
||||
|
|
@ -266,7 +293,7 @@ async def _fetch_usage_data(start_date: str, end_date: str, user_id: str | None
|
|||
return resp.model_dump(mode="json")
|
||||
|
||||
|
||||
async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: str | None = None) -> dict[str, Any]:
|
||||
async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: str | None = None) -> _UsageDump:
|
||||
resp: Final = await _query_activity(
|
||||
TABLE_DAILY_TEAM_SPEND,
|
||||
ENTITY_FIELD_TEAM,
|
||||
|
|
@ -277,7 +304,7 @@ async def _fetch_team_usage_data(start_date: str, end_date: str, team_ids: str |
|
|||
return resp.model_dump(mode="json")
|
||||
|
||||
|
||||
async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: str | None = None) -> dict[str, Any]:
|
||||
async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: str | None = None) -> _UsageDump:
|
||||
resp: Final = await _query_activity(
|
||||
TABLE_DAILY_TAG_SPEND,
|
||||
ENTITY_FIELD_TAG,
|
||||
|
|
@ -294,7 +321,7 @@ async def _fetch_tag_usage_data(start_date: str, end_date: str, tags: str | None
|
|||
|
||||
|
||||
def _accumulate_breakdown(
|
||||
results: list[dict[str, Any]], dimension: str, fields: list[str]
|
||||
results: Sequence[_DayDump], dimension: str, fields: Sequence[str]
|
||||
) -> dict[str, dict[str, float]]:
|
||||
"""Aggregate a single breakdown dimension across days."""
|
||||
totals: Final[dict[str, dict[str, float]]] = {}
|
||||
|
|
@ -317,7 +344,7 @@ def _ranked_lines(
|
|||
return [fmt(name, vals) for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[:limit]]
|
||||
|
||||
|
||||
def _summarise_usage_data(data: dict[str, Any]) -> str:
|
||||
def _summarise_usage_data(data: _UsageDump) -> str:
|
||||
meta: Final = data.get("metadata", {})
|
||||
results: Final = data.get("results", [])
|
||||
|
||||
|
|
@ -349,7 +376,7 @@ def _summarise_usage_data(data: dict[str, Any]) -> str:
|
|||
return "\n".join(sections)
|
||||
|
||||
|
||||
def _summarise_entity_data(data: dict[str, Any], entity_label: str) -> str:
|
||||
def _summarise_entity_data(data: _UsageDump, entity_label: str) -> str:
|
||||
"""Summarise team/tag entity usage data."""
|
||||
results: Final = data.get("results", [])
|
||||
if not results:
|
||||
|
|
@ -409,16 +436,16 @@ def _sse(event: SSEEvent) -> str:
|
|||
|
||||
def _resolve_fetch_kwargs(
|
||||
fn_name: str,
|
||||
fn_args: dict[str, str],
|
||||
fn_args: Mapping[str, str],
|
||||
user_id: str | None,
|
||||
is_admin: bool,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, str]:
|
||||
"""Build keyword arguments for a tool's fetch function."""
|
||||
start_date: Final = fn_args.get("start_date", "")
|
||||
end_date: Final = fn_args.get("end_date", "")
|
||||
if not start_date or not end_date:
|
||||
raise ValueError("Missing required start_date or end_date from tool arguments")
|
||||
kwargs: Final[dict[str, Any]] = {"start_date": start_date, "end_date": end_date}
|
||||
kwargs: Final[dict[str, str]] = {"start_date": start_date, "end_date": end_date}
|
||||
if fn_name == "get_usage_data":
|
||||
if not is_admin:
|
||||
if user_id is None:
|
||||
|
|
@ -443,7 +470,7 @@ def _resolve_fetch_kwargs(
|
|||
async def _execute_tool_call(
|
||||
handler: ToolHandler,
|
||||
fn_name: str,
|
||||
fn_args: dict[str, str],
|
||||
fn_args: Mapping[str, str],
|
||||
user_id: str | None,
|
||||
is_admin: bool,
|
||||
) -> str:
|
||||
|
|
@ -455,13 +482,13 @@ async def _execute_tool_call(
|
|||
|
||||
async def _process_tool_call(
|
||||
tc: Any,
|
||||
chat_messages: list[dict[str, Any]],
|
||||
chat_messages: list[Mapping[str, object]],
|
||||
user_id: str | None,
|
||||
is_admin: bool,
|
||||
) -> AsyncIterator[str]:
|
||||
"""Execute a single tool call, yielding SSE events for status."""
|
||||
fn_name: Final = tc.function.name
|
||||
fn_args: Final = json.loads(tc.function.arguments)
|
||||
fn_name: Final[str] = tc.function.name
|
||||
fn_args: Final[Mapping[str, str]] = json.loads(tc.function.arguments)
|
||||
|
||||
allowed_names: Final = {t["function"]["name"] for t in get_tools_for_role(is_admin)}
|
||||
handler: Final = TOOL_HANDLERS.get(fn_name)
|
||||
|
|
@ -495,7 +522,7 @@ async def _process_tool_call(
|
|||
chat_messages.append({"role": "tool", "tool_call_id": tc.id, "content": tool_result})
|
||||
|
||||
|
||||
async def _stream_final_response(model: str, chat_messages: list[dict[str, Any]]) -> AsyncIterator[str]:
|
||||
async def _stream_final_response(model: str, chat_messages: list[Mapping[str, object]]) -> AsyncIterator[str]:
|
||||
"""Stream the final LLM response after tool results are appended."""
|
||||
yield _sse({"type": "status", "message": "Analyzing results..."})
|
||||
|
||||
|
|
@ -520,7 +547,7 @@ async def stream_usage_ai_chat(
|
|||
"""Stream SSE events: status → tool_call → chunk → done."""
|
||||
resolved_model: Final = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL
|
||||
truncated: Final = messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages
|
||||
chat_messages: Final[list[dict[str, Any]]] = [
|
||||
chat_messages: Final[list[Mapping[str, object]]] = [
|
||||
{"role": "system", "content": _build_system_prompt(is_admin)},
|
||||
*truncated,
|
||||
]
|
||||
|
|
|
|||
|
|
@ -11,11 +11,19 @@ These endpoints use optimized single SQL queries with joins to efficiently calcu
|
|||
user metrics from tag activity data and return time series for dashboard visualization.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final, Protocol, TypeVar, overload
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import LiteLLM_DailyTagSpend as PrismaDailyTagSpendRow
|
||||
from prisma.models import LiteLLM_UserTable as PrismaUserRow
|
||||
from prisma.models import LiteLLM_VerificationToken as PrismaVerificationTokenRow
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
|
@ -103,6 +111,54 @@ class PerUserAnalyticsResponse(BaseModel):
|
|||
total_pages: int
|
||||
|
||||
|
||||
class _DistinctTagRow(BaseModel):
|
||||
tag: str
|
||||
|
||||
|
||||
class _ActiveUsersRow(BaseModel):
|
||||
tag: str
|
||||
active_users: int
|
||||
date: str
|
||||
period_start: str | None = None
|
||||
period_end: str | None = None
|
||||
|
||||
|
||||
class _TagSummaryRow(BaseModel):
|
||||
tag: str
|
||||
unique_users: int | None = None
|
||||
total_requests: float | int | str | None = None
|
||||
successful_requests: float | int | str | None = None
|
||||
failed_requests: float | int | str | None = None
|
||||
total_tokens: float | int | str | None = None
|
||||
total_spend: float | int | str | None = None
|
||||
|
||||
|
||||
_DISTINCT_TAG_ROWS: Final = TypeAdapter(list[_DistinctTagRow])
|
||||
_ACTIVE_USERS_ROWS: Final = TypeAdapter(list[_ActiveUsersRow])
|
||||
_TAG_SUMMARY_ROWS: Final = TypeAdapter(list[_TagSummaryRow])
|
||||
|
||||
_RowT_co: Final = TypeVar("_RowT_co", covariant=True)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _TableOps(Protocol[_RowT_co]):
|
||||
async def find_many(self, where: Mapping[str, object] | None = None) -> Sequence[_RowT_co]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def _typed_table(repo: DailyTagSpendRepository) -> "_TableOps[PrismaDailyTagSpendRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: VerificationTokenRepository) -> "_TableOps[PrismaVerificationTokenRow]": ...
|
||||
@overload
|
||||
def _typed_table(repo: UserRepository) -> "_TableOps[PrismaUserRow]": ...
|
||||
def _typed_table(repo: DailyTagSpendRepository | VerificationTokenRepository | UserRepository) -> object:
|
||||
return repo.table
|
||||
|
||||
|
||||
async def _query_raw(prisma_client: "PrismaClient", sql_query: str, *params: object) -> object:
|
||||
return await prisma_client.db.query_raw(sql_query, *params)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/tag/distinct",
|
||||
response_model=DistinctTagsResponse,
|
||||
|
|
@ -141,9 +197,9 @@ async def get_distinct_user_agent_tags(
|
|||
LIMIT {MAX_TAGS}
|
||||
"""
|
||||
|
||||
db_response: Final = await prisma_client.db.query_raw(sql_query)
|
||||
db_response: Final = _DISTINCT_TAG_ROWS.validate_python(await _query_raw(prisma_client, sql_query))
|
||||
|
||||
results: Final = [DistinctTagResponse(tag=row["tag"]) for row in db_response]
|
||||
results: Final = [DistinctTagResponse(tag=row.tag) for row in db_response]
|
||||
|
||||
return DistinctTagsResponse(results=results)
|
||||
|
||||
|
|
@ -231,11 +287,10 @@ async def get_daily_active_users(
|
|||
ORDER BY dts.date DESC, active_users DESC
|
||||
"""
|
||||
|
||||
db_response: Final = await prisma_client.db.query_raw(sql_query, *params)
|
||||
db_response: Final = _ACTIVE_USERS_ROWS.validate_python(await _query_raw(prisma_client, sql_query, *params))
|
||||
|
||||
results: Final = [
|
||||
TagActiveUsersResponse(tag=row["tag"], active_users=row["active_users"], date=row["date"])
|
||||
for row in db_response
|
||||
TagActiveUsersResponse(tag=row.tag, active_users=row.active_users, date=row.date) for row in db_response
|
||||
]
|
||||
|
||||
return ActiveUsersAnalyticsResponse(results=results)
|
||||
|
|
@ -346,15 +401,15 @@ async def get_weekly_active_users(
|
|||
ORDER BY week_offset DESC, active_users DESC
|
||||
"""
|
||||
|
||||
db_response: Final = await prisma_client.db.query_raw(sql_query, *params)
|
||||
db_response: Final = _ACTIVE_USERS_ROWS.validate_python(await _query_raw(prisma_client, sql_query, *params))
|
||||
|
||||
results: Final = [
|
||||
TagActiveUsersResponse(
|
||||
tag=row["tag"],
|
||||
active_users=row["active_users"],
|
||||
date=row["date"], # This will be "Week 1 (Jan 15)", "Week 2 (Jan 8)", etc.
|
||||
period_start=row["period_start"],
|
||||
period_end=row["period_end"],
|
||||
tag=row.tag,
|
||||
active_users=row.active_users,
|
||||
date=row.date, # This will be "Week 1 (Jan 15)", "Week 2 (Jan 8)", etc.
|
||||
period_start=row.period_start,
|
||||
period_end=row.period_end,
|
||||
)
|
||||
for row in db_response
|
||||
]
|
||||
|
|
@ -467,15 +522,15 @@ async def get_monthly_active_users(
|
|||
ORDER BY month_offset DESC, active_users DESC
|
||||
"""
|
||||
|
||||
db_response: Final = await prisma_client.db.query_raw(sql_query, *params)
|
||||
db_response: Final = _ACTIVE_USERS_ROWS.validate_python(await _query_raw(prisma_client, sql_query, *params))
|
||||
|
||||
results: Final = [
|
||||
TagActiveUsersResponse(
|
||||
tag=row["tag"],
|
||||
active_users=row["active_users"],
|
||||
date=row["date"], # This will be "Month 1 (Jan)", "Month 2 (Dec)", etc.
|
||||
period_start=row["period_start"],
|
||||
period_end=row["period_end"],
|
||||
tag=row.tag,
|
||||
active_users=row.active_users,
|
||||
date=row.date, # This will be "Month 1 (Jan)", "Month 2 (Dec)", etc.
|
||||
period_start=row.period_start,
|
||||
period_end=row.period_end,
|
||||
)
|
||||
for row in db_response
|
||||
]
|
||||
|
|
@ -565,17 +620,17 @@ async def get_tag_summary(
|
|||
ORDER BY total_requests DESC
|
||||
"""
|
||||
|
||||
db_response: Final = await prisma_client.db.query_raw(sql_query, *params)
|
||||
db_response: Final = _TAG_SUMMARY_ROWS.validate_python(await _query_raw(prisma_client, sql_query, *params))
|
||||
|
||||
results: Final = [
|
||||
TagSummaryMetrics(
|
||||
tag=row["tag"],
|
||||
unique_users=row["unique_users"] or 0,
|
||||
total_requests=int(row["total_requests"] or 0),
|
||||
successful_requests=int(row["successful_requests"] or 0),
|
||||
failed_requests=int(row["failed_requests"] or 0),
|
||||
total_tokens=int(row["total_tokens"] or 0),
|
||||
total_spend=float(row["total_spend"] or 0.0),
|
||||
tag=row.tag,
|
||||
unique_users=row.unique_users or 0,
|
||||
total_requests=int(row.total_requests or 0),
|
||||
successful_requests=int(row.successful_requests or 0),
|
||||
failed_requests=int(row.failed_requests or 0),
|
||||
total_tokens=int(row.total_tokens or 0),
|
||||
total_spend=float(row.total_spend or 0.0),
|
||||
)
|
||||
for row in db_response
|
||||
]
|
||||
|
|
@ -648,7 +703,7 @@ async def get_per_user_analytics(
|
|||
start_date: Final = start_dt.strftime("%Y-%m-%d")
|
||||
|
||||
# Build where clause with date range
|
||||
where_clause: Final[dict[str, Any]] = {"date": {"gte": start_date, "lte": end_date}}
|
||||
where_clause: Final[dict[str, object]] = {"date": {"gte": start_date, "lte": end_date}}
|
||||
|
||||
# Add tag filtering if provided
|
||||
if tag_filters and len(tag_filters) > 0:
|
||||
|
|
@ -657,7 +712,7 @@ async def get_per_user_analytics(
|
|||
where_clause["tag"] = {"contains": tag_filter}
|
||||
|
||||
# Get all tag records in the date range with optional tag filtering
|
||||
tag_records: Final = await DailyTagSpendRepository(prisma_client).table.find_many(where=where_clause)
|
||||
tag_records: Final = await _typed_table(DailyTagSpendRepository(prisma_client)).find_many(where=where_clause)
|
||||
|
||||
# Get unique api_keys
|
||||
api_keys: Final = set(record.api_key for record in tag_records if record.api_key)
|
||||
|
|
@ -672,7 +727,7 @@ async def get_per_user_analytics(
|
|||
)
|
||||
|
||||
# Lookup user_id for each api_key
|
||||
api_key_records: Final = await VerificationTokenRepository(prisma_client).table.find_many(
|
||||
api_key_records: Final = await _typed_table(VerificationTokenRepository(prisma_client)).find_many(
|
||||
where={"token": {"in": list(api_keys)}}
|
||||
)
|
||||
|
||||
|
|
@ -681,7 +736,9 @@ async def get_per_user_analytics(
|
|||
|
||||
# Get user emails for the user_ids
|
||||
user_ids: Final = list(set(api_key_to_user_id.values()))
|
||||
user_records: Final = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": user_ids}})
|
||||
user_records: Final = await _typed_table(UserRepository(prisma_client)).find_many(
|
||||
where={"user_id": {"in": user_ids}}
|
||||
)
|
||||
|
||||
# Create mapping from user_id to user_email
|
||||
user_id_to_email: Final = {record.user_id: record.user_email for record in user_records}
|
||||
|
|
|
|||
|
|
@ -3,8 +3,10 @@ CRUD ENDPOINTS FOR PROMPTS
|
|||
"""
|
||||
|
||||
import tempfile
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
|
||||
|
||||
from fastapi import (
|
||||
APIRouter,
|
||||
|
|
@ -38,9 +40,68 @@ from litellm.types.prompts.init_prompts import (
|
|||
)
|
||||
from litellm.types.proxy.prompt_endpoints import TestPromptRequest
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.prompts.prompt_registry import InMemoryPromptRegistry
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
class _PromptRow(Protocol):
|
||||
@property
|
||||
def id(self) -> str: ...
|
||||
@property
|
||||
def prompt_id(self) -> str: ...
|
||||
@property
|
||||
def version(self) -> int: ...
|
||||
@property
|
||||
def environment(self) -> str: ...
|
||||
@property
|
||||
def created_by(self) -> str | None: ...
|
||||
@property
|
||||
def created_at(self) -> "datetime": ...
|
||||
@property
|
||||
def updated_at(self) -> "datetime": ...
|
||||
@property
|
||||
def litellm_params(self) -> str | Mapping[str, object]: ...
|
||||
@property
|
||||
def prompt_info(self) -> str | Mapping[str, object] | None: ...
|
||||
|
||||
def model_dump(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class _PromptRowData(BaseModel):
|
||||
prompt_id: str
|
||||
version: int = 1
|
||||
environment: str = "development"
|
||||
created_by: str | None = None
|
||||
litellm_params: str | Mapping[str, object] | None = None
|
||||
prompt_info: str | Mapping[str, object] | None = None
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class _PromptTableActions(Protocol):
|
||||
def find_many(
|
||||
self,
|
||||
*,
|
||||
where: Mapping[str, str | int],
|
||||
order: Mapping[str, str] = ...,
|
||||
take: int = ...,
|
||||
distinct: Sequence[str] = ...,
|
||||
) -> Awaitable[Sequence[_PromptRow]]: ...
|
||||
|
||||
def create(self, *, data: Mapping[str, str | int | None]) -> Awaitable[_PromptRow]: ...
|
||||
|
||||
def update(self, *, where: Mapping[str, str | int], data: Mapping[str, str]) -> Awaitable[_PromptRow]: ...
|
||||
|
||||
def delete_many(self, *, where: Mapping[str, str]) -> Awaitable[int]: ...
|
||||
|
||||
|
||||
def _prompt_table(prisma_client: "PrismaClient") -> _PromptTableActions:
|
||||
return PromptRepository(prisma_client).table
|
||||
|
||||
|
||||
def get_base_prompt_id(prompt_id: str) -> str:
|
||||
"""
|
||||
Extract the base prompt ID by stripping the version suffix if present.
|
||||
|
|
@ -132,7 +193,7 @@ def construct_versioned_prompt_id(prompt_id: str, version: int | None = None) ->
|
|||
return f"{base_id}.v{version}"
|
||||
|
||||
|
||||
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: dict[str, Any]) -> str:
|
||||
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Find the latest version of a prompt from available prompt IDs.
|
||||
|
||||
|
|
@ -198,7 +259,9 @@ def get_latest_prompt_versions(prompts: list[PromptSpec]) -> list[PromptSpec]:
|
|||
return list(latest_prompts.values())
|
||||
|
||||
|
||||
async def get_next_version_for_prompt(prisma_client, prompt_id: str, environment: str = "development") -> int:
|
||||
async def get_next_version_for_prompt(
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = "development"
|
||||
) -> int:
|
||||
"""
|
||||
Get the next version number for a prompt in a specific environment.
|
||||
|
||||
|
|
@ -210,7 +273,7 @@ async def get_next_version_for_prompt(prisma_client, prompt_id: str, environment
|
|||
Returns:
|
||||
Next version number (1 if no versions exist, max_version + 1 otherwise)
|
||||
"""
|
||||
existing_prompts: Final = await PromptRepository(prisma_client).table.find_many(
|
||||
existing_prompts: Final = await _prompt_table(prisma_client).find_many(
|
||||
where={"prompt_id": prompt_id, "environment": environment}
|
||||
)
|
||||
|
||||
|
|
@ -221,7 +284,7 @@ async def get_next_version_for_prompt(prisma_client, prompt_id: str, environment
|
|||
return 1
|
||||
|
||||
|
||||
def create_versioned_prompt_spec(db_prompt) -> PromptSpec:
|
||||
def create_versioned_prompt_spec(db_prompt: _PromptRow) -> PromptSpec:
|
||||
"""
|
||||
Helper function to create a PromptSpec with versioned prompt_id from a DB prompt entry.
|
||||
|
||||
|
|
@ -235,38 +298,33 @@ def create_versioned_prompt_spec(db_prompt) -> PromptSpec:
|
|||
|
||||
from litellm.types.prompts.init_prompts import PromptLiteLLMParams
|
||||
|
||||
prompt_dict: Final = db_prompt.model_dump()
|
||||
base_prompt_id: Final = prompt_dict["prompt_id"]
|
||||
version: Final = prompt_dict.get("version", 1)
|
||||
environment: Final = prompt_dict.get("environment", "development")
|
||||
created_by: Final = prompt_dict.get("created_by")
|
||||
row: Final = _PromptRowData.model_validate(db_prompt.model_dump())
|
||||
|
||||
# Parse litellm_params
|
||||
litellm_params_data = prompt_dict.get("litellm_params")
|
||||
if isinstance(litellm_params_data, str):
|
||||
litellm_params_data = json.loads(litellm_params_data)
|
||||
litellm_params: Final = PromptLiteLLMParams(**litellm_params_data)
|
||||
litellm_params_data: Final = row.litellm_params
|
||||
litellm_params_dict: Final[Mapping[str, object] | None] = (
|
||||
json.loads(litellm_params_data) if isinstance(litellm_params_data, str) else litellm_params_data
|
||||
)
|
||||
litellm_params: Final = PromptLiteLLMParams.model_validate(litellm_params_dict)
|
||||
|
||||
# Parse prompt_info
|
||||
prompt_info_data = prompt_dict.get("prompt_info")
|
||||
prompt_info_data: Final = row.prompt_info
|
||||
if prompt_info_data:
|
||||
if isinstance(prompt_info_data, str):
|
||||
prompt_info_data = json.loads(prompt_info_data)
|
||||
prompt_info = PromptInfo(**prompt_info_data)
|
||||
prompt_info_dict: Final[Mapping[str, object]] = (
|
||||
json.loads(prompt_info_data) if isinstance(prompt_info_data, str) else prompt_info_data
|
||||
)
|
||||
prompt_info = PromptInfo.model_validate(prompt_info_dict)
|
||||
else:
|
||||
prompt_info = PromptInfo(prompt_type="db")
|
||||
|
||||
# Create versioned prompt_id
|
||||
versioned_prompt_id: Final = f"{base_prompt_id}.v{version}"
|
||||
versioned_prompt_id: Final = f"{row.prompt_id}.v{row.version}"
|
||||
|
||||
return PromptSpec(
|
||||
prompt_id=versioned_prompt_id,
|
||||
litellm_params=litellm_params,
|
||||
prompt_info=prompt_info,
|
||||
created_at=prompt_dict.get("created_at"),
|
||||
updated_at=prompt_dict.get("updated_at"),
|
||||
environment=environment,
|
||||
created_by=created_by,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
environment=row.environment,
|
||||
created_by=row.created_by,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -431,10 +489,10 @@ async def get_prompt_versions(
|
|||
# Query DB for versions
|
||||
versioned_prompts: Final = []
|
||||
if prisma_client is not None:
|
||||
where_clause: Final[dict[str, Any]] = {"prompt_id": base_prompt_id}
|
||||
where_clause: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
|
||||
if environment:
|
||||
where_clause["environment"] = environment
|
||||
db_prompts: Final = await PromptRepository(prisma_client).table.find_many(
|
||||
db_prompts: Final = await _prompt_table(prisma_client).find_many(
|
||||
where=where_clause,
|
||||
order={"version": "desc"},
|
||||
)
|
||||
|
|
@ -590,7 +648,7 @@ async def get_prompt_info(
|
|||
# Query all environments this prompt exists in (lightweight: distinct on environment)
|
||||
all_environments: list[str] = []
|
||||
if prisma_client is not None:
|
||||
all_prompt_rows: Final = await PromptRepository(prisma_client).table.find_many(
|
||||
all_prompt_rows: Final = await _prompt_table(prisma_client).find_many(
|
||||
where={"prompt_id": base_prompt_id},
|
||||
distinct=["environment"],
|
||||
)
|
||||
|
|
@ -602,13 +660,13 @@ async def get_prompt_info(
|
|||
prompt_spec = None
|
||||
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
|
||||
if environment and prisma_client is not None:
|
||||
where_clause: Final[dict[str, Any]] = {
|
||||
where_clause: Final[dict[str, str | int]] = {
|
||||
"prompt_id": base_prompt_id,
|
||||
"environment": environment,
|
||||
}
|
||||
if requested_version is not None:
|
||||
where_clause["version"] = requested_version
|
||||
env_prompts: Final = await PromptRepository(prisma_client).table.find_many(
|
||||
env_prompts: Final = await _prompt_table(prisma_client).find_many(
|
||||
where=where_clause,
|
||||
order={"version": "desc"},
|
||||
take=1,
|
||||
|
|
@ -721,7 +779,7 @@ async def create_prompt(
|
|||
)
|
||||
|
||||
# Store prompt in db with version
|
||||
prompt_db_entry: Final = await PromptRepository(prisma_client).table.create(
|
||||
prompt_db_entry: Final = await _prompt_table(prisma_client).create(
|
||||
data={
|
||||
"prompt_id": request.prompt_id,
|
||||
"version": new_version,
|
||||
|
|
@ -811,7 +869,7 @@ async def update_prompt(
|
|||
)
|
||||
|
||||
# Check if any version of this prompt exists (in any environment)
|
||||
existing_prompts = await PromptRepository(prisma_client).table.find_many(where={"prompt_id": base_prompt_id})
|
||||
existing_prompts = await _prompt_table(prisma_client).find_many(where={"prompt_id": base_prompt_id})
|
||||
|
||||
if not existing_prompts:
|
||||
raise HTTPException(
|
||||
|
|
@ -835,7 +893,7 @@ async def update_prompt(
|
|||
)
|
||||
|
||||
# Store new version in db
|
||||
prompt_db_entry: Final = await PromptRepository(prisma_client).table.create(
|
||||
prompt_db_entry: Final = await _prompt_table(prisma_client).create(
|
||||
data={
|
||||
"prompt_id": base_prompt_id,
|
||||
"version": new_version,
|
||||
|
|
@ -936,12 +994,12 @@ async def delete_prompt(
|
|||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Build delete filter; scope to environment if provided
|
||||
delete_where: Final[dict[str, Any]] = {"prompt_id": base_prompt_id}
|
||||
delete_where: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
|
||||
if environment:
|
||||
delete_where["environment"] = environment
|
||||
|
||||
# Delete versions from the database (scoped to environment if provided)
|
||||
await PromptRepository(prisma_client).table.delete_many(where=delete_where)
|
||||
await _prompt_table(prisma_client).delete_many(where=delete_where)
|
||||
|
||||
# Remove matching prompts from memory — scope to environment if provided
|
||||
if environment:
|
||||
|
|
@ -967,7 +1025,9 @@ async def delete_prompt(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _reload_prompt_in_registry(registry: Any, versioned_id: str, updated_prompt_spec: PromptSpec) -> PromptSpec:
|
||||
def _reload_prompt_in_registry(
|
||||
registry: "InMemoryPromptRegistry", versioned_id: str, updated_prompt_spec: PromptSpec
|
||||
) -> PromptSpec:
|
||||
"""Remove stale entry and re-initialize the prompt in the in-memory registry."""
|
||||
if versioned_id in registry.IN_MEMORY_PROMPTS:
|
||||
del registry.IN_MEMORY_PROMPTS[versioned_id]
|
||||
|
|
@ -1033,14 +1093,14 @@ async def patch_prompt(
|
|||
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
|
||||
|
||||
# Build query to find the exact row by composite unique key
|
||||
find_where: Final[dict[str, Any]] = {
|
||||
find_where: Final[dict[str, str | int]] = {
|
||||
"prompt_id": base_prompt_id,
|
||||
"environment": env,
|
||||
}
|
||||
if requested_version is not None:
|
||||
find_where["version"] = requested_version
|
||||
|
||||
db_rows: Final = await PromptRepository(prisma_client).table.find_many(
|
||||
db_rows: Final = await _prompt_table(prisma_client).find_many(
|
||||
where=find_where,
|
||||
order={"version": "desc"},
|
||||
take=1,
|
||||
|
|
@ -1084,7 +1144,7 @@ async def patch_prompt(
|
|||
raise HTTPException(status_code=400, detail="litellm_params cannot be None")
|
||||
|
||||
# Build update data dict
|
||||
update_data: Final[dict[str, Any]] = {
|
||||
update_data: Final[dict[str, str]] = {
|
||||
"litellm_params": updated_litellm_params.model_dump_json(),
|
||||
"prompt_info": updated_prompt_info.model_dump_json(),
|
||||
}
|
||||
|
|
@ -1092,7 +1152,7 @@ async def patch_prompt(
|
|||
update_data["created_by"] = user_api_key_dict.user_id
|
||||
|
||||
# Update by primary key (id) to target exactly one row
|
||||
updated_prompt_db_entry: Final = await PromptRepository(prisma_client).table.update(
|
||||
updated_prompt_db_entry: Final = await _prompt_table(prisma_client).update(
|
||||
where={"id": target_row.id},
|
||||
data=update_data,
|
||||
)
|
||||
|
|
@ -1216,7 +1276,7 @@ async def test_prompt(
|
|||
|
||||
# Use ProxyBaseLLMRequestProcessing to go through all proxy logic
|
||||
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
result: Final = await base_llm_response_processor.base_process_llm_request(
|
||||
result: Final[object] = await base_llm_response_processor.base_process_llm_request(
|
||||
request=fastapi_request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -2329,8 +2329,11 @@ def load_from_azure_key_vault(use_azure_key_vault: bool = False):
|
|||
def cost_tracking():
|
||||
global prisma_client
|
||||
if prisma_client is not None:
|
||||
from litellm.integrations.shadow_eval_logger import ShadowEvalLogger
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger())
|
||||
litellm.logging_callback_manager.add_litellm_async_success_callback(_ProxyDBLogger())
|
||||
litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger())
|
||||
|
||||
|
||||
# Bounds authoritative DB re-reads when enforcing a budget against a
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from importlib.resources import files
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -32,14 +34,66 @@ from litellm.types.proxy.public_endpoints.public_endpoints import (
|
|||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from datetime import datetime
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
class _ProviderSupportEntry(TypedDict, total=False):
|
||||
display_name: ReadOnly[str]
|
||||
endpoints: ReadOnly[Mapping[str, bool]]
|
||||
|
||||
|
||||
class _ProvidersFile(TypedDict, total=False):
|
||||
providers: ReadOnly[Mapping[str, _ProviderSupportEntry]]
|
||||
|
||||
|
||||
class _EndpointProviderEntry(TypedDict):
|
||||
slug: ReadOnly[str]
|
||||
display_name: ReadOnly[str]
|
||||
|
||||
|
||||
class _EndpointEntry(TypedDict):
|
||||
key: ReadOnly[str]
|
||||
label: ReadOnly[str]
|
||||
endpoint: ReadOnly[str]
|
||||
providers: ReadOnly[Sequence[_EndpointProviderEntry]]
|
||||
|
||||
|
||||
class _PluginRow(Protocol):
|
||||
@property
|
||||
def id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def name(self) -> str: ...
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool: ...
|
||||
|
||||
@property
|
||||
def created_at(self) -> "datetime | None": ...
|
||||
|
||||
@property
|
||||
def updated_at(self) -> "datetime | None": ...
|
||||
|
||||
@property
|
||||
def manifest_json(self) -> str | None: ...
|
||||
|
||||
|
||||
class _PluginTableActions(Protocol):
|
||||
def find_many(self, *, where: Mapping[str, bool]) -> Awaitable[Sequence[_PluginRow]]: ...
|
||||
|
||||
|
||||
def _plugin_table(prisma_client: object) -> _PluginTableActions:
|
||||
return ClaudeCodePluginRepository(prisma_client).table
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /public/endpoints — helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_ENDPOINT_METADATA: Final[dict[str, dict[str, str]]] = {
|
||||
_ENDPOINT_METADATA: Final[Mapping[str, Mapping[str, str]]] = {
|
||||
"chat_completions": {"label": "Chat Completions", "endpoint": "/chat/completions"},
|
||||
"messages": {"label": "Messages", "endpoint": "/messages"},
|
||||
"responses": {"label": "Responses", "endpoint": "/responses"},
|
||||
|
|
@ -108,12 +162,12 @@ def _clean_display_name(raw: str) -> str:
|
|||
return _SLUG_SUFFIX_RE.sub("", raw).strip()
|
||||
|
||||
|
||||
def _build_endpoints(raw: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
def _build_endpoints(raw: _ProvidersFile) -> list[_EndpointEntry]:
|
||||
"""Transform raw provider_endpoints_support_backup.json into the response shape."""
|
||||
providers: Final[dict[str, Any]] = raw.get("providers", {})
|
||||
providers: Final = raw.get("providers", {})
|
||||
|
||||
# Collect endpoint keys in insertion order (union across all providers).
|
||||
seen: Final[set] = set()
|
||||
seen: Final[set[str]] = set()
|
||||
all_keys: Final[list[str]] = []
|
||||
for provider_data in providers.values():
|
||||
for key in provider_data.get("endpoints", {}):
|
||||
|
|
@ -121,13 +175,13 @@ def _build_endpoints(raw: dict[str, Any]) -> list[dict[str, Any]]:
|
|||
seen.add(key)
|
||||
all_keys.append(key)
|
||||
|
||||
result: Final[list[dict[str, Any]]] = []
|
||||
result: Final[list[_EndpointEntry]] = []
|
||||
for key in all_keys:
|
||||
meta = _ENDPOINT_METADATA.get(key)
|
||||
label = meta["label"] if meta else key.replace("_", " ").title()
|
||||
path = meta["endpoint"] if meta else "/" + key.replace("_", "/")
|
||||
|
||||
supporting: list[dict[str, str]] = [
|
||||
supporting: list[_EndpointProviderEntry] = [
|
||||
{
|
||||
"slug": slug,
|
||||
"display_name": _clean_display_name(pd.get("display_name", slug)),
|
||||
|
|
@ -140,8 +194,10 @@ def _build_endpoints(raw: dict[str, Any]) -> list[dict[str, Any]]:
|
|||
return result
|
||||
|
||||
|
||||
def _load_endpoints() -> list[dict[str, Any]]:
|
||||
raw = json.loads(files("litellm").joinpath("provider_endpoints_support_backup.json").read_text(encoding="utf-8"))
|
||||
def _load_endpoints() -> list[_EndpointEntry]:
|
||||
raw: Final[_ProvidersFile] = json.loads(
|
||||
files("litellm").joinpath("provider_endpoints_support_backup.json").read_text(encoding="utf-8")
|
||||
)
|
||||
return _build_endpoints(raw)
|
||||
|
||||
|
||||
|
|
@ -235,12 +291,7 @@ async def get_mcp_servers():
|
|||
)
|
||||
|
||||
public_mcp_servers: Final = global_mcp_server_manager.get_public_mcp_servers()
|
||||
return [
|
||||
MCPPublicServer(
|
||||
**server.model_dump(),
|
||||
)
|
||||
for server in public_mcp_servers
|
||||
]
|
||||
return [MCPPublicServer.model_validate(server.model_dump()) for server in public_mcp_servers]
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -259,7 +310,7 @@ async def public_skill_hub():
|
|||
|
||||
try:
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
plugins: Final = await ClaudeCodePluginRepository(prisma_client).table.find_many(where={"enabled": True})
|
||||
plugins: Final = await _plugin_table(prisma_client).find_many(where={"enabled": True})
|
||||
items: Final = []
|
||||
for plugin in plugins:
|
||||
raw = plugin.manifest_json or {}
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ Provides:
|
|||
"""
|
||||
|
||||
import base64
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
|
|
@ -31,6 +32,9 @@ from litellm.proxy.vector_store_endpoints.utils import (
|
|||
)
|
||||
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -58,7 +62,7 @@ def _append_payload_to_scan_stack(
|
|||
payload_stack.append((value, next_depth))
|
||||
|
||||
|
||||
def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
|
||||
def _collect_vector_store_ids_from_payload(payload: object) -> set[str]:
|
||||
vector_store_ids: Final[set[str]] = set()
|
||||
payload_stack: Final = [(payload, 0)]
|
||||
|
||||
|
|
@ -95,7 +99,7 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
|
|||
|
||||
|
||||
async def _authorize_nested_vector_store_ids(
|
||||
payload: Any,
|
||||
payload: object,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
for vector_store_id in sorted(_collect_vector_store_ids_from_payload(payload)):
|
||||
|
|
@ -109,7 +113,7 @@ def _build_file_metadata_entry(
|
|||
response: Any,
|
||||
file_data: tuple[str, bytes, str] | None = None,
|
||||
file_url: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> Mapping[str, str | int | None]:
|
||||
"""
|
||||
Build a file metadata entry for storing in vector_store_metadata.
|
||||
|
||||
|
|
@ -159,8 +163,8 @@ def _build_file_metadata_entry(
|
|||
|
||||
async def _save_vector_store_to_db_from_rag_ingest(
|
||||
response: Any,
|
||||
ingest_options: dict[str, Any],
|
||||
prisma_client,
|
||||
ingest_options: Mapping[str, dict[str, str | None]],
|
||||
prisma_client: "PrismaClient",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
file_data: tuple[str, bytes, str] | None = None,
|
||||
file_url: str | None = None,
|
||||
|
|
@ -299,9 +303,9 @@ async def parse_rag_ingest_request(
|
|||
headers: Final = _safe_get_request_headers(request)
|
||||
content_type = headers.get("content-type", "")
|
||||
|
||||
file_data = None
|
||||
file_url = None
|
||||
file_id = None
|
||||
file_data: tuple[str, bytes, str] | None = None
|
||||
file_url: str | None = None
|
||||
file_id: str | None = None
|
||||
ingest_options: dict[str, Any] = {}
|
||||
|
||||
if "multipart/form-data" in content_type:
|
||||
|
|
@ -315,7 +319,7 @@ async def parse_rag_ingest_request(
|
|||
file_data = (file_obj.filename, file_content, file_obj.content_type)
|
||||
|
||||
# Parse JSON from 'request' form field (contains full request body as JSON)
|
||||
request_json_str: Final = form_data.get("request")
|
||||
request_json_str: Final[str | bytes | None] = form_data.get("request")
|
||||
if request_json_str:
|
||||
request_data: Final = orjson.loads(request_json_str)
|
||||
ingest_options = request_data.get("ingest_options", {})
|
||||
|
|
@ -382,7 +386,7 @@ async def parse_rag_ingest_request(
|
|||
"api_key",
|
||||
"api_base",
|
||||
}
|
||||
vector_store_opts: Final = ingest_options.get("vector_store", {})
|
||||
vector_store_opts: Final[object] = ingest_options.get("vector_store", {})
|
||||
if isinstance(vector_store_opts, dict):
|
||||
for field in _BLOCKED_VECTOR_STORE_CREDENTIAL_PARAMS:
|
||||
if field in vector_store_opts:
|
||||
|
|
@ -658,7 +662,7 @@ async def rag_query(
|
|||
)
|
||||
|
||||
# Add litellm data
|
||||
request_data: dict[str, Any] = {}
|
||||
request_data: dict[str, object] = {}
|
||||
request_data = await add_litellm_data_to_request(
|
||||
data=request_data,
|
||||
request=request,
|
||||
|
|
|
|||
|
|
@ -10,9 +10,10 @@ https://platform.openai.com/docs/api-reference/responses-streaming
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from fastapi import Request, Response
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
|
|
@ -20,25 +21,30 @@ from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessin
|
|||
from litellm.proxy.response_polling.polling_handler import ResponsePollingHandler
|
||||
from litellm.types.llms.openai import ResponsesAPIStatus
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
async def background_streaming_task(
|
||||
polling_id: str,
|
||||
data: dict,
|
||||
data,
|
||||
polling_handler: ResponsePollingHandler,
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: dict,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
general_settings,
|
||||
llm_router: "Router | None",
|
||||
proxy_config: "ProxyConfig",
|
||||
proxy_logging_obj: "ProxyLogging",
|
||||
select_data_generator,
|
||||
user_model,
|
||||
user_temperature,
|
||||
user_request_timeout,
|
||||
user_max_tokens,
|
||||
user_api_base,
|
||||
version,
|
||||
user_temperature: float | None,
|
||||
user_request_timeout: float | None,
|
||||
user_max_tokens: int | None,
|
||||
user_api_base: str | None,
|
||||
version: str | None,
|
||||
):
|
||||
"""
|
||||
Background task to stream response and update cache
|
||||
|
|
@ -69,7 +75,7 @@ async def background_streaming_task(
|
|||
# Make streaming request.
|
||||
# Pre-call checks (rate limits, guardrails, budget) were already run
|
||||
# before polling ID creation, so skip them here to avoid double-counting.
|
||||
response: Final = await processor.base_process_llm_request(
|
||||
response: Final[StreamingResponse] = await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -1450,6 +1450,44 @@ model LiteLLM_AutoRouterSession {
|
|||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic.
|
||||
// A sampled slice of requests is duplicated through the router in a detached task and an
|
||||
// LLM judge compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// stopped_at; every count, status, and spend figure is derived from the append-only
|
||||
// attempt rows, so nothing can disagree across pods or stop races.
|
||||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
api_key_id String // hashed virtual key whose traffic is shadowed
|
||||
router_name String
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample budget: judge at most this many turns
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
ends_at DateTime
|
||||
stopped_at DateTime?
|
||||
|
||||
@@index([api_key_id])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
// One row per sampled pipeline: a blind verdict (real | shadow | tie) or an error.
|
||||
model LiteLLM_ShadowEvalAttempt {
|
||||
id String @id @default(cuid())
|
||||
job_id String
|
||||
request_id String // the judged real request
|
||||
outcome String // real | shadow | tie | error
|
||||
tier String? // router's tier for the prompt, when classified
|
||||
real_model String?
|
||||
shadow_model String?
|
||||
confidence Float?
|
||||
judge_cost Float @default(0)
|
||||
error String?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
@@index([job_id])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
|
|
|
|||
|
|
@ -14,16 +14,19 @@ Flow:
|
|||
import json
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, Final, cast
|
||||
from collections.abc import Iterable, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeAlias, cast
|
||||
|
||||
from litellm._internal_context import is_internal_call
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.openai import ResponseOutputItem, ResponsesAPIResponse
|
||||
from litellm.types.vector_stores import VectorStoreSearchResult
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
|
||||
# Keep ToolParam broad so we stay compatible with both dict and Pydantic forms
|
||||
ToolParam = Any
|
||||
ToolParam: TypeAlias = object
|
||||
|
||||
FILE_SEARCH_FUNCTION_NAME: Final = "litellm_file_search"
|
||||
|
||||
|
|
@ -35,7 +38,7 @@ FILE_SEARCH_FUNCTION_NAME: Final = "litellm_file_search"
|
|||
|
||||
def should_use_emulated_file_search(
|
||||
tools: Iterable[ToolParam] | None,
|
||||
provider_config: Any, # BaseResponsesAPIConfig
|
||||
provider_config: "BaseResponsesAPIConfig | None",
|
||||
) -> bool:
|
||||
"""Return True when there is a file_search tool and the provider can't handle it natively."""
|
||||
if not tools:
|
||||
|
|
@ -51,7 +54,7 @@ def should_use_emulated_file_search(
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _build_function_tool(vector_store_ids: list[str]) -> dict[str, Any]:
|
||||
def _build_function_tool(vector_store_ids: list[str]) -> dict[str, object]:
|
||||
"""
|
||||
Create a Responses API function-tool definition that describes file search.
|
||||
The function accepts one or more natural-language queries (like OpenAI's native
|
||||
|
|
@ -96,14 +99,14 @@ def _build_function_tool(vector_store_ids: list[str]) -> dict[str, Any]:
|
|||
|
||||
def _replace_file_search_tools(
|
||||
tools: Iterable[ToolParam] | None,
|
||||
) -> tuple[list[dict[str, Any]], list[str]]:
|
||||
) -> tuple[list[object], list[str]]:
|
||||
"""
|
||||
Replace all file_search tools with a single function tool.
|
||||
|
||||
Returns:
|
||||
(new_tools_list, all_vector_store_ids)
|
||||
"""
|
||||
non_file_search: Final[list[dict[str, Any]]] = []
|
||||
non_file_search: Final[list[object]] = []
|
||||
vector_store_ids: Final[list[str]] = []
|
||||
|
||||
for tool in tools or []:
|
||||
|
|
@ -172,7 +175,7 @@ async def _run_vector_searches(
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_field(result: Any, key: str, default: Any = None) -> Any:
|
||||
def _get_field(result: object, key: str, default: object = None) -> Any:
|
||||
"""Read a field from either a dict/TypedDict or an attribute-based object."""
|
||||
if isinstance(result, dict):
|
||||
return result.get(key, default)
|
||||
|
|
@ -211,7 +214,7 @@ def _format_search_results_as_tool_output(
|
|||
|
||||
def _build_search_results_for_include(
|
||||
results: list[VectorStoreSearchResult],
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[dict[str, object]]:
|
||||
"""
|
||||
Convert VectorStoreSearchResult objects to the format expected in
|
||||
file_search_call.search_results (mirrors OpenAI's include= format).
|
||||
|
|
@ -220,7 +223,7 @@ def _build_search_results_for_include(
|
|||
behaviour of OpenAI's native file_search which surfaces every relevant
|
||||
chunk even when multiple chunks originate from the same document.
|
||||
"""
|
||||
formatted: Final[list[dict[str, Any]]] = []
|
||||
formatted: Final[list[dict[str, object]]] = []
|
||||
for result in results:
|
||||
file_id = _get_field(result, "file_id") or ""
|
||||
content_items = _get_field(result, "content") or []
|
||||
|
|
@ -243,7 +246,7 @@ def _build_file_search_call_output(
|
|||
queries: list[str],
|
||||
results: list[VectorStoreSearchResult] | None = None,
|
||||
include_search_results: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Build the file_search_call output item (mirrors OpenAI's format).
|
||||
|
||||
Args:
|
||||
|
|
@ -268,14 +271,14 @@ def _build_file_search_call_output(
|
|||
def _build_file_citation_annotations(
|
||||
results: list[VectorStoreSearchResult],
|
||||
text: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[dict[str, object]]:
|
||||
"""
|
||||
Build file_citation annotations for the text.
|
||||
Each result with a file_id gets a citation at the end of the text.
|
||||
"""
|
||||
annotations: Final[list[dict[str, Any]]] = []
|
||||
annotations: Final[list[dict[str, object]]] = []
|
||||
index: Final = len(text) # cite at end of text block
|
||||
seen_file_ids: Final[set] = set()
|
||||
seen_file_ids: Final[set[object]] = set()
|
||||
|
||||
for result in results:
|
||||
file_id = _get_field(result, "file_id")
|
||||
|
|
@ -298,7 +301,7 @@ def _build_file_citation_annotations(
|
|||
def _build_message_output(
|
||||
response_text: str,
|
||||
results: list[VectorStoreSearchResult],
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Build the message output item with optional file_citation annotations."""
|
||||
annotations: Final = _build_file_citation_annotations(results, response_text)
|
||||
return {
|
||||
|
|
@ -330,8 +333,8 @@ def _extract_text_from_responses_output(response: ResponsesAPIResponse) -> str:
|
|||
|
||||
def _synthesize_responses_api_response(
|
||||
original_response: ResponsesAPIResponse,
|
||||
file_search_call_output: dict[str, Any],
|
||||
message_output: dict[str, Any],
|
||||
file_search_call_output: dict[str, object],
|
||||
message_output: dict[str, object],
|
||||
first_response: ResponsesAPIResponse | None = None,
|
||||
) -> ResponsesAPIResponse:
|
||||
"""
|
||||
|
|
@ -343,7 +346,7 @@ def _synthesize_responses_api_response(
|
|||
synthesized _hidden_params so that billing callbacks see the total cost of
|
||||
both provider calls that the emulated flow makes.
|
||||
"""
|
||||
synthesized_output: Final[list[dict[str, Any]]] = [file_search_call_output, message_output]
|
||||
synthesized_output: Final[list[dict[str, object]]] = [file_search_call_output, message_output]
|
||||
synthesized: Final = ResponsesAPIResponse(
|
||||
id=getattr(original_response, "id", f"resp_{uuid.uuid4().hex}"),
|
||||
object="response",
|
||||
|
|
@ -383,12 +386,12 @@ async def _call_aresponses(input, model, tools, **kwargs): # pragma: no cover
|
|||
|
||||
def _prepare_emulated_file_search_call(
|
||||
kwargs: dict[str, Any],
|
||||
) -> tuple[bool, dict[str, Any]]:
|
||||
) -> tuple[bool, dict[str, object]]:
|
||||
include_items: Final[list[str]] = list(kwargs.get("include") or [])
|
||||
include_search_results: Final = "file_search_call.results" in include_items
|
||||
|
||||
original_stream: Final = kwargs.get("stream")
|
||||
updated_kwargs = kwargs
|
||||
updated_kwargs: dict[str, object] = kwargs
|
||||
if original_stream:
|
||||
verbose_logger.debug(
|
||||
"Streaming is not yet supported for emulated file_search. Disabling stream for this request."
|
||||
|
|
@ -398,7 +401,7 @@ def _prepare_emulated_file_search_call(
|
|||
return include_search_results, updated_kwargs
|
||||
|
||||
|
||||
def _extract_tool_call_fields(tool_call: Any, fallback_call_id: str) -> tuple[str, str]:
|
||||
def _extract_tool_call_fields(tool_call: object, fallback_call_id: str) -> tuple[str, str]:
|
||||
"""Extract (call_id, raw_arguments_string) from a dict or Pydantic tool_call item."""
|
||||
if isinstance(tool_call, dict):
|
||||
call_id = str(tool_call.get("call_id") or tool_call.get("id") or fallback_call_id)
|
||||
|
|
@ -410,7 +413,7 @@ def _extract_tool_call_fields(tool_call: Any, fallback_call_id: str) -> tuple[st
|
|||
return call_id, raw_args
|
||||
|
||||
|
||||
def _resolve_queries_from_args(args: dict[str, Any], input: Any) -> list[str]:
|
||||
def _resolve_queries_from_args(args: dict[str, Any], input: object) -> list[str]:
|
||||
"""Pull the queries list out of parsed tool-call arguments, with backward-compat fallbacks."""
|
||||
queries_from_call: Final = args.get("queries")
|
||||
if not queries_from_call:
|
||||
|
|
@ -423,13 +426,13 @@ def _resolve_queries_from_args(args: dict[str, Any], input: Any) -> list[str]:
|
|||
|
||||
|
||||
async def _execute_file_search_tool_calls(
|
||||
file_search_calls: list[Any],
|
||||
file_search_calls: Sequence[object],
|
||||
all_vs_ids: list[str],
|
||||
input: Any,
|
||||
input: object,
|
||||
file_search_call_id: str,
|
||||
) -> tuple[list[dict[str, Any]], list[str], list[VectorStoreSearchResult]]:
|
||||
) -> tuple[list[object], list[str], list[VectorStoreSearchResult]]:
|
||||
"""Run the vector search for each file_search tool_call and collect results."""
|
||||
tool_results: Final[list[dict[str, Any]]] = []
|
||||
tool_results: Final[list[object]] = []
|
||||
all_queries: Final[list[str]] = []
|
||||
all_results: Final[list[VectorStoreSearchResult]] = []
|
||||
|
||||
|
|
@ -465,17 +468,17 @@ async def _execute_file_search_tool_calls(
|
|||
|
||||
|
||||
def _build_follow_up_input(
|
||||
input: Any,
|
||||
input: object,
|
||||
first_response: ResponsesAPIResponse,
|
||||
tool_results: list[dict[str, Any]],
|
||||
) -> list[Any]:
|
||||
tool_results: list[object],
|
||||
) -> list[object]:
|
||||
"""Assemble the follow-up call input: original messages + first-response output + tool results.
|
||||
|
||||
Including all output items (text blocks, reasoning, non-file-search calls) ensures providers
|
||||
like Anthropic that emit text before the tool call have complete conversation context.
|
||||
Serializes Pydantic model instances to plain dicts so the transformation layer can call .get().
|
||||
"""
|
||||
original_input_items: Final = (
|
||||
original_input_items: Final[list[object]] = (
|
||||
list(input) if isinstance(input, (list, tuple)) else [{"role": "user", "content": str(input)}]
|
||||
)
|
||||
first_response_output_items: Final[list[Any]] = []
|
||||
|
|
@ -491,7 +494,7 @@ def _build_follow_up_input(
|
|||
|
||||
|
||||
async def aresponses_with_emulated_file_search(
|
||||
input: Any,
|
||||
input: object,
|
||||
model: str,
|
||||
tools: Iterable[ToolParam] | None = None,
|
||||
# Pass-through params — forwarded as-is to the underlying aresponses call
|
||||
|
|
|
|||
|
|
@ -68,7 +68,7 @@ async def create_mcp_list_tools_events(
|
|||
# Convert tools to dict format for the event
|
||||
_mcp_tools_dict: Final = [
|
||||
tool.model_dump()
|
||||
if hasattr(tool, "model_dump") and callable(getattr(tool, "model_dump"))
|
||||
if hasattr(tool, "model_dump") and callable(getattr(tool, "model_dump", None))
|
||||
else tool.__dict__
|
||||
if hasattr(tool, "__dict__")
|
||||
else {"name": getattr(tool, "name", str(tool))}
|
||||
|
|
@ -356,7 +356,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers_obj)
|
||||
|
||||
# Also check if headers are provided in tools array (from request body)
|
||||
tools: Final = self.original_request_params.get("tools")
|
||||
tools: Final[Sequence[object] | None] = self.original_request_params.get("tools")
|
||||
if tools:
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict) and tool.get("type") == "mcp":
|
||||
|
|
@ -395,7 +395,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
def _make_stream_error_event(self) -> ResponsesAPIStreamingResponse:
|
||||
err: Final = self._stream_error
|
||||
status_code: Final = getattr(err, "status_code", None)
|
||||
status_code: Final[object] = getattr(err, "status_code", None)
|
||||
return ErrorEvent(
|
||||
type=ResponsesAPIStreamEvents.ERROR,
|
||||
sequence_number=self._last_sequence_number + 1,
|
||||
|
|
@ -515,7 +515,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
# Capture the response ID from the first event to ensure consistency
|
||||
if self._cached_response_id is None and hasattr(chunk, "response"):
|
||||
response_obj = getattr(chunk, "response", None)
|
||||
response_obj: ResponsesAPIResponse | None = getattr(chunk, "response", None)
|
||||
if response_obj and hasattr(response_obj, "id"):
|
||||
self._cached_response_id = response_obj.id
|
||||
verbose_logger.debug("Cached response ID: %s", self._cached_response_id)
|
||||
|
|
@ -559,7 +559,8 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
"""Check if this chunk indicates the response is completed"""
|
||||
from litellm.types.llms.openai import ResponsesAPIStreamEvents
|
||||
|
||||
return getattr(chunk, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
chunk_type: Final[object] = getattr(chunk, "type", None)
|
||||
return chunk_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
|
||||
async def _process_base_iterator_chunk(self) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
|
|
@ -571,14 +572,14 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
chunk: Final = await cast(Any, self.base_iterator).__anext__()
|
||||
|
||||
if self._cached_response_id is None and hasattr(chunk, "response"):
|
||||
new_response: Final = getattr(chunk, "response", None)
|
||||
new_response: Final[ResponsesAPIResponse | None] = getattr(chunk, "response", None)
|
||||
new_response_id: Final = getattr(new_response, "id", None) if new_response is not None else None
|
||||
if new_response_id:
|
||||
self._cached_response_id = new_response_id
|
||||
|
||||
# Ensure response ID consistency - update chunk if needed
|
||||
if self._cached_response_id and hasattr(chunk, "response"):
|
||||
response_obj = getattr(chunk, "response", None)
|
||||
response_obj: ResponsesAPIResponse | None = getattr(chunk, "response", None)
|
||||
if response_obj and hasattr(response_obj, "id"):
|
||||
if response_obj.id != self._cached_response_id:
|
||||
verbose_logger.debug(
|
||||
|
|
@ -605,7 +606,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
from litellm.responses.main import aresponses
|
||||
|
||||
# Make the initial response API call - but avoid the MCP wrapper
|
||||
params: Final = self.original_request_params.copy()
|
||||
params: Final[dict[str, object]] = self.original_request_params.copy()
|
||||
params["stream"] = True # Ensure streaming
|
||||
|
||||
# Use the pre-fetched all_tools from original_request_params (no re-processing needed)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import json
|
|||
import time
|
||||
import traceback
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
|
|
@ -1035,7 +1035,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
@runtime_checkable
|
||||
class _HasModelDump(Protocol):
|
||||
def model_dump(self, *, exclude_none: bool = ...) -> Mapping[str, object]: ...
|
||||
def model_dump(self, *, exclude_none: bool = ...) -> dict[str, object]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
|
|
@ -1043,8 +1043,8 @@ class _HasModelDumpJson(Protocol):
|
|||
def model_dump_json(self, *, exclude_none: bool = ...) -> str: ...
|
||||
|
||||
|
||||
def _dump_response_object(obj: Any) -> dict[str, Any]:
|
||||
if hasattr(obj, "model_dump"):
|
||||
def _dump_response_object(obj: object) -> dict[str, Any]:
|
||||
if isinstance(obj, _HasModelDump):
|
||||
return obj.model_dump()
|
||||
if _is_json_object(obj):
|
||||
return obj
|
||||
|
|
@ -1134,7 +1134,8 @@ def _add_text_like_part_events(
|
|||
delta=text[i : i + chunk_size],
|
||||
)
|
||||
)
|
||||
for annotation_index, annotation in enumerate(part_payload.get("annotations", []) or []):
|
||||
annotations_payload: Final[Sequence[dict[str, object]]] = part_payload.get("annotations", []) or []
|
||||
for annotation_index, annotation in enumerate(annotations_payload):
|
||||
events.append(
|
||||
openai_types.OutputTextAnnotationAddedEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED,
|
||||
|
|
@ -1200,7 +1201,8 @@ def _build_synthetic_response_events(
|
|||
]
|
||||
|
||||
sequence_number = 0
|
||||
for output_index, output_item in enumerate(getattr(transformed, "output", []) or []):
|
||||
output_items: Final[Sequence[object]] = getattr(transformed, "output", []) or []
|
||||
for output_index, output_item in enumerate(output_items):
|
||||
output_item_payload = _dump_response_object(output_item)
|
||||
item_id = str(output_item_payload.get("id") or transformed.id)
|
||||
item_type = output_item_payload.get("type")
|
||||
|
|
@ -1214,7 +1216,8 @@ def _build_synthetic_response_events(
|
|||
)
|
||||
|
||||
if item_type == "message":
|
||||
for content_index, part in enumerate(output_item_payload.get("content", []) or []):
|
||||
content_parts: Sequence[object] = output_item_payload.get("content", []) or []
|
||||
for content_index, part in enumerate(content_parts):
|
||||
part_payload = _dump_response_object(part)
|
||||
events.append(
|
||||
openai_types.ContentPartAddedEvent(
|
||||
|
|
@ -1261,7 +1264,8 @@ def _build_synthetic_response_events(
|
|||
)
|
||||
)
|
||||
elif item_type == "reasoning":
|
||||
for summary_index, summary in enumerate(output_item_payload.get("summary", []) or []):
|
||||
summaries: Sequence[object] = output_item_payload.get("summary", []) or []
|
||||
for summary_index, summary in enumerate(summaries):
|
||||
summary_payload = _dump_response_object(summary)
|
||||
summary_text = str(summary_payload.get("text") or "")
|
||||
for i in range(0, len(summary_text), chunk_size):
|
||||
|
|
@ -1463,7 +1467,8 @@ class ResponsesWebSocketStreaming:
|
|||
# masked response.completed.
|
||||
if self.output_guardrail_callbacks:
|
||||
try:
|
||||
_evt_type = json.loads(response_str).get("type")
|
||||
_evt_payload: Mapping[str, object] = json.loads(response_str)
|
||||
_evt_type = _evt_payload.get("type")
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
_evt_type = None
|
||||
if _evt_type in self._DELTA_EVENT_TYPES or _evt_type in self._OUTPUT_DONE_EVENT_TYPES:
|
||||
|
|
@ -1527,7 +1532,7 @@ class ResponsesWebSocketStreaming:
|
|||
Non-``response.create`` messages are returned unchanged.
|
||||
"""
|
||||
try:
|
||||
msg_obj: Final = json.loads(message)
|
||||
msg_obj: Final[dict[str, object]] = json.loads(message)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return message
|
||||
|
||||
|
|
@ -1544,7 +1549,8 @@ class ResponsesWebSocketStreaming:
|
|||
self.request_data["metadata"] = {}
|
||||
|
||||
modified = model_modified
|
||||
for cb in self.guardrail_callbacks:
|
||||
guardrail_cbs: Final[tuple[PresidioGuardrailCallback, ...]] = tuple(self.guardrail_callbacks)
|
||||
for cb in guardrail_cbs:
|
||||
presidio_config = cb.get_presidio_settings_from_request_data(self.request_data)
|
||||
# response.create carries client text in two shapes:
|
||||
# flat: {"type": "response.create", "input": ..., "instructions": ...}
|
||||
|
|
@ -1655,7 +1661,7 @@ class ResponsesWebSocketStreaming:
|
|||
return response_str
|
||||
|
||||
try:
|
||||
evt_obj: Final = json.loads(response_str)
|
||||
evt_obj: Final[dict[str, object]] = json.loads(response_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return response_str
|
||||
|
||||
|
|
@ -2012,7 +2018,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
async def _parse_message(self, raw_message: str) -> dict[str, object] | None:
|
||||
"""Parse raw WS text; return the message dict or None (JSON error / ignored type)."""
|
||||
try:
|
||||
msg_obj: Final = json.loads(raw_message)
|
||||
msg_obj: Final[dict[str, object]] = json.loads(raw_message)
|
||||
except json.JSONDecodeError:
|
||||
await self._send_error("Invalid JSON in response.create event", "invalid_request_error")
|
||||
return None
|
||||
|
|
@ -2293,11 +2299,10 @@ class ManagedResponsesWebSocketHandler:
|
|||
# reuse the router-resolved self.model; passing the alias raw to
|
||||
# litellm.aresponses fails in get_llm_provider. A genuinely different
|
||||
# provider-prefixed per-frame model is still honored.
|
||||
requested_model: Final = call_kwargs.pop("model", None)
|
||||
if requested_model is None or requested_model == self.model_group:
|
||||
model = self.model
|
||||
else:
|
||||
model = requested_model
|
||||
requested_model: Final[str | None] = call_kwargs.pop("model", None)
|
||||
model: Final[str] = (
|
||||
self.model if requested_model is None or requested_model == self.model_group else requested_model
|
||||
)
|
||||
|
||||
previous_response_id: Final[str | None] = call_kwargs.pop("previous_response_id", None)
|
||||
current_messages: Final = self._input_to_messages(call_kwargs.get("input"))
|
||||
|
|
|
|||
|
|
@ -93,9 +93,9 @@ class ResponsesAPIRequestUtils:
|
|||
|
||||
@staticmethod
|
||||
def merge_client_forwarded_headers(
|
||||
extra_headers: dict[str, Any] | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
client_headers: dict[str, str] | None,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Merge headers forwarded by the proxy (`headers` kwarg, set when
|
||||
`forward_client_headers_to_llm_api` is enabled) into `extra_headers`.
|
||||
|
|
@ -210,9 +210,9 @@ class ResponsesAPIRequestUtils:
|
|||
|
||||
valid_keys: Final = get_type_hints(ResponsesAPIOptionalRequestParams).keys()
|
||||
custom_llm_provider: Final = params.pop("custom_llm_provider", None)
|
||||
special_params: Final = params.pop("kwargs", {})
|
||||
special_params: Final[dict[str, object]] = params.pop("kwargs", {})
|
||||
|
||||
additional_drop_params: Final = params.pop("additional_drop_params", None)
|
||||
additional_drop_params: Final[list[str] | None] = params.pop("additional_drop_params", None)
|
||||
non_default_params: Final = PreProcessNonDefaultParams.base_pre_process_non_default_params(
|
||||
passed_params=params,
|
||||
special_params=special_params,
|
||||
|
|
@ -401,9 +401,9 @@ class ResponsesAPIRequestUtils:
|
|||
|
||||
@staticmethod
|
||||
def _update_encrypted_content_item_ids_in_response(
|
||||
response: Union["ResponsesAPIResponse", dict[str, Any]],
|
||||
response: Union["ResponsesAPIResponse", dict[str, object]],
|
||||
model_id: str | None,
|
||||
) -> Union["ResponsesAPIResponse", dict[str, Any]]:
|
||||
) -> Union["ResponsesAPIResponse", dict[str, object]]:
|
||||
"""Rewrite item IDs for output items that contain ``encrypted_content``.
|
||||
|
||||
Encodes ``model_id`` into the item ID so that follow-up requests can be
|
||||
|
|
@ -415,7 +415,7 @@ class ResponsesAPIRequestUtils:
|
|||
if not model_id:
|
||||
return response
|
||||
|
||||
output: list | None = None
|
||||
output: object = None
|
||||
if isinstance(response, dict):
|
||||
output = response.get("output")
|
||||
else:
|
||||
|
|
@ -459,7 +459,7 @@ class ResponsesAPIRequestUtils:
|
|||
return response
|
||||
|
||||
@staticmethod
|
||||
def _restore_encrypted_content_item_ids_in_input(request_input: Any) -> Any:
|
||||
def _restore_encrypted_content_item_ids_in_input(request_input: object) -> Any:
|
||||
"""Decode litellm-encoded item IDs in request input back to original IDs.
|
||||
|
||||
Called before forwarding the request to the upstream provider so the
|
||||
|
|
@ -867,7 +867,7 @@ class ResponsesAPIRequestUtils:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def collect_container_ids_from_responses_response(response: Any) -> list[str]:
|
||||
def collect_container_ids_from_responses_response(response: object) -> list[str]:
|
||||
"""Return unique container IDs referenced in a Responses API payload."""
|
||||
if response is None:
|
||||
return []
|
||||
|
|
@ -953,7 +953,7 @@ class ResponsesAPIRequestUtils:
|
|||
@staticmethod
|
||||
def extract_mcp_headers_from_request(
|
||||
secret_fields: dict[str, Any] | None,
|
||||
tools: Iterable[Any] | None,
|
||||
tools: Iterable[object] | None,
|
||||
) -> tuple[
|
||||
str | None,
|
||||
dict[str, dict[str, str]] | None,
|
||||
|
|
|
|||
|
|
@ -26,8 +26,9 @@ from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
|
|||
from pydantic import BaseModel, create_model
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY, RETURN_RAW_MODEL_NAME_METADATA_KEY
|
||||
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
|
||||
from litellm.llms.base_llm.base_utils import type_to_response_format_param
|
||||
from litellm.types.utils import (
|
||||
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
|
||||
|
|
@ -206,40 +207,6 @@ def _append_custom_keywords(base_keywords: list[str], custom_keywords: list[str]
|
|||
return [*base_keywords, *deduped_custom.values()]
|
||||
|
||||
|
||||
# Metadata keys that carry only the parent request's budget reservation state. These
|
||||
# must not reach internal sub-calls (classifier, embedding): the reservation belongs to
|
||||
# the routed completion being decided on, not to the sub-call itself, and forwarding it
|
||||
# would let the sub-call's cost callback finalize the reservation, causing the routed
|
||||
# completion's callback to skip incrementing key/team budget counters.
|
||||
#
|
||||
# Note: user_api_key_auth itself is intentionally kept; it is required by
|
||||
# _filter_deployments_by_model_access_groups to scope embedding/classifier model
|
||||
# selection to the caller's authorized access groups. It is forwarded as a sanitized
|
||||
# copy with its budget_reservation sub-field removed, because the proxy cost callback
|
||||
# (_get_budget_reservation_from_metadata) falls back to reading the reservation from
|
||||
# inside the auth object when the top-level key is absent; forwarding it unsanitized
|
||||
# would re-create the exact double-finalization this stripping exists to prevent.
|
||||
_BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
|
||||
|
||||
|
||||
def _sanitize_user_api_key_auth(auth: Any) -> Any:
|
||||
if isinstance(auth, dict):
|
||||
return {k: v for k, v in auth.items() if k != "budget_reservation"}
|
||||
if getattr(auth, "budget_reservation", None) is not None and hasattr(auth, "model_copy"):
|
||||
return auth.model_copy(update={"budget_reservation": None})
|
||||
return auth
|
||||
|
||||
|
||||
def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any]:
|
||||
if not metadata:
|
||||
return {}
|
||||
return {
|
||||
k: _sanitize_user_api_key_auth(v) if k == "user_api_key_auth" else v
|
||||
for k, v in metadata.items()
|
||||
if k not in _BUDGET_RESERVATION_METADATA_KEYS
|
||||
} | {INTERNAL_CALL_ORIGIN_METADATA_KEY: AUTOROUTER_CLASSIFIER_CALL_ORIGIN}
|
||||
|
||||
|
||||
def _parent_session_kwargs(request_kwargs: Mapping[str, Any] | None) -> Mapping[str, Any]:
|
||||
kwargs: Final = request_kwargs or {}
|
||||
return {k: kwargs[k] for k in ("litellm_session_id", "litellm_trace_id") if kwargs.get(k) is not None}
|
||||
|
|
@ -1069,7 +1036,7 @@ class ComplexityRouter(CustomLogger):
|
|||
)
|
||||
|
||||
request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata")
|
||||
metadata: Final = _classifier_call_metadata(request_metadata)
|
||||
metadata: Final = forwarded_internal_call_metadata(request_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN)
|
||||
turn_off_message_logging: Final = _effective_turn_off_message_logging(request_kwargs)
|
||||
|
||||
labeled_tiers: Final = self.config.labeled_tiers()
|
||||
|
|
@ -1562,8 +1529,12 @@ class ComplexityRouter(CustomLogger):
|
|||
# embedding call. Forwarding it would let the embedding's cost callback finalize the
|
||||
# reservation, so the routed completion's own callback then skips incrementing the
|
||||
# key/team budget. Key/team attribution fields are preserved for spend logging.
|
||||
metadata: Final = _classifier_call_metadata(request_kwargs.get("metadata"))
|
||||
litellm_metadata: Final = _classifier_call_metadata(request_kwargs.get("litellm_metadata"))
|
||||
metadata: Final = forwarded_internal_call_metadata(
|
||||
request_kwargs.get("metadata"), AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
)
|
||||
litellm_metadata: Final = forwarded_internal_call_metadata(
|
||||
request_kwargs.get("litellm_metadata"), AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
)
|
||||
turn_off_message_logging: Final = _effective_turn_off_message_logging(request_kwargs)
|
||||
proxy_server_request: Final = {"body": {"model": self.config.embedding_model, "input": [user_message]}}
|
||||
query_vector: Final = (
|
||||
|
|
|
|||
|
|
@ -8,9 +8,11 @@ Use this to route requests between Teams
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import CONSUMED_REQUEST_TAGS_METADATA_KEY
|
||||
|
|
@ -25,9 +27,39 @@ else:
|
|||
LitellmRouter = Any
|
||||
|
||||
|
||||
class _TagRoutingLitellmParams(TypedDict, total=False):
|
||||
tags: ReadOnly[Sequence[str] | None]
|
||||
tag_regex: ReadOnly[Sequence[str] | None]
|
||||
|
||||
|
||||
class _TagRoutingDeployment(TypedDict, total=False):
|
||||
model_name: ReadOnly[str]
|
||||
litellm_params: ReadOnly[_TagRoutingLitellmParams]
|
||||
model_info: ReadOnly[Mapping[str, object] | None]
|
||||
|
||||
|
||||
class _TagRoutingMatchStamp(TypedDict):
|
||||
matched_deployment: ReadOnly[str | None]
|
||||
matched_via: ReadOnly[str]
|
||||
matched_value: ReadOnly[str]
|
||||
request_tags: ReadOnly[Sequence[str]]
|
||||
user_agent: ReadOnly[str]
|
||||
|
||||
|
||||
class _TagRoutingMetadata(TypedDict, total=False):
|
||||
tags: ReadOnly[Sequence[str] | None]
|
||||
inherited_tags: ReadOnly[Sequence[str] | None]
|
||||
user_agent: ReadOnly[str]
|
||||
tag_routing: ReadOnly[_TagRoutingMatchStamp]
|
||||
_consumed_request_tags: ReadOnly[object]
|
||||
|
||||
|
||||
_EMPTY_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _is_valid_deployment_tag_regex(
|
||||
tag_regexes: list[str],
|
||||
header_strings: list[str],
|
||||
tag_regexes: Sequence[str],
|
||||
header_strings: Sequence[str],
|
||||
) -> str | None:
|
||||
"""
|
||||
Test compiled regex patterns against "Header-Name: value" strings.
|
||||
|
|
@ -77,11 +109,11 @@ def is_valid_deployment_tag(
|
|||
|
||||
|
||||
def _match_deployment(
|
||||
deployment: Any,
|
||||
request_tags: list[str] | None,
|
||||
header_strings: list[str],
|
||||
deployment: _TagRoutingDeployment,
|
||||
request_tags: Sequence[str] | None,
|
||||
header_strings: Sequence[str],
|
||||
match_any: bool,
|
||||
) -> dict[str, str] | None:
|
||||
) -> Mapping[str, str] | None:
|
||||
"""
|
||||
Determine whether *deployment* matches the current request.
|
||||
|
||||
|
|
@ -94,8 +126,8 @@ def _match_deployment(
|
|||
ran and failed, so the regex cannot override strict-tag policy.
|
||||
"""
|
||||
litellm_params: Final = deployment.get("litellm_params", {})
|
||||
deployment_tags: Final[list[str] | None] = litellm_params.get("tags")
|
||||
deployment_tag_regex: Final[list[str] | None] = litellm_params.get("tag_regex")
|
||||
deployment_tags: Final[Sequence[str] | None] = litellm_params.get("tags")
|
||||
deployment_tag_regex: Final[Sequence[str] | None] = litellm_params.get("tag_regex")
|
||||
|
||||
# 1. Exact tag match (existing behaviour).
|
||||
if deployment_tags and request_tags:
|
||||
|
|
@ -166,38 +198,38 @@ def _split_tags(tags: Sequence[str]) -> tuple[tuple[str, ...], list[str], tuple[
|
|||
|
||||
|
||||
def _exclude_deployments(
|
||||
deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
deployments: Iterable[_TagRoutingDeployment],
|
||||
excluded_set: frozenset[str],
|
||||
) -> list[Any]:
|
||||
) -> list[_TagRoutingDeployment]:
|
||||
if not excluded_set:
|
||||
return list(deployments)
|
||||
return [d for d in deployments if not excluded_set.intersection(d.get("litellm_params", {}).get("tags") or [])]
|
||||
|
||||
|
||||
def _require_all_tags(
|
||||
deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
deployments: Iterable[_TagRoutingDeployment],
|
||||
required_set: frozenset[str],
|
||||
) -> tuple[Any, ...]:
|
||||
) -> tuple[_TagRoutingDeployment, ...]:
|
||||
if not required_set:
|
||||
return tuple(deployments)
|
||||
return tuple(d for d in deployments if required_set.issubset(d.get("litellm_params", {}).get("tags") or []))
|
||||
|
||||
|
||||
def _default_tagged_pool(
|
||||
deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
) -> tuple[Any, ...]:
|
||||
deployments: Iterable[_TagRoutingDeployment],
|
||||
) -> tuple[_TagRoutingDeployment, ...]:
|
||||
defaults: Final = tuple(d for d in deployments if "default" in (d.get("litellm_params", {}).get("tags") or []))
|
||||
return defaults if defaults else tuple(deployments)
|
||||
|
||||
|
||||
def _known_tag_values(deployments: Sequence[Any] | Mapping[Any, Any]) -> frozenset[str]:
|
||||
def _known_tag_values(deployments: Iterable[_TagRoutingDeployment]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
tag for d in deployments for tag in (d.get("litellm_params", MappingProxyType({})).get("tags") or ())
|
||||
tag for d in deployments for tag in (d.get("litellm_params", _TagRoutingLitellmParams()).get("tags") or ())
|
||||
)
|
||||
|
||||
|
||||
def _unknown_required_tag_hides_an_answer(
|
||||
healthy_deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
healthy_deployments: Iterable[_TagRoutingDeployment],
|
||||
excluded_set: frozenset[str],
|
||||
required_set: frozenset[str],
|
||||
routing_confirmed: frozenset[str],
|
||||
|
|
@ -221,23 +253,23 @@ def _unknown_required_tag_hides_an_answer(
|
|||
|
||||
|
||||
def _chain_allows_fail_open(
|
||||
healthy_deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
healthy_deployments: Iterable[_TagRoutingDeployment],
|
||||
excluded_set: frozenset[str],
|
||||
required_set: frozenset[str],
|
||||
routing_confirmed: frozenset[str],
|
||||
) -> bool:
|
||||
if _unknown_required_tag_hides_an_answer(healthy_deployments, excluded_set, required_set, routing_confirmed):
|
||||
return False
|
||||
return any((d.get("model_info") or {}).get("allow_fail_open") is True for d in healthy_deployments)
|
||||
return any((d.get("model_info") or _EMPTY_MODEL_INFO).get("allow_fail_open") is True for d in healthy_deployments)
|
||||
|
||||
|
||||
def _trusted_only_pool(
|
||||
healthy_deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
healthy_deployments: Iterable[_TagRoutingDeployment],
|
||||
excluded_set: frozenset[str],
|
||||
required_set: frozenset[str],
|
||||
inherited_excluded_set: frozenset[str] | None,
|
||||
inherited_required_set: frozenset[str] | None,
|
||||
) -> tuple[Any, ...]:
|
||||
) -> tuple[_TagRoutingDeployment, ...]:
|
||||
# inherited_*_set is None only when this request carries no origin information
|
||||
# at all (e.g. direct SDK Router usage, bypassing the proxy layer that
|
||||
# populates metadata.inherited_tags) -- treat every constraint as
|
||||
|
|
@ -264,8 +296,8 @@ def _trusted_only_pool(
|
|||
|
||||
|
||||
def _resolve_or_fail_open(
|
||||
pool: Sequence[Any],
|
||||
healthy_deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
pool: Sequence[_TagRoutingDeployment],
|
||||
healthy_deployments: Iterable[_TagRoutingDeployment],
|
||||
excluded_set: frozenset[str],
|
||||
required_set: frozenset[str],
|
||||
inherited_excluded_set: frozenset[str] | None,
|
||||
|
|
@ -273,7 +305,7 @@ def _resolve_or_fail_open(
|
|||
routing_confirmed: frozenset[str],
|
||||
model: str,
|
||||
request_tags: object,
|
||||
) -> tuple[Any, ...]:
|
||||
) -> tuple[_TagRoutingDeployment, ...]:
|
||||
if pool:
|
||||
return tuple(pool)
|
||||
if _chain_allows_fail_open(healthy_deployments, excluded_set, required_set, routing_confirmed):
|
||||
|
|
@ -293,7 +325,7 @@ def _resolve_or_fail_open(
|
|||
|
||||
|
||||
def _resolve_constraint_only_pool(
|
||||
healthy_deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
healthy_deployments: Iterable[_TagRoutingDeployment],
|
||||
excluded_set: frozenset[str],
|
||||
required_set: frozenset[str],
|
||||
inherited_excluded_set: frozenset[str] | None,
|
||||
|
|
@ -301,7 +333,7 @@ def _resolve_constraint_only_pool(
|
|||
routing_confirmed: frozenset[str],
|
||||
model: str,
|
||||
request_tags: object,
|
||||
) -> tuple[Any, ...]:
|
||||
) -> tuple[_TagRoutingDeployment, ...]:
|
||||
pool: Final = (
|
||||
_require_all_tags(_exclude_deployments(healthy_deployments, excluded_set), required_set)
|
||||
if required_set
|
||||
|
|
@ -323,8 +355,8 @@ def _resolve_constraint_only_pool(
|
|||
def _all_deployments_or_fallback(
|
||||
llm_router_instance: LitellmRouter,
|
||||
model: str,
|
||||
fallback: Sequence[Any] | Mapping[Any, Any],
|
||||
) -> Sequence[Any] | Mapping[Any, Any]:
|
||||
fallback: Iterable[_TagRoutingDeployment],
|
||||
) -> Iterable[_TagRoutingDeployment]:
|
||||
try:
|
||||
return llm_router_instance._get_all_deployments(model_name=model)
|
||||
except Exception: # noqa: BLE001 # fail safe toward today's healthy-only behavior on lookup errors
|
||||
|
|
@ -334,8 +366,8 @@ def _all_deployments_or_fallback(
|
|||
def _chain_tag_filtering_override(
|
||||
llm_router_instance: LitellmRouter,
|
||||
model: str,
|
||||
healthy_deployments: Sequence[Any] | Mapping[Any, Any],
|
||||
) -> bool | None:
|
||||
healthy_deployments: Iterable[_TagRoutingDeployment],
|
||||
) -> object:
|
||||
# Resolved from every deployment configured for this model group, not just the
|
||||
# ones that survived cooldown/health filtering (async_get_healthy_deployments
|
||||
# filters cooldowns before calling get_deployments_for_tag) -- otherwise the
|
||||
|
|
@ -347,14 +379,14 @@ def _chain_tag_filtering_override(
|
|||
# than crashing the request.
|
||||
all_deployments: Final = _all_deployments_or_fallback(llm_router_instance, model, healthy_deployments)
|
||||
for d in all_deployments:
|
||||
value = (d.get("model_info") or MappingProxyType({})).get("enable_tag_filtering")
|
||||
value = (d.get("model_info") or _EMPTY_MODEL_INFO).get("enable_tag_filtering")
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _inherited_constraint_sets(
|
||||
inherited_tags: object, routing_prefix: str
|
||||
inherited_tags: Sequence[str] | None, routing_prefix: str
|
||||
) -> tuple[frozenset[str] | None, frozenset[str] | None]:
|
||||
# None means no origin information is available at all (e.g. this request
|
||||
# bypassed the proxy layer that populates metadata.inherited_tags, as direct
|
||||
|
|
@ -385,15 +417,18 @@ def _tag_known_to_group(
|
|||
if tag_set & routing_confirmed:
|
||||
return True
|
||||
try:
|
||||
all_deployments: Final = llm_router_instance._get_all_deployments(model_name=model)
|
||||
all_deployments: Final[Sequence[_TagRoutingDeployment]] = llm_router_instance._get_all_deployments(
|
||||
model_name=model
|
||||
)
|
||||
except Exception: # noqa: BLE001 # fail safe toward "unrecognized" so lookup errors preserve the existing silent-fallback behavior
|
||||
return False
|
||||
return any(
|
||||
tag_set.intersection(d.get("litellm_params", MappingProxyType({})).get("tags") or ()) for d in all_deployments
|
||||
tag_set.intersection(d.get("litellm_params", _TagRoutingLitellmParams()).get("tags") or ())
|
||||
for d in all_deployments
|
||||
)
|
||||
|
||||
|
||||
def _request_tags_after_router_consumption(metadata: Mapping[Any, Any], model: str) -> Sequence[str] | None:
|
||||
def _request_tags_after_router_consumption(metadata: _TagRoutingMetadata, model: str) -> Sequence[str] | None:
|
||||
# The pre-routing hook stamps which tags selected the router it rewrote the request
|
||||
# to: those tags already did their job and must not also constrain deployment choice
|
||||
# inside the routed group. The request's other tags still apply there, on top of the
|
||||
|
|
@ -451,7 +486,8 @@ async def get_deployments_for_tag(
|
|||
|
||||
verbose_logger.debug("request metadata: %s", request_kwargs.get(metadata_variable_name))
|
||||
if metadata_variable_name in request_kwargs:
|
||||
metadata: Final = request_kwargs[metadata_variable_name]
|
||||
metadata: Final[_TagRoutingMetadata] = request_kwargs[metadata_variable_name]
|
||||
stampable_metadata: Final[dict[str, object]] = request_kwargs[metadata_variable_name]
|
||||
request_tags: Final = _request_tags_after_router_consumption(metadata, model)
|
||||
match_any: Final = llm_router_instance.tag_filtering_match_any
|
||||
routing_prefix: Final = llm_router_instance.tag_routing_prefix or ""
|
||||
|
|
@ -496,8 +532,8 @@ async def get_deployments_for_tag(
|
|||
request_tags,
|
||||
)
|
||||
|
||||
new_healthy_deployments: Final[list[Any]] = []
|
||||
default_deployments: Final[list[Any]] = []
|
||||
new_healthy_deployments: Final[list[_TagRoutingDeployment]] = []
|
||||
default_deployments: Final[list[_TagRoutingDeployment]] = []
|
||||
|
||||
if has_positive_filter:
|
||||
verbose_logger.debug(
|
||||
|
|
@ -523,7 +559,7 @@ async def get_deployments_for_tag(
|
|||
match_result["matched_value"],
|
||||
)
|
||||
if "tag_routing" not in metadata:
|
||||
metadata["tag_routing"] = {
|
||||
stampable_metadata["tag_routing"] = {
|
||||
"matched_deployment": deployment.get("model_name"),
|
||||
"matched_via": match_result["matched_via"],
|
||||
"matched_value": match_result["matched_value"],
|
||||
|
|
@ -568,7 +604,7 @@ async def get_deployments_for_tag(
|
|||
return new_healthy_deployments if len(new_healthy_deployments) > 0 else default_deployments
|
||||
|
||||
# for Untagged requests use default deployments if set
|
||||
_default_deployments_with_tags: Final = []
|
||||
_default_deployments_with_tags: Final[list[_TagRoutingDeployment]] = []
|
||||
for deployment in healthy_deployments:
|
||||
if "default" in deployment.get("litellm_params", {}).get("tags", []):
|
||||
_default_deployments_with_tags.append(deployment)
|
||||
|
|
@ -603,7 +639,7 @@ def _tags_in_metadata(metadata: object) -> list[str]:
|
|||
|
||||
|
||||
def _get_tags_from_request_kwargs(
|
||||
request_kwargs: Mapping[Any, Any] | None = None,
|
||||
request_kwargs: Mapping[str, object] | None = None,
|
||||
metadata_variable_name: Literal["metadata", "litellm_metadata"] | None = None,
|
||||
) -> list[str]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -3,9 +3,10 @@ Types for auto-router management endpoints
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, computed_field, field_validator
|
||||
|
||||
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
|
||||
from litellm.types.utils import StandardLoggingRoutingDecision
|
||||
|
|
@ -141,3 +142,112 @@ class AutoRouterBenchmarksResponse(BaseModel):
|
|||
routers_in_scope: int
|
||||
totals: AutoRouterBenchmarkTotals
|
||||
groups: tuple[AutoRouterBenchmarkGroup, ...]
|
||||
|
||||
|
||||
ShadowEvalStatus: TypeAlias = Literal["running", "completed", "stopped"]
|
||||
|
||||
DEFAULT_SHADOW_EVAL_JUDGE_MODEL: Final[str] = "anthropic/claude-sonnet-5"
|
||||
|
||||
|
||||
class StartShadowEvalRequest(BaseModel):
|
||||
"""Start shadowing a key's traffic through an auto-router for blind comparison."""
|
||||
|
||||
api_key_id: str = Field(
|
||||
description=(
|
||||
"The hashed virtual key whose traffic will be shadowed. Shadow evaluation runs ONLY on this "
|
||||
"key's traffic; requests made with any other key are not sampled."
|
||||
)
|
||||
)
|
||||
router_name: str = Field(description="The auto-router config to shadow requests through")
|
||||
shadow_percentage: float = Field(
|
||||
ge=0.1,
|
||||
le=100.0,
|
||||
description="Percentage of the key's requests to duplicate through the router",
|
||||
)
|
||||
judge_model: str = Field(
|
||||
default=DEFAULT_SHADOW_EVAL_JUDGE_MODEL,
|
||||
description=(
|
||||
"Model used to blindly judge real vs. shadow responses. The judge only compares two answers, so a "
|
||||
"mid-tier model (Claude Sonnet or GPT-4o class) is the sweet spot: small/nano-class models produce "
|
||||
"unreliable or malformed verdicts, while frontier reasoning models add cost without changing outcomes."
|
||||
),
|
||||
)
|
||||
duration_days: int = Field(
|
||||
default=7,
|
||||
ge=1,
|
||||
le=30,
|
||||
description="How many days the job samples traffic before completing on its own",
|
||||
)
|
||||
max_turns: int = Field(
|
||||
default=200,
|
||||
ge=1,
|
||||
le=2000,
|
||||
description=(
|
||||
"Sample budget: the job judges at most this many turns, then completes. This is also the spend "
|
||||
"bound; expected judge cost is roughly max_turns times one judge call"
|
||||
),
|
||||
)
|
||||
|
||||
@field_validator("shadow_percentage")
|
||||
@classmethod
|
||||
def _round_percentage(cls, value: float) -> float:
|
||||
return round(value, 2)
|
||||
|
||||
|
||||
class ShadowEvalSlice(BaseModel):
|
||||
"""Judge outcomes for one slice of a job's verdicts (a router tier, or one of the
|
||||
models the shadowed key currently uses)."""
|
||||
|
||||
group: str
|
||||
turn_count: int
|
||||
real_win_rate_pct: float = Field(description="Share of judged turns where the real (control) model won")
|
||||
shadow_win_rate_pct: float = Field(description="Share of judged turns where the shadowed router's pick won")
|
||||
tie_rate_pct: float
|
||||
avg_judge_confidence: float
|
||||
|
||||
|
||||
class ShadowEvalResult(BaseModel):
|
||||
"""Stratified results of a shadow-eval job's verdicts so far."""
|
||||
|
||||
by_tier: tuple[ShadowEvalSlice, ...]
|
||||
by_current_model: tuple[ShadowEvalSlice, ...]
|
||||
overall_shadow_win_rate_pct: float
|
||||
overall_tie_rate_pct: float
|
||||
|
||||
|
||||
class ShadowEvalJobResponse(BaseModel):
|
||||
"""A shadow-eval job. Validates directly from the prisma record (job_id reads the
|
||||
row's id); status is derived from stopped_at and ends_at, never stored, so no writer
|
||||
anywhere can produce an inconsistent one. Aggregate fields are populated by the
|
||||
detail endpoint only and stay None on list responses."""
|
||||
|
||||
model_config = ConfigDict(from_attributes=True, populate_by_name=True)
|
||||
|
||||
job_id: str = Field(validation_alias=AliasChoices("id", "job_id"))
|
||||
api_key_id: str = Field(description="The hashed virtual key whose traffic this job evaluates, and only that key's")
|
||||
router_name: str
|
||||
judge_model: str
|
||||
shadow_percentage: float
|
||||
max_turns: int
|
||||
created_at: datetime
|
||||
ends_at: datetime
|
||||
stopped_at: datetime | None = None
|
||||
|
||||
judged_count: int | None = Field(default=None, description="Verdicts recorded; detail endpoint only")
|
||||
error_count: int | None = Field(default=None, description="Sampled attempts that errored; detail endpoint only")
|
||||
judge_spend: float | None = Field(default=None, description="Judge cost so far; detail endpoint only")
|
||||
last_error: str | None = Field(default=None, description="Most recent attempt error; detail endpoint only")
|
||||
results: ShadowEvalResult | None = Field(default=None, description="Stratified verdicts; detail endpoint only")
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def status(self) -> ShadowEvalStatus:
|
||||
"""A job whose window has passed reads completed even if a later sweep stamped
|
||||
stopped_at; stopped means sampling ended before the window did."""
|
||||
if datetime.now(timezone.utc) >= (
|
||||
self.ends_at if self.ends_at.tzinfo else self.ends_at.replace(tzinfo=timezone.utc)
|
||||
):
|
||||
return "completed"
|
||||
if self.stopped_at is not None:
|
||||
return "stopped"
|
||||
return "running"
|
||||
|
|
|
|||
|
|
@ -2782,11 +2782,13 @@ RoutingDecisionCause = Literal[
|
|||
]
|
||||
|
||||
|
||||
InternalCallOrigin = Literal["autorouter_classifier"]
|
||||
InternalCallOrigin = Literal["autorouter_classifier", "shadow_eval_router", "shadow_eval_judge"]
|
||||
"""Which internal litellm feature originated a billed sub-call, so a spend log row
|
||||
records that it is not traffic the caller sent."""
|
||||
|
||||
AUTOROUTER_CLASSIFIER_CALL_ORIGIN: Final[InternalCallOrigin] = "autorouter_classifier"
|
||||
SHADOW_EVAL_ROUTER_CALL_ORIGIN: Final[InternalCallOrigin] = "shadow_eval_router"
|
||||
SHADOW_EVAL_JUDGE_CALL_ORIGIN: Final[InternalCallOrigin] = "shadow_eval_judge"
|
||||
|
||||
|
||||
class StandardLoggingRoutingDecision(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3058
|
||||
"limit": 3046
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 133
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 1384
|
||||
"limit": 1342
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 11
|
||||
|
|
@ -39,7 +39,7 @@
|
|||
"limit": 505
|
||||
},
|
||||
"B009": {
|
||||
"limit": 64
|
||||
"limit": 60
|
||||
},
|
||||
"B010": {
|
||||
"limit": 190
|
||||
|
|
@ -234,7 +234,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 1224
|
||||
"limit": 1220
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 524
|
||||
|
|
|
|||
|
|
@ -1450,6 +1450,44 @@ model LiteLLM_AutoRouterSession {
|
|||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic.
|
||||
// A sampled slice of requests is duplicated through the router in a detached task and an
|
||||
// LLM judge compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// stopped_at; every count, status, and spend figure is derived from the append-only
|
||||
// attempt rows, so nothing can disagree across pods or stop races.
|
||||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
api_key_id String // hashed virtual key whose traffic is shadowed
|
||||
router_name String
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample budget: judge at most this many turns
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
ends_at DateTime
|
||||
stopped_at DateTime?
|
||||
|
||||
@@index([api_key_id])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
// One row per sampled pipeline: a blind verdict (real | shadow | tie) or an error.
|
||||
model LiteLLM_ShadowEvalAttempt {
|
||||
id String @id @default(cuid())
|
||||
job_id String
|
||||
request_id String // the judged real request
|
||||
outcome String // real | shadow | tie | error
|
||||
tier String? // router's tier for the prompt, when classified
|
||||
real_model String?
|
||||
shadow_model String?
|
||||
confidence Float?
|
||||
judge_cost Float @default(0)
|
||||
error String?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
@@index([job_id])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
|
|
|
|||
466
tests/test_litellm/integrations/test_shadow_eval_logger.py
Normal file
466
tests/test_litellm/integrations/test_shadow_eval_logger.py
Normal file
|
|
@ -0,0 +1,466 @@
|
|||
"""Unit tests for the shadow-eval logger: sampling, unmasking, the hook's skip chain,
|
||||
the detached pipeline's single attempt-row write, and the cache-first job lookup."""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.integrations.shadow_eval_logger import (
|
||||
_MAX_CONCURRENT_SHADOW_TASKS,
|
||||
_MAX_JUDGE_PROMPT_CHARS,
|
||||
JUDGE_MAX_OUTPUT_TOKENS,
|
||||
ActiveShadowEvalJob,
|
||||
ShadowEvalLogger,
|
||||
_judge_user_prompt,
|
||||
_sample_hits,
|
||||
_unmask_preference,
|
||||
)
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
|
||||
|
||||
def _job(**overrides) -> ActiveShadowEvalJob:
|
||||
defaults = dict(
|
||||
id="job-1",
|
||||
router_name="my-router",
|
||||
shadow_percentage=100.0,
|
||||
judge_model="judge-model",
|
||||
max_turns=200,
|
||||
ends_at=datetime.now(timezone.utc) + timedelta(days=1),
|
||||
attempts=0,
|
||||
)
|
||||
return ActiveShadowEvalJob(**{**defaults, **overrides})
|
||||
|
||||
|
||||
def _prisma(jobs=(), attempt_counts=()) -> MagicMock:
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=list(jobs))
|
||||
prisma.db.litellm_shadowevalattempt.group_by = AsyncMock(
|
||||
return_value=[{"job_id": job_id, "_count": {"_all": count}} for job_id, count in attempt_counts]
|
||||
)
|
||||
prisma.db.litellm_shadowevalattempt.create = AsyncMock()
|
||||
return prisma
|
||||
|
||||
|
||||
def _job_record(job: ActiveShadowEvalJob, api_key_id="key-hash") -> MagicMock:
|
||||
record = MagicMock()
|
||||
for field, value in dict(
|
||||
id=job.id,
|
||||
api_key_id=api_key_id,
|
||||
router_name=job.router_name,
|
||||
shadow_percentage=job.shadow_percentage,
|
||||
judge_model=job.judge_model,
|
||||
max_turns=job.max_turns,
|
||||
ends_at=job.ends_at,
|
||||
).items():
|
||||
setattr(record, field, value)
|
||||
return record
|
||||
|
||||
|
||||
def _router(shadow_text="shadow answer", judge_json='{"preference": "A", "confidence": 0.9, "reasoning": "x"}'):
|
||||
"""One mock router serving the shadow call first, the judge call second. The shadow
|
||||
call's metadata receives the routing decision write-back, like the real router."""
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
if kwargs["model"] == "my-router":
|
||||
kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": "cheap-model"}
|
||||
return {"choices": [{"message": {"content": shadow_text}}], "usage": {"completion_tokens": 5}}
|
||||
return {"choices": [{"message": {"content": judge_json}}]}
|
||||
|
||||
router.acompletion = MagicMock(side_effect=acompletion)
|
||||
return router
|
||||
|
||||
|
||||
def _logger(router=None, prisma=None, job=None) -> ShadowEvalLogger:
|
||||
cache = InMemoryCache(max_size_in_memory=4, default_ttl=60)
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: router,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=cache,
|
||||
)
|
||||
if job is not None:
|
||||
cache.set_cache("shadow_eval:active_jobs", {"key-hash": job})
|
||||
return logger
|
||||
|
||||
|
||||
def _success_kwargs(request_id="req-1", api_key_hash="key-hash", request_metadata=None, call_type="acompletion"):
|
||||
return {
|
||||
"standard_logging_object": {
|
||||
"id": request_id,
|
||||
"call_type": call_type,
|
||||
"model": "claude-opus",
|
||||
"metadata": {"user_api_key_hash": api_key_hash},
|
||||
"model_parameters": {"temperature": 0.5, "stream": True},
|
||||
},
|
||||
"litellm_params": {"metadata": request_metadata or {}},
|
||||
"messages": [{"role": "user", "content": "what is 2+2"}],
|
||||
}
|
||||
|
||||
|
||||
RESPONSE = {"choices": [{"message": {"content": "real answer"}}]}
|
||||
|
||||
|
||||
async def _drain(logger: ShadowEvalLogger, target: int = 0):
|
||||
for _ in range(100):
|
||||
if logger._inflight_shadow_tasks == target:
|
||||
return
|
||||
await asyncio.sleep(0.01)
|
||||
raise AssertionError("shadow tasks never drained")
|
||||
|
||||
|
||||
class TestSampling:
|
||||
def test_boundaries_and_determinism(self):
|
||||
assert not any(_sample_hits(f"req-{i}", "job", 0.0) for i in range(100))
|
||||
assert all(_sample_hits(f"req-{i}", "job", 100.0) for i in range(100))
|
||||
assert len({_sample_hits("req-1", "job-1", 50.0) for _ in range(10)}) == 1
|
||||
|
||||
def test_distribution_close_to_percentage(self):
|
||||
hits = sum(_sample_hits(f"req-{i}", "job-x", 10.0) for i in range(10_000))
|
||||
assert 800 < hits < 1200
|
||||
|
||||
def test_different_jobs_sample_independently(self):
|
||||
agreements = sum(
|
||||
_sample_hits(f"req-{i}", "job-a", 50.0) == _sample_hits(f"req-{i}", "job-b", 50.0) for i in range(1000)
|
||||
)
|
||||
assert 300 < agreements < 700
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw,real_is_a,expected",
|
||||
[
|
||||
("A", True, "real"),
|
||||
("a", True, "real"),
|
||||
("A", False, "shadow"),
|
||||
("B", True, "shadow"),
|
||||
("B", False, "real"),
|
||||
("tie", True, "tie"),
|
||||
("garbage", True, "tie"),
|
||||
("", False, "tie"),
|
||||
],
|
||||
)
|
||||
def test_unmask_preference(raw, real_is_a, expected):
|
||||
assert _unmask_preference(raw, real_is_a) == expected
|
||||
|
||||
|
||||
def test_judge_prompt_is_bounded_however_large_the_inputs():
|
||||
prompt = _judge_user_prompt("c" * 200_000, "a" * 200_000, "b" * 200_000)
|
||||
assert len(prompt) < _MAX_JUDGE_PROMPT_CHARS + 100
|
||||
assert prompt.endswith("Which response is better?")
|
||||
small = _judge_user_prompt("conv", "alpha", "beta")
|
||||
assert "conv" in small and "alpha" in small and "beta" in small
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestSuccessHookSkipChain:
|
||||
async def test_happy_path_writes_exactly_one_attempt_row(self, monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm as litellm_module
|
||||
|
||||
monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005)
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, job=_job())
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
create = prisma.db.litellm_shadowevalattempt.create
|
||||
create.assert_awaited_once()
|
||||
row = create.call_args.kwargs["data"]
|
||||
assert row["job_id"] == "job-1"
|
||||
assert row["request_id"] == "req-1"
|
||||
assert row["outcome"] in ("real", "shadow")
|
||||
assert row["tier"] == "SIMPLE"
|
||||
assert row["real_model"] == "claude-opus"
|
||||
assert row["shadow_model"] == "cheap-model"
|
||||
assert row["confidence"] == 0.9
|
||||
assert row["judge_cost"] == 0.005
|
||||
assert row["error"] is None
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 0
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs_mutation,job_mutation",
|
||||
[
|
||||
({"request_metadata": {INTERNAL_CALL_ORIGIN_METADATA_KEY: "shadow_eval_router"}}, {}),
|
||||
({"api_key_hash": "other-key"}, {}),
|
||||
({"call_type": "aembedding"}, {}),
|
||||
({"call_type": None}, {}),
|
||||
({"request_metadata": {"routing_decision": {"router_model_name": "my-router"}}}, {}),
|
||||
({}, {"ends_at": datetime.now(timezone.utc) - timedelta(seconds=1)}),
|
||||
({}, {"attempts": 200}),
|
||||
({}, {"attempts": 199, "max_turns": 200, "_starts": 1}),
|
||||
],
|
||||
ids=[
|
||||
"internal-origin",
|
||||
"no-job-for-key",
|
||||
"non-chat",
|
||||
"missing-call-type",
|
||||
"self-shadow",
|
||||
"past-end",
|
||||
"turn-budget-reached",
|
||||
"budget-consumed-by-started-tasks",
|
||||
],
|
||||
)
|
||||
async def test_skip_paths_store_nothing(self, kwargs_mutation, job_mutation):
|
||||
starts = job_mutation.pop("_starts", 0)
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, job=_job(**job_mutation))
|
||||
logger._job_starts = {"job-1": starts}
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(**kwargs_mutation), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
|
||||
assert logger._job_starts.get("job-1", 0) == starts
|
||||
|
||||
async def test_completed_pipelines_hold_turn_budget_within_a_cache_generation(self):
|
||||
"""A finished pipeline frees its concurrency slot but not its slice of the turn
|
||||
budget; the budget only reopens when a cache refill absorbs the written rows."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, job=_job(attempts=199, max_turns=200))
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
await logger.async_log_success_event(_success_kwargs(request_id="req-2"), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
assert prisma.db.litellm_shadowevalattempt.create.await_count == 1
|
||||
|
||||
async def test_v1_messages_surface_forwards_identity_from_litellm_metadata(self):
|
||||
"""/v1/messages stores identity in litellm_params.litellm_metadata, so the hook
|
||||
resolves the bucket through the shared helper; every surface forwards the same
|
||||
identity to the shadow and judge calls."""
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, job=_job())
|
||||
|
||||
hook_kwargs = _success_kwargs()
|
||||
hook_kwargs["litellm_params"] = {
|
||||
"litellm_metadata": {"user_api_key_hash": "key-hash", "user_api_key_team_id": "team-1"}
|
||||
}
|
||||
await logger.async_log_success_event(hook_kwargs, RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
shadow_call = router.acompletion.call_args_list[0].kwargs
|
||||
assert shadow_call["metadata"]["user_api_key_hash"] == "key-hash"
|
||||
assert shadow_call["metadata"]["user_api_key_team_id"] == "team-1"
|
||||
|
||||
async def test_redacted_requests_are_never_shadowed(self):
|
||||
"""Redaction rewrites the logged messages before callbacks run, so this hook only
|
||||
ever sees placeholders for opted-out traffic; the skip uses the redactor's own
|
||||
predicate, so every redaction source counts."""
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, job=_job())
|
||||
|
||||
hook_kwargs = _success_kwargs()
|
||||
hook_kwargs["standard_callback_dynamic_params"] = {"turn_off_message_logging": True}
|
||||
await logger.async_log_success_event(hook_kwargs, RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
||||
router.acompletion.assert_not_called()
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
|
||||
|
||||
async def test_inflight_cap_sheds_instead_of_queueing(self):
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, job=_job())
|
||||
logger._inflight_shadow_tasks = _MAX_CONCURRENT_SHADOW_TASKS
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None)
|
||||
|
||||
assert logger._inflight_shadow_tasks == _MAX_CONCURRENT_SHADOW_TASKS
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestActiveJobsCache:
|
||||
async def test_cache_miss_reads_db_once_then_serves_from_cache(self):
|
||||
job = _job()
|
||||
prisma = _prisma(jobs=[_job_record(job)], attempt_counts=[("job-1", 7)])
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: None,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60),
|
||||
)
|
||||
|
||||
first = await logger._active_jobs()
|
||||
second = await logger._active_jobs()
|
||||
|
||||
assert first["key-hash"].id == "job-1"
|
||||
assert second["key-hash"].attempts == 7
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
|
||||
where = prisma.db.litellm_shadowevaljob.find_many.call_args.kwargs["where"]
|
||||
assert where["stopped_at"] is None
|
||||
assert "gt" in where["ends_at"]
|
||||
count_where = prisma.db.litellm_shadowevalattempt.group_by.call_args.kwargs["where"]
|
||||
assert count_where == {"job_id": {"in": ["job-1"]}}
|
||||
|
||||
async def test_no_active_jobs_is_cached_too(self):
|
||||
prisma = _prisma(jobs=[])
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: None,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60),
|
||||
)
|
||||
|
||||
assert await logger._active_jobs() == {}
|
||||
assert await logger._active_jobs() == {}
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
|
||||
prisma.db.litellm_shadowevalattempt.group_by.assert_not_called()
|
||||
|
||||
async def test_db_fault_returns_empty_without_caching_the_fault(self):
|
||||
prisma = _prisma()
|
||||
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(side_effect=RuntimeError("db blip"))
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: None,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60),
|
||||
)
|
||||
|
||||
assert await logger._active_jobs() == {}
|
||||
assert await logger._active_jobs() == {}
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 2
|
||||
|
||||
async def test_cache_refill_resets_the_starts_counter(self):
|
||||
job = _job()
|
||||
prisma = _prisma(jobs=[_job_record(job)], attempt_counts=[("job-1", 7)])
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: None,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60),
|
||||
)
|
||||
logger._job_starts = {"job-1": 5}
|
||||
|
||||
await logger._active_jobs()
|
||||
|
||||
assert logger._job_starts == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestShadowPipeline:
|
||||
async def test_no_prisma_means_no_provider_spend(self):
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=None)
|
||||
|
||||
await logger._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
model_parameters={},
|
||||
parent_metadata={},
|
||||
)
|
||||
|
||||
router.acompletion.assert_not_called()
|
||||
|
||||
async def test_over_budget_key_skips_before_any_call(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""The gate delegates to the auth path's own budget owner, so an over-budget
|
||||
verdict there (BudgetExceededError) skips the shadow before any provider call."""
|
||||
import litellm.proxy.auth.auth_checks as auth_checks
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
monkeypatch.setattr(
|
||||
auth_checks,
|
||||
"_virtual_key_max_budget_check",
|
||||
AsyncMock(side_effect=BudgetExceededError(current_cost=11.0, max_budget=10.0)),
|
||||
)
|
||||
router = _router()
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=router, prisma=prisma)
|
||||
|
||||
await logger._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
model_parameters={},
|
||||
parent_metadata={"user_api_key_auth": UserAPIKeyAuth(api_key="sk-abc", max_budget=10.0)},
|
||||
)
|
||||
|
||||
router.acompletion.assert_not_called()
|
||||
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"router_factory,expected_error,expected_cost",
|
||||
[
|
||||
(lambda: _failing_router(), "provider exploded", 0.0),
|
||||
(lambda: _router(judge_json="I prefer response A, definitely"), "unparseable judge verdict", 0.007),
|
||||
],
|
||||
ids=["shadow-call-fails", "judge-verdict-unparseable"],
|
||||
)
|
||||
async def test_failures_become_error_rows_and_keep_billed_judge_cost(
|
||||
self, router_factory, expected_error, expected_cost, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
import litellm as litellm_module
|
||||
|
||||
monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.007)
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=router_factory(), prisma=prisma)
|
||||
|
||||
await logger._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
model_parameters={},
|
||||
parent_metadata={},
|
||||
)
|
||||
|
||||
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
|
||||
assert row["outcome"] == "error"
|
||||
assert expected_error in row["error"]
|
||||
assert row["confidence"] is None
|
||||
assert row["judge_cost"] == expected_cost
|
||||
|
||||
async def test_sub_calls_carry_identity_and_origin_but_never_parent_request_state(self):
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma)
|
||||
parent_metadata = {
|
||||
"user_api_key_hash": "key-hash",
|
||||
"user_api_key_team_id": "team-1",
|
||||
"user_api_key_budget_reservation": {"amount": 1.0},
|
||||
"routing_decision": {"router_model_name": "other-router"},
|
||||
}
|
||||
|
||||
await logger._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
model_parameters={"stream": True, "temperature": 0.2, "metadata": {"x": 1}},
|
||||
parent_metadata=parent_metadata,
|
||||
)
|
||||
|
||||
shadow_call = router.acompletion.call_args_list[0].kwargs
|
||||
judge_call = router.acompletion.call_args_list[1].kwargs
|
||||
for call in (shadow_call, judge_call):
|
||||
assert call["num_retries"] == 0
|
||||
assert call["fallbacks"] == []
|
||||
assert call["metadata"]["user_api_key_hash"] == "key-hash"
|
||||
assert call["metadata"]["user_api_key_team_id"] == "team-1"
|
||||
assert "user_api_key_budget_reservation" not in call["metadata"]
|
||||
assert shadow_call["metadata"][INTERNAL_CALL_ORIGIN_METADATA_KEY] == SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
assert judge_call["metadata"][INTERNAL_CALL_ORIGIN_METADATA_KEY] == SHADOW_EVAL_JUDGE_CALL_ORIGIN
|
||||
assert "routing_decision" not in judge_call["metadata"]
|
||||
assert "stream" not in shadow_call
|
||||
assert shadow_call["temperature"] == 0.2
|
||||
assert judge_call["max_tokens"] == JUDGE_MAX_OUTPUT_TOKENS
|
||||
|
||||
|
||||
def _failing_router():
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=None)
|
||||
router.acompletion = AsyncMock(side_effect=RuntimeError("provider exploded"))
|
||||
return router
|
||||
|
|
@ -0,0 +1,121 @@
|
|||
"""Unit tests for internal-call metadata forwarding: budget-reservation stripping and origin stamping."""
|
||||
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.litellm_core_utils.internal_call_metadata import (
|
||||
forwarded_internal_call_metadata,
|
||||
sanitized_forwardable_call_metadata,
|
||||
)
|
||||
from litellm.types.utils import SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
|
||||
PARENT = {
|
||||
"user_api_key": "sk-hash",
|
||||
"user_api_key_hash": "sk-hash",
|
||||
"user_api_key_team_id": "team-1",
|
||||
"user_api_key_budget_reservation": {"amount": 1.0},
|
||||
"user_api_key_auth": {"api_key": "sk-hash", "budget_reservation": {"amount": 1.0}},
|
||||
"routing_decision": {"router_model_name": "my-router"},
|
||||
"headers": {"x-request-id": "abc"},
|
||||
}
|
||||
|
||||
|
||||
def test_forwarded_metadata_strips_reservation_everywhere_and_stamps_origin():
|
||||
result = forwarded_internal_call_metadata(PARENT, "autorouter_classifier")
|
||||
|
||||
assert result[INTERNAL_CALL_ORIGIN_METADATA_KEY] == "autorouter_classifier"
|
||||
assert "user_api_key_budget_reservation" not in result
|
||||
assert result["user_api_key_auth"] == {"api_key": "sk-hash"}
|
||||
assert result["routing_decision"] == {"router_model_name": "my-router"}
|
||||
assert PARENT["user_api_key_auth"]["budget_reservation"] is not None
|
||||
|
||||
|
||||
def test_forwarded_metadata_empty_parent_stays_unstamped():
|
||||
assert forwarded_internal_call_metadata(None, "autorouter_classifier") == {}
|
||||
assert forwarded_internal_call_metadata({}, "autorouter_classifier") == {}
|
||||
|
||||
|
||||
def test_sanitized_forwardable_metadata_keeps_only_identity_and_always_stamps():
|
||||
result = sanitized_forwardable_call_metadata(PARENT, SHADOW_EVAL_ROUTER_CALL_ORIGIN)
|
||||
|
||||
assert result[INTERNAL_CALL_ORIGIN_METADATA_KEY] == SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
assert result["user_api_key"] == "sk-hash"
|
||||
assert result["user_api_key_team_id"] == "team-1"
|
||||
assert result["user_api_key_auth"] == {"api_key": "sk-hash"}
|
||||
assert "routing_decision" not in result
|
||||
assert "headers" not in result
|
||||
assert "user_api_key_budget_reservation" not in result
|
||||
|
||||
assert sanitized_forwardable_call_metadata({}, SHADOW_EVAL_ROUTER_CALL_ORIGIN) == {
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY: SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
}
|
||||
|
||||
|
||||
class TestSubCallMetadataSanitization:
|
||||
"""The proxy cost callback must not be able to recover the parent budget reservation
|
||||
from sub-call metadata, in either of the shapes it knows how to read."""
|
||||
|
||||
def test_cost_callback_cannot_recover_reservation_from_sanitized_metadata(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.proxy_track_cost_callback import (
|
||||
_get_budget_reservation_from_metadata,
|
||||
)
|
||||
|
||||
reservation = {"reserved_cost": 1.0}
|
||||
auth_shapes = (
|
||||
{"models": ["gpt-4o"], "budget_reservation": dict(reservation)},
|
||||
UserAPIKeyAuth(api_key="sk-abc", budget_reservation=dict(reservation)),
|
||||
)
|
||||
for auth in auth_shapes:
|
||||
metadata = {
|
||||
"user_api_key_hash": "hash-abc",
|
||||
"user_api_key_budget_reservation": dict(reservation),
|
||||
"user_api_key_auth": auth,
|
||||
}
|
||||
assert _get_budget_reservation_from_metadata(metadata) == reservation
|
||||
|
||||
sanitized = forwarded_internal_call_metadata(metadata, "autorouter_classifier")
|
||||
assert sanitized is not None
|
||||
assert sanitized["user_api_key_auth"] is not None
|
||||
assert _get_budget_reservation_from_metadata(sanitized) is None
|
||||
|
||||
def test_classifier_buckets_keep_non_spend_fields_on_a_chat_completions_parent(self):
|
||||
"""Drives the real resolver over the buckets the embedding classifier builds.
|
||||
|
||||
An absent bucket must stay empty rather than carry a lone origin stamp:
|
||||
get_litellm_metadata_from_kwargs prefers litellm_metadata whenever truthy, so an
|
||||
origin-only dict would make an empty litellm_metadata win and silently drop
|
||||
requester_ip_address, tags and spend_logs_metadata from the classifier's row."""
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
|
||||
parent = {
|
||||
"user_api_key": "sk-abc",
|
||||
"requester_ip_address": "10.0.0.1",
|
||||
"spend_logs_metadata": {"team_note": "keep me"},
|
||||
"tags": ["prod"],
|
||||
}
|
||||
resolved = get_litellm_metadata_from_kwargs(
|
||||
{
|
||||
"litellm_params": {
|
||||
"metadata": forwarded_internal_call_metadata(parent, "autorouter_classifier"),
|
||||
"litellm_metadata": forwarded_internal_call_metadata(None, "autorouter_classifier"),
|
||||
}
|
||||
}
|
||||
)
|
||||
assert resolved["internal_call_origin"] == "autorouter_classifier"
|
||||
assert resolved["requester_ip_address"] == "10.0.0.1"
|
||||
assert resolved["spend_logs_metadata"] == {"team_note": "keep me"}
|
||||
assert resolved["tags"] == ["prod"]
|
||||
|
||||
def test_sanitized_auth_keeps_access_group_fields_and_leaves_original_untouched(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
auth = UserAPIKeyAuth(
|
||||
api_key="sk-abc",
|
||||
team_id="team-1",
|
||||
budget_reservation={"reserved_cost": 1.0},
|
||||
)
|
||||
sanitized = forwarded_internal_call_metadata({"user_api_key_auth": auth}, "autorouter_classifier")
|
||||
sanitized_auth = sanitized["user_api_key_auth"]
|
||||
assert sanitized_auth.budget_reservation is None
|
||||
assert sanitized_auth.team_id == "team-1"
|
||||
assert sanitized_auth.api_key == auth.api_key
|
||||
assert auth.budget_reservation == {"reserved_cost": 1.0}
|
||||
92
tests/test_litellm/litellm_core_utils/test_llm_judge.py
Normal file
92
tests/test_litellm/litellm_core_utils/test_llm_judge.py
Normal file
|
|
@ -0,0 +1,92 @@
|
|||
"""Unit tests for the shared LLM-judge primitives: verdict parsing, router resolution, dispatch."""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.llm_judge import (
|
||||
extract_text_from_content,
|
||||
judge_acompletion,
|
||||
parse_json_verdict,
|
||||
router_resolves_model,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw,expected",
|
||||
[
|
||||
('{"preference": "A", "confidence": 0.9}', "A"),
|
||||
('Here it is:\n```json\n{"preference": "B"}\n```\nDone.', "B"),
|
||||
('```\n{"preference": "tie"}\n```', "tie"),
|
||||
('Verdict: {"preference": "A", "confidence": 0.5} final.', "A"),
|
||||
],
|
||||
)
|
||||
def test_parse_json_verdict_tolerates_fences_and_prose(raw, expected):
|
||||
assert parse_json_verdict(raw)["preference"] == expected
|
||||
|
||||
|
||||
def test_parse_json_verdict_rejects_non_object():
|
||||
with pytest.raises(ValueError):
|
||||
parse_json_verdict('["not", "an", "object"]')
|
||||
with pytest.raises((json.JSONDecodeError, ValueError)):
|
||||
parse_json_verdict("no json here at all")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"content,expected",
|
||||
[
|
||||
("hello", "hello"),
|
||||
([{"type": "text", "text": "a"}, {"type": "image_url", "image_url": {}}, {"type": "text", "text": "b"}], "a b"),
|
||||
(42, ""),
|
||||
(None, ""),
|
||||
],
|
||||
)
|
||||
def test_extract_text_from_content(content, expected):
|
||||
assert extract_text_from_content(content) == expected
|
||||
|
||||
|
||||
def _router(alias=(), deployments=False) -> MagicMock:
|
||||
router = MagicMock()
|
||||
router.model_group_alias = dict.fromkeys(alias, "x")
|
||||
router.get_model_list = MagicMock(
|
||||
return_value=[{"litellm_params": {"model": "openai/gpt-4o"}}] if deployments else None
|
||||
)
|
||||
router.acompletion = AsyncMock(return_value={"choices": [{"message": {"content": "router answer"}}]})
|
||||
return router
|
||||
|
||||
|
||||
def test_router_resolves_model_matrix():
|
||||
assert router_resolves_model(None, "gpt-4o") is False
|
||||
assert router_resolves_model(_router(), "gpt-4o") is False
|
||||
assert router_resolves_model(_router(alias=("gpt-4o",)), "gpt-4o") is True
|
||||
assert router_resolves_model(_router(deployments=True), "gpt-4o") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_judge_acompletion_prefers_router_and_disables_retries():
|
||||
router = _router(deployments=True)
|
||||
response = await judge_acompletion(router, "judge-model", [{"role": "user", "content": "hi"}], temperature=0)
|
||||
assert response == {"choices": [{"message": {"content": "router answer"}}]}
|
||||
_, kwargs = router.acompletion.call_args
|
||||
assert kwargs["num_retries"] == 0
|
||||
assert kwargs["fallbacks"] == []
|
||||
assert kwargs["temperature"] == 0
|
||||
assert kwargs["drop_params"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_judge_acompletion_falls_back_to_sdk_for_unconfigured_model(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm as litellm_module
|
||||
|
||||
sdk = AsyncMock(return_value={"choices": [{"message": {"content": "sdk answer"}}]})
|
||||
monkeypatch.setattr(litellm_module, "acompletion", sdk)
|
||||
router = _router()
|
||||
|
||||
response = await judge_acompletion(router, "anthropic/claude-sonnet-5", [{"role": "user", "content": "hi"}])
|
||||
|
||||
assert response == {"choices": [{"message": {"content": "sdk answer"}}]}
|
||||
router.acompletion.assert_not_called()
|
||||
assert sdk.call_args.kwargs["model"] == "anthropic/claude-sonnet-5"
|
||||
assert sdk.call_args.kwargs["num_retries"] == 0
|
||||
assert sdk.call_args.kwargs["drop_params"] is True
|
||||
|
|
@ -279,3 +279,10 @@ def test_every_drain_trigger_reads_the_one_queue_census_owner():
|
|||
assert queue in owner_source, queue
|
||||
for site in (proxy_utils.update_spend, proxy_utils.update_spend_logs_job, proxy_utils._monitor_spend_logs_queue):
|
||||
assert "_total_queued_spend_transactions" in inspect.getsource(site), site.__name__
|
||||
|
||||
|
||||
def test_internal_call_origin_never_reaches_the_rollup():
|
||||
"""A shadow eval's duplicate carries a real routing_decision, so the decision-presence
|
||||
gate alone would count it; the internal_call_origin stamp must exclude it."""
|
||||
assert _build(metadata=_metadata(internal_call_origin="shadow_eval_router")) is None
|
||||
assert _build() is not None
|
||||
|
|
|
|||
|
|
@ -2221,3 +2221,50 @@ async def test_commit_spend_updates_to_db_does_not_stamp_key_settings_updated_at
|
|||
assert call_kwargs["where"] == {"token": token}
|
||||
assert set(call_kwargs["data"]) == {"spend", "last_active"}
|
||||
assert call_kwargs["data"]["spend"] == {"increment": response_cost}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_daily_transaction_internal_call_keeps_spend_but_not_request_counts():
|
||||
"""Internal sub-calls (auto-router classifier, shadow eval's shadow and judge) bill
|
||||
spend and tokens to the key but are not requests the caller made: api_requests,
|
||||
successful_requests, and autorouter_savings_spend must all stay zero for them."""
|
||||
writer = DBSpendUpdateWriter()
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.get_request_status = MagicMock(return_value="success")
|
||||
|
||||
def _payload(metadata: dict) -> dict:
|
||||
return {
|
||||
"request_id": "req-internal-1",
|
||||
"user": "test-user",
|
||||
"startTime": "2026-08-11T00:00:00",
|
||||
"api_key": "test-key",
|
||||
"model": "claude-sonnet-5",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"model_group": "claude-sonnet-5",
|
||||
"call_type": "acompletion",
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 10,
|
||||
"spend": 0.05,
|
||||
"metadata": json.dumps(metadata),
|
||||
}
|
||||
|
||||
internal = await writer._common_add_spend_log_transaction_to_daily_transaction(
|
||||
payload=_payload({"internal_call_origin": "shadow_eval_judge"}),
|
||||
prisma_client=mock_prisma,
|
||||
type="user",
|
||||
)
|
||||
user_sent = await writer._common_add_spend_log_transaction_to_daily_transaction(
|
||||
payload=_payload({}),
|
||||
prisma_client=mock_prisma,
|
||||
type="user",
|
||||
)
|
||||
|
||||
assert internal is not None and user_sent is not None
|
||||
assert internal["spend"] == 0.05
|
||||
assert internal["prompt_tokens"] == 100
|
||||
assert internal["api_requests"] == 0
|
||||
assert internal["successful_requests"] == 0
|
||||
assert internal["failed_requests"] == 0
|
||||
assert internal["autorouter_savings_spend"] == 0.0
|
||||
assert user_sent["api_requests"] == 1
|
||||
assert user_sent["successful_requests"] == 1
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from fastapi import HTTPException
|
|||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
|
||||
|
|
@ -5554,3 +5555,41 @@ async def test_configured_estimate_blocks_the_overrun_the_static_floor_admits(mo
|
|||
|
||||
assert await admitted({}) == 7
|
||||
assert await admitted({"default_estimated_output_tokens": 3000}) == 2
|
||||
|
||||
|
||||
def test_internal_call_origin_success_ops_are_skipped():
|
||||
"""Internal sub-calls (auto-router classifier, shadow eval shadow/judge) bill spend
|
||||
to the caller's key but must not consume its TPM counters: the same kwargs charge
|
||||
ops without the origin stamp and none with it."""
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
response = ModelResponse(
|
||||
id="internal-origin-tpm",
|
||||
object="chat.completion",
|
||||
created=int(datetime.now().timestamp()),
|
||||
model="gpt-4o-mini",
|
||||
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
|
||||
choices=[],
|
||||
)
|
||||
|
||||
def _kwargs(metadata: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {
|
||||
"standard_logging_object": {
|
||||
"metadata": {"user_api_key_hash": hash_token("sk-internal-origin")}
|
||||
},
|
||||
"litellm_params": {"metadata": metadata},
|
||||
"model": "gpt-4o-mini",
|
||||
}
|
||||
|
||||
charged = handler._build_success_event_pipeline_operations(
|
||||
kwargs=_kwargs({}), response_obj=response, rate_limit_type="output"
|
||||
)
|
||||
skipped = handler._build_success_event_pipeline_operations(
|
||||
kwargs=_kwargs({INTERNAL_CALL_ORIGIN_METADATA_KEY: "shadow_eval_judge"}),
|
||||
response_obj=response,
|
||||
rate_limit_type="output",
|
||||
)
|
||||
|
||||
assert charged
|
||||
assert skipped == []
|
||||
|
|
|
|||
|
|
@ -466,3 +466,276 @@ class TestAutoRouterBenchmarks:
|
|||
end_date="2026-08-01",
|
||||
)
|
||||
assert response.groups[0].tier_turns == expected
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shadow eval endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.management_endpoints.auto_router_endpoints import (
|
||||
get_shadow_eval_job,
|
||||
list_shadow_eval_jobs,
|
||||
start_shadow_eval,
|
||||
stop_shadow_eval_job,
|
||||
)
|
||||
from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalJobResponse, StartShadowEvalRequest
|
||||
|
||||
VIEWER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, api_key="sk-view", user_id="viewer")
|
||||
NON_ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user", user_id="user")
|
||||
|
||||
|
||||
def _shadow_router() -> MagicMock:
|
||||
router = MagicMock()
|
||||
router.auto_routers = {}
|
||||
router.complexity_routers = {"my-router": [MagicMock()]}
|
||||
router.adaptive_routers = {}
|
||||
router.quality_routers = {}
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=None)
|
||||
return router
|
||||
|
||||
|
||||
def _job_record(**overrides: object) -> MagicMock:
|
||||
"""Spec'd like a real prisma row: only the table's columns exist as attributes, so
|
||||
from_attributes validation falls back to model defaults for everything else."""
|
||||
defaults = {
|
||||
"id": "job-1",
|
||||
"api_key_id": "key-hash",
|
||||
"router_name": "my-router",
|
||||
"judge_model": "anthropic/claude-sonnet-5",
|
||||
"shadow_percentage": 10.0,
|
||||
"max_turns": 200,
|
||||
"created_at": datetime(2026, 8, 11, tzinfo=timezone.utc),
|
||||
"ends_at": datetime.now(timezone.utc) + timedelta(days=7),
|
||||
"stopped_at": None,
|
||||
}
|
||||
fields = {**defaults, **overrides}
|
||||
record = MagicMock(spec=list(fields))
|
||||
for key, value in fields.items():
|
||||
setattr(record, key, value)
|
||||
return record
|
||||
|
||||
|
||||
def _shadow_prisma(active_job=None, agg_rows=None) -> MagicMock:
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=MagicMock())
|
||||
prisma.db.execute_raw = AsyncMock(return_value=0)
|
||||
prisma.db.litellm_shadowevaljob.find_first = AsyncMock(return_value=active_job)
|
||||
prisma.db.litellm_shadowevaljob.find_unique = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_shadowevaljob.create = AsyncMock(return_value=_job_record())
|
||||
prisma.db.litellm_shadowevaljob.update = AsyncMock(
|
||||
return_value=_job_record(stopped_at=datetime.now(timezone.utc))
|
||||
)
|
||||
prisma.db.litellm_shadowevalattempt.find_first = AsyncMock(return_value=None)
|
||||
|
||||
async def query_raw(sql: str, *params: object):
|
||||
if "FILTER (WHERE outcome != 'error')::int AS judged_count" in sql:
|
||||
return [{"judged_count": 10, "error_count": 2, "judge_spend": 0.031}]
|
||||
return agg_rows if agg_rows is not None else []
|
||||
|
||||
prisma.db.query_raw = AsyncMock(side_effect=query_raw)
|
||||
return prisma
|
||||
|
||||
|
||||
def _start_request(**overrides: object) -> StartShadowEvalRequest:
|
||||
payload = {
|
||||
"api_key_id": "key-hash",
|
||||
"router_name": "my-router",
|
||||
"shadow_percentage": 10.0,
|
||||
"judge_model": "anthropic/claude-sonnet-5",
|
||||
"duration_days": 7,
|
||||
"max_turns": 200,
|
||||
}
|
||||
payload.update(overrides)
|
||||
return StartShadowEvalRequest.model_validate(payload)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_creates_job_and_frees_expired_or_exhausted_ones(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Expiry and turn-budget exhaustion both end sampling on their own; either must
|
||||
release the one-active-per-key index so a new eval can start."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
response = await start_shadow_eval(_start_request(), ADMIN)
|
||||
|
||||
assert response.status == "running"
|
||||
assert response.max_turns == 200
|
||||
assert response.judged_count is None
|
||||
sweep_sql, sweep_key = prisma.db.execute_raw.call_args.args
|
||||
assert "stopped_at IS NULL" in sweep_sql
|
||||
assert "ends_at <= NOW()" in sweep_sql
|
||||
assert ">= j.max_turns" in sweep_sql
|
||||
assert sweep_key == "key-hash"
|
||||
create_data = prisma.db.litellm_shadowevaljob.create.call_args.kwargs["data"]
|
||||
assert create_data["api_key_id"] == "key-hash"
|
||||
assert create_data["created_by"] == "admin"
|
||||
assert "status" not in create_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"caller,request_overrides,active,expected_status",
|
||||
[
|
||||
(NON_ADMIN, {}, None, 403),
|
||||
(VIEWER, {}, None, 403),
|
||||
(ADMIN, {"router_name": "not-a-router"}, None, 400),
|
||||
(ADMIN, {"judge_model": "not/a real model!"}, None, 400),
|
||||
(ADMIN, {"judge_model": "my-router"}, None, 400),
|
||||
(ADMIN, {}, "active", 409),
|
||||
],
|
||||
ids=["non-admin", "view-only", "unknown-router", "unresolvable-judge", "router-as-judge", "already-active"],
|
||||
)
|
||||
async def test_start_shadow_eval_rejections(
|
||||
monkeypatch: pytest.MonkeyPatch, caller, request_overrides, active, expected_status
|
||||
):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma(active_job=_job_record() if active else None)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await start_shadow_eval(_start_request(**request_overrides), caller)
|
||||
assert exc.value.status_code == expected_status
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_rejects_a_key_this_proxy_does_not_know(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A typo'd api_key_id would otherwise create a job no traffic can ever match."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await start_shadow_eval(_start_request(), ADMIN)
|
||||
assert exc.value.status_code == 400
|
||||
assert "not a key on this proxy" in exc.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_concurrent_unique_violation_is_a_409(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from prisma.errors import UniqueViolationError
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
prisma.db.litellm_shadowevaljob.create = AsyncMock(
|
||||
side_effect=UniqueViolationError(MagicMock(message="unique constraint"))
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await start_shadow_eval(_start_request(), ADMIN)
|
||||
assert exc.value.status_code == 409
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_shadow_eval_job_derives_counts_spend_and_stratified_results(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
tier_rows = [
|
||||
{"grp": "SIMPLE", "turn_count": 8, "real_wins": 2, "shadow_wins": 4, "ties": 2, "avg_confidence": 0.8},
|
||||
{"grp": "REASONING", "turn_count": 2, "real_wins": 2, "shadow_wins": 0, "ties": 0, "avg_confidence": 0.9},
|
||||
]
|
||||
prisma = _shadow_prisma(agg_rows=tier_rows)
|
||||
prisma.db.litellm_shadowevaljob.find_unique = AsyncMock(return_value=_job_record())
|
||||
prisma.db.litellm_shadowevalattempt.find_first = AsyncMock(
|
||||
return_value=MagicMock(error="judge call failed: boom")
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
|
||||
response = await get_shadow_eval_job("job-1", VIEWER)
|
||||
|
||||
assert response.job_id == "job-1"
|
||||
assert response.status == "running"
|
||||
assert response.judged_count == 10
|
||||
assert response.error_count == 2
|
||||
assert response.judge_spend == 0.031
|
||||
assert response.last_error == "judge call failed: boom"
|
||||
assert [s.group for s in response.results.by_tier] == ["SIMPLE", "REASONING"]
|
||||
assert response.results.by_tier[0].shadow_win_rate_pct == 50.0
|
||||
assert response.results.overall_shadow_win_rate_pct == 40.0
|
||||
assert response.results.overall_tie_rate_pct == 20.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_shadow_eval_job_404s_and_gates_on_role(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", _shadow_prisma())
|
||||
|
||||
with pytest.raises(HTTPException) as missing:
|
||||
await get_shadow_eval_job("nope", VIEWER)
|
||||
assert missing.value.status_code == 404
|
||||
|
||||
with pytest.raises(HTTPException) as forbidden:
|
||||
await get_shadow_eval_job("job-1", NON_ADMIN)
|
||||
assert forbidden.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_shadow_eval_jobs_returns_derived_status_without_aggregates(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
prisma.db.litellm_shadowevaljob.find_many = AsyncMock(
|
||||
return_value=[
|
||||
_job_record(),
|
||||
_job_record(id="job-2", ends_at=datetime.now(timezone.utc) - timedelta(days=1)),
|
||||
_job_record(id="job-3", stopped_at=datetime.now(timezone.utc)),
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
|
||||
jobs = await list_shadow_eval_jobs(VIEWER, api_key_id=None, limit=50)
|
||||
|
||||
assert [job.status for job in jobs] == ["running", "completed", "stopped"]
|
||||
swept = ShadowEvalJobResponse.model_validate(
|
||||
_job_record(
|
||||
id="job-4",
|
||||
ends_at=datetime.now(timezone.utc) - timedelta(days=1),
|
||||
stopped_at=datetime.now(timezone.utc),
|
||||
),
|
||||
from_attributes=True,
|
||||
)
|
||||
assert swept.status == "completed"
|
||||
assert all(job.judged_count is None and job.results is None for job in jobs)
|
||||
assert prisma.db.query_raw.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_shadow_eval_sets_stopped_at_and_rejects_non_running(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
prisma.db.litellm_shadowevaljob.find_unique = AsyncMock(return_value=_job_record())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
|
||||
stopped = await stop_shadow_eval_job("job-1", ADMIN)
|
||||
assert stopped.status == "stopped"
|
||||
update = prisma.db.litellm_shadowevaljob.update.call_args.kwargs
|
||||
assert set(update["data"]) == {"stopped_at"}
|
||||
|
||||
prisma.db.litellm_shadowevaljob.find_unique = AsyncMock(
|
||||
return_value=_job_record(ends_at=datetime.now(timezone.utc) - timedelta(days=1))
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await stop_shadow_eval_job("job-1", ADMIN)
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
with pytest.raises(HTTPException) as forbidden:
|
||||
await stop_shadow_eval_job("job-1", VIEWER)
|
||||
assert forbidden.value.status_code == 403
|
||||
|
|
|
|||
|
|
@ -524,8 +524,9 @@ def test_load_from_azure_key_vault_missing_uri_failure_is_swallowed(monkeypatch)
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_cost_tracking_adds_two_callbacks_when_prisma_set(monkeypatch):
|
||||
def test_cost_tracking_adds_db_and_shadow_eval_callbacks_when_prisma_set(monkeypatch):
|
||||
import litellm
|
||||
from litellm.integrations.shadow_eval_logger import ShadowEvalLogger
|
||||
|
||||
fake_prisma = MagicMock()
|
||||
monkeypatch.setattr(ps, "prisma_client", fake_prisma, raising=False)
|
||||
|
|
@ -535,16 +536,19 @@ def test_cost_tracking_adds_two_callbacks_when_prisma_set(monkeypatch):
|
|||
before_callbacks = len(litellm.callbacks)
|
||||
before_async = len(litellm._async_success_callback)
|
||||
|
||||
cost_tracking()
|
||||
cost_tracking()
|
||||
|
||||
observed = {
|
||||
"added_to_callbacks": len(litellm.callbacks) - before_callbacks,
|
||||
"added_to_async_success": len(litellm._async_success_callback) - before_async,
|
||||
"shadow_eval_loggers": sum(isinstance(cb, ShadowEvalLogger) for cb in litellm.callbacks),
|
||||
"prisma_was_set": True,
|
||||
}
|
||||
assert normalize(observed) == {
|
||||
"added_to_callbacks": 1,
|
||||
"added_to_callbacks": 2,
|
||||
"added_to_async_success": 1,
|
||||
"shadow_eval_loggers": 1,
|
||||
"prisma_was_set": True,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -3449,98 +3449,6 @@ class TestKeywordOverrideEdgeCases:
|
|||
assert result.model in {"gpt-4o-mini", "gpt-4o", "claude-sonnet-4-20250514", "o1-preview"}
|
||||
|
||||
|
||||
class TestSubCallMetadataSanitization:
|
||||
"""The proxy cost callback must not be able to recover the parent budget reservation
|
||||
from sub-call metadata, in either of the shapes it knows how to read."""
|
||||
|
||||
def test_cost_callback_cannot_recover_reservation_from_sanitized_metadata(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.proxy_track_cost_callback import (
|
||||
_get_budget_reservation_from_metadata,
|
||||
)
|
||||
from litellm.router_strategy.complexity_router.complexity_router import (
|
||||
_classifier_call_metadata,
|
||||
)
|
||||
|
||||
reservation = {"reserved_cost": 1.0}
|
||||
auth_shapes = (
|
||||
{"models": ["gpt-4o"], "budget_reservation": dict(reservation)},
|
||||
UserAPIKeyAuth(api_key="sk-abc", budget_reservation=dict(reservation)),
|
||||
)
|
||||
for auth in auth_shapes:
|
||||
metadata = {
|
||||
"user_api_key_hash": "hash-abc",
|
||||
"user_api_key_budget_reservation": dict(reservation),
|
||||
"user_api_key_auth": auth,
|
||||
}
|
||||
assert _get_budget_reservation_from_metadata(metadata) == reservation
|
||||
|
||||
sanitized = _classifier_call_metadata(metadata)
|
||||
assert sanitized is not None
|
||||
assert sanitized["user_api_key_auth"] is not None
|
||||
assert _get_budget_reservation_from_metadata(sanitized) is None
|
||||
|
||||
def test_absent_parent_bucket_stays_empty(self):
|
||||
"""An absent bucket must not be materialized just to carry the origin.
|
||||
|
||||
The embedding path passes both buckets, and get_litellm_metadata_from_kwargs
|
||||
prefers litellm_metadata whenever it is truthy, backfilling only user_api_key*
|
||||
keys from metadata. Returning an origin-only dict here would make a chat
|
||||
completions parent's empty litellm_metadata win and silently drop
|
||||
requester_ip_address, tags and spend_logs_metadata from the classifier's row."""
|
||||
from litellm.router_strategy.complexity_router.complexity_router import (
|
||||
_classifier_call_metadata,
|
||||
)
|
||||
|
||||
for absent in (None, {}):
|
||||
assert _classifier_call_metadata(absent) == {}
|
||||
|
||||
def test_classifier_buckets_keep_non_spend_fields_on_a_chat_completions_parent(self):
|
||||
"""Drives the real resolver over the buckets the embedding classifier builds."""
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
from litellm.router_strategy.complexity_router.complexity_router import (
|
||||
_classifier_call_metadata,
|
||||
)
|
||||
|
||||
parent = {
|
||||
"user_api_key": "sk-abc",
|
||||
"requester_ip_address": "10.0.0.1",
|
||||
"spend_logs_metadata": {"team_note": "keep me"},
|
||||
"tags": ["prod"],
|
||||
}
|
||||
resolved = get_litellm_metadata_from_kwargs(
|
||||
{
|
||||
"litellm_params": {
|
||||
"metadata": _classifier_call_metadata(parent),
|
||||
"litellm_metadata": _classifier_call_metadata(None),
|
||||
}
|
||||
}
|
||||
)
|
||||
assert resolved["internal_call_origin"] == "autorouter_classifier"
|
||||
assert resolved["requester_ip_address"] == "10.0.0.1"
|
||||
assert resolved["spend_logs_metadata"] == {"team_note": "keep me"}
|
||||
assert resolved["tags"] == ["prod"]
|
||||
|
||||
def test_sanitized_auth_keeps_access_group_fields_and_leaves_original_untouched(self):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.router_strategy.complexity_router.complexity_router import (
|
||||
_classifier_call_metadata,
|
||||
)
|
||||
|
||||
auth = UserAPIKeyAuth(
|
||||
api_key="sk-abc",
|
||||
team_id="team-1",
|
||||
budget_reservation={"reserved_cost": 1.0},
|
||||
)
|
||||
sanitized = _classifier_call_metadata({"user_api_key_auth": auth})
|
||||
assert sanitized is not None
|
||||
sanitized_auth = sanitized["user_api_key_auth"]
|
||||
assert sanitized_auth.budget_reservation is None
|
||||
assert sanitized_auth.team_id == "team-1"
|
||||
assert sanitized_auth.api_key == auth.api_key
|
||||
assert auth.budget_reservation == {"reserved_cost": 1.0}
|
||||
|
||||
|
||||
class TestRoutingDecisionCauseLogging:
|
||||
"""The info log must name what drove each routing decision so an operator can tell a
|
||||
literal keyword match, a semantic keyword match, and the complexity scorer apart.
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23001
|
||||
"limit": 22943
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27146
|
||||
"limit": 27141
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT006": {
|
||||
"limit": 1077
|
||||
"limit": 1074
|
||||
},
|
||||
"LIT007": {
|
||||
"limit": 0
|
||||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16731
|
||||
"limit": 16722
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5596
|
||||
|
|
|
|||
|
|
@ -3794,26 +3794,11 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/CollapsibleMessage.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/DrawerHeader.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/HistoryTree.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/JsonViewer.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx": {
|
||||
"no-nested-ternary": {
|
||||
"count": 3
|
||||
|
|
@ -3833,11 +3818,6 @@
|
|||
"count": 2
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/OutputCard.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.test.tsx": {
|
||||
"unused-imports/no-unused-imports": {
|
||||
"count": 2
|
||||
|
|
@ -3848,26 +3828,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/SimpleMessageBlock.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/SimpleToolCallBlock.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/TokenFlow.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/TruncatedValue.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts": {
|
||||
"react-hooks/immutability": {
|
||||
"count": 2
|
||||
|
|
@ -4018,4 +3978,4 @@
|
|||
"count": 1
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -37,4 +37,18 @@ describe("CollapsibleMessage", () => {
|
|||
await user.click(screen.getByText("SYSTEM"));
|
||||
expect(screen.getByText("Toggle me")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should expand with Enter and collapse with Space from the keyboard", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<CollapsibleMessage label="SYSTEM" content="Toggle me" defaultExpanded={false} />);
|
||||
|
||||
expect(screen.getByText("Toggle me")).not.toBeVisible();
|
||||
|
||||
await user.tab();
|
||||
await user.keyboard("{Enter}");
|
||||
expect(screen.getByText("Toggle me")).toBeVisible();
|
||||
|
||||
await user.keyboard(" ");
|
||||
expect(screen.getByText("Toggle me")).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -4,10 +4,8 @@
|
|||
*/
|
||||
|
||||
import { useState } from "react";
|
||||
import { Typography } from "antd";
|
||||
import { DownOutlined, RightOutlined } from "@ant-design/icons";
|
||||
|
||||
const { Text } = Typography;
|
||||
import { ChevronDown, ChevronRight } from "lucide-react";
|
||||
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
|
||||
|
||||
interface CollapsibleMessageProps {
|
||||
label: string;
|
||||
|
|
@ -17,7 +15,6 @@ interface CollapsibleMessageProps {
|
|||
|
||||
export function CollapsibleMessage({ label, content, defaultExpanded = false }: CollapsibleMessageProps) {
|
||||
const [isExpanded, setIsExpanded] = useState(defaultExpanded);
|
||||
const [isHovered, setIsHovered] = useState(false);
|
||||
const charCount = content?.length || 0;
|
||||
|
||||
if (!content || charCount === 0) {
|
||||
|
|
@ -25,60 +22,23 @@ export function CollapsibleMessage({ label, content, defaultExpanded = false }:
|
|||
}
|
||||
|
||||
return (
|
||||
<div style={{ marginBottom: 8 }}>
|
||||
{/* Clickable Header with hover state */}
|
||||
<div
|
||||
onClick={() => setIsExpanded(!isExpanded)}
|
||||
onMouseEnter={() => setIsHovered(true)}
|
||||
onMouseLeave={() => setIsHovered(false)}
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 6,
|
||||
cursor: "pointer",
|
||||
padding: "4px 0",
|
||||
borderRadius: 4,
|
||||
background: isHovered ? "#f5f5f5" : "transparent",
|
||||
transition: "background 0.15s ease",
|
||||
marginBottom: isExpanded ? 4 : 0,
|
||||
}}
|
||||
>
|
||||
<Collapsible open={isExpanded} onOpenChange={setIsExpanded} className="mb-2">
|
||||
<CollapsibleTrigger className="flex w-full items-center gap-1.5 rounded py-1 text-left transition-colors hover:bg-muted">
|
||||
{isExpanded ? (
|
||||
<DownOutlined style={{ fontSize: 10, color: "#8c8c8c" }} />
|
||||
<ChevronDown className="size-3 shrink-0 text-muted-foreground" />
|
||||
) : (
|
||||
<RightOutlined style={{ fontSize: 10, color: "#8c8c8c" }} />
|
||||
<ChevronRight className="size-3 shrink-0 text-muted-foreground" />
|
||||
)}
|
||||
<Text type="secondary" style={{ fontSize: 10, letterSpacing: "0.5px", textTransform: "uppercase" }}>
|
||||
{label}
|
||||
</Text>
|
||||
<Text type="secondary" style={{ fontSize: 10 }}>
|
||||
({charCount.toLocaleString()} chars)
|
||||
</Text>
|
||||
</div>
|
||||
<span className="text-[10px] uppercase tracking-[0.5px] text-muted-foreground">{label}</span>
|
||||
<span className="text-[10px] text-muted-foreground">({charCount.toLocaleString()} chars)</span>
|
||||
</CollapsibleTrigger>
|
||||
|
||||
{/* Content with smooth animation */}
|
||||
<div
|
||||
style={{
|
||||
maxHeight: isExpanded ? "2000px" : "0px",
|
||||
overflow: "hidden",
|
||||
transition: "max-height 0.2s ease-out, opacity 0.2s ease-out",
|
||||
opacity: isExpanded ? 1 : 0,
|
||||
}}
|
||||
<CollapsibleContent
|
||||
keepMounted
|
||||
className="mt-1 border-l border-border pl-4 text-[13px] leading-[1.7] break-words whitespace-pre-wrap text-foreground"
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
paddingLeft: 16,
|
||||
fontSize: 13,
|
||||
lineHeight: 1.7,
|
||||
color: "#262626",
|
||||
borderLeft: "1px solid #f0f0f0",
|
||||
whiteSpace: "pre-wrap",
|
||||
wordBreak: "break-word",
|
||||
}}
|
||||
>
|
||||
{content}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{content}
|
||||
</CollapsibleContent>
|
||||
</Collapsible>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -41,4 +41,22 @@ describe("HistoryTree", () => {
|
|||
expect(screen.getByText("Hello")).toBeInTheDocument();
|
||||
expect(screen.getByText("Hi there")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should expand with Enter and collapse with Space from the keyboard", async () => {
|
||||
const user = userEvent.setup();
|
||||
const messages: ParsedMessage[] = [
|
||||
{ role: "user", content: "Hello" },
|
||||
{ role: "assistant", content: "Hi there" },
|
||||
];
|
||||
render(<HistoryTree messages={messages} />);
|
||||
|
||||
expect(screen.getByText("Hello")).not.toBeVisible();
|
||||
|
||||
await user.tab();
|
||||
await user.keyboard("{Enter}");
|
||||
expect(screen.getByText("Hello")).toBeVisible();
|
||||
|
||||
await user.keyboard(" ");
|
||||
expect(screen.getByText("Hello")).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -4,80 +4,46 @@
|
|||
*/
|
||||
|
||||
import { useState } from "react";
|
||||
import { Typography } from "antd";
|
||||
import { DownOutlined, RightOutlined } from "@ant-design/icons";
|
||||
import { ChevronDown, ChevronRight } from "lucide-react";
|
||||
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
|
||||
import { ParsedMessage } from "./prettyMessagesTypes";
|
||||
import { SimpleMessageBlock } from "./SimpleMessageBlock";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface HistoryTreeProps {
|
||||
messages: ParsedMessage[];
|
||||
}
|
||||
|
||||
export function HistoryTree({ messages }: HistoryTreeProps) {
|
||||
const [isExpanded, setIsExpanded] = useState(false);
|
||||
const [isHovered, setIsHovered] = useState(false);
|
||||
|
||||
if (messages.length === 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return (
|
||||
<div style={{ marginBottom: 8 }}>
|
||||
{/* Clickable Header with hover state */}
|
||||
<div
|
||||
onClick={() => setIsExpanded(!isExpanded)}
|
||||
onMouseEnter={() => setIsHovered(true)}
|
||||
onMouseLeave={() => setIsHovered(false)}
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 6,
|
||||
cursor: "pointer",
|
||||
padding: "4px 0",
|
||||
borderRadius: 4,
|
||||
background: isHovered ? "#f5f5f5" : "transparent",
|
||||
transition: "background 0.15s ease",
|
||||
marginBottom: isExpanded ? 4 : 0,
|
||||
}}
|
||||
>
|
||||
<Collapsible open={isExpanded} onOpenChange={setIsExpanded} className="mb-2">
|
||||
<CollapsibleTrigger className="flex w-full items-center gap-1.5 rounded py-1 text-left transition-colors hover:bg-muted">
|
||||
{isExpanded ? (
|
||||
<DownOutlined style={{ fontSize: 10, color: "#8c8c8c" }} />
|
||||
<ChevronDown className="size-3 shrink-0 text-muted-foreground" />
|
||||
) : (
|
||||
<RightOutlined style={{ fontSize: 10, color: "#8c8c8c" }} />
|
||||
<ChevronRight className="size-3 shrink-0 text-muted-foreground" />
|
||||
)}
|
||||
<Text type="secondary" style={{ fontSize: 10, letterSpacing: "0.5px", textTransform: "uppercase" }}>
|
||||
<span className="text-[10px] uppercase tracking-[0.5px] text-muted-foreground">
|
||||
HISTORY ({messages.length} message{messages.length !== 1 ? "s" : ""})
|
||||
</Text>
|
||||
</div>
|
||||
</span>
|
||||
</CollapsibleTrigger>
|
||||
|
||||
{/* Expanded Tree Content with smooth animation */}
|
||||
<div
|
||||
style={{
|
||||
maxHeight: isExpanded ? "2000px" : "0px",
|
||||
overflow: "hidden",
|
||||
transition: "max-height 0.2s ease-out, opacity 0.2s ease-out",
|
||||
opacity: isExpanded ? 1 : 0,
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
paddingLeft: 16,
|
||||
borderLeft: "1px solid #f0f0f0",
|
||||
}}
|
||||
>
|
||||
{messages.map((msg, index) => (
|
||||
<SimpleMessageBlock
|
||||
key={index}
|
||||
label={msg.role.toUpperCase()}
|
||||
content={msg.content}
|
||||
toolCalls={msg.toolCalls}
|
||||
isCompact={true}
|
||||
/>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<CollapsibleContent keepMounted className="mt-1 border-l border-border pl-4">
|
||||
{messages.map((msg, index) => (
|
||||
<SimpleMessageBlock
|
||||
key={index}
|
||||
label={msg.role.toUpperCase()}
|
||||
content={msg.content}
|
||||
toolCalls={msg.toolCalls}
|
||||
isCompact={true}
|
||||
/>
|
||||
))}
|
||||
</CollapsibleContent>
|
||||
</Collapsible>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,28 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, expect, it } from "vitest";
|
||||
import { JsonViewer } from "./JsonViewer";
|
||||
|
||||
describe("JsonViewer", () => {
|
||||
it("should render a placeholder and no tree when the log entry carries no payload", () => {
|
||||
render(<JsonViewer data={null} mode="formatted" />);
|
||||
|
||||
expect(screen.getByText("No data")).toBeInTheDocument();
|
||||
expect(screen.queryByRole("tree")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the payload as a tree exposing its keys", () => {
|
||||
render(<JsonViewer data={{ model: "claude-opus-4-5", stream: true }} mode="formatted" />);
|
||||
|
||||
expect(screen.getByRole("tree")).toBeInTheDocument();
|
||||
expect(screen.getByText(/model/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/stream/)).toBeInTheDocument();
|
||||
expect(screen.queryByText("No data")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should treat an empty payload as data rather than showing the placeholder", () => {
|
||||
render(<JsonViewer data={{}} mode="formatted" />);
|
||||
|
||||
expect(screen.getByRole("tree")).toBeInTheDocument();
|
||||
expect(screen.queryByText("No data")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,10 +1,7 @@
|
|||
import { Typography } from "antd";
|
||||
import { JsonView, defaultStyles } from "react-json-view-lite";
|
||||
import "react-json-view-lite/dist/index.css";
|
||||
import { JSON_MAX_HEIGHT, COLOR_BG_LIGHT, SPACING_LARGE } from "./constants";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface JsonViewerProps {
|
||||
data: any;
|
||||
mode: "formatted";
|
||||
|
|
@ -15,7 +12,7 @@ interface JsonViewerProps {
|
|||
* Uses an interactive tree component for easy navigation.
|
||||
*/
|
||||
export function JsonViewer({ data }: JsonViewerProps) {
|
||||
if (!data) return <Text type="secondary">No data</Text>;
|
||||
if (!data) return <span className="text-muted-foreground">No data</span>;
|
||||
|
||||
return (
|
||||
<div
|
||||
|
|
|
|||
|
|
@ -4,14 +4,12 @@
|
|||
*/
|
||||
|
||||
import { useState } from "react";
|
||||
import { Typography } from "antd";
|
||||
import MessageManager from "@/components/molecules/message_manager";
|
||||
import { COLOR_BORDER } from "./constants";
|
||||
import { ParsedMessage } from "./prettyMessagesTypes";
|
||||
import { SectionHeader } from "./SectionHeader";
|
||||
import { SimpleMessageBlock } from "./SimpleMessageBlock";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface OutputCardProps {
|
||||
message: ParsedMessage | null;
|
||||
completionTokens?: number;
|
||||
|
|
@ -24,55 +22,12 @@ export function OutputCard({ message, completionTokens, outputCost }: OutputCard
|
|||
const handleCopy = () => {
|
||||
if (!message) return;
|
||||
|
||||
const content = message.content || "";
|
||||
navigator.clipboard.writeText(content);
|
||||
navigator.clipboard.writeText(message.content || "");
|
||||
MessageManager.success("Output copied");
|
||||
};
|
||||
|
||||
if (!message) {
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #f0f0f0",
|
||||
borderRadius: 6,
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
<SectionHeader
|
||||
type="output"
|
||||
tokens={completionTokens}
|
||||
cost={outputCost}
|
||||
onCopy={handleCopy}
|
||||
isCollapsed={isCollapsed}
|
||||
onToggleCollapse={() => setIsCollapsed(!isCollapsed)}
|
||||
/>
|
||||
<div
|
||||
style={{
|
||||
maxHeight: isCollapsed ? "0px" : "10000px",
|
||||
overflow: "hidden",
|
||||
transition: "max-height 0.3s ease-out, opacity 0.3s ease-out",
|
||||
opacity: isCollapsed ? 0 : 1,
|
||||
}}
|
||||
>
|
||||
<div style={{ padding: "12px 16px" }}>
|
||||
<Text type="secondary" style={{ fontSize: 13, fontStyle: "italic" }}>
|
||||
No response data available
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #f0f0f0",
|
||||
borderRadius: 6,
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
{/* Datadog-style Header */}
|
||||
<div className="overflow-hidden rounded-md" style={{ border: `1px solid ${COLOR_BORDER}` }}>
|
||||
<SectionHeader
|
||||
type="output"
|
||||
tokens={completionTokens}
|
||||
|
|
@ -82,17 +37,16 @@ export function OutputCard({ message, completionTokens, outputCost }: OutputCard
|
|||
onToggleCollapse={() => setIsCollapsed(!isCollapsed)}
|
||||
/>
|
||||
|
||||
{/* Content */}
|
||||
<div
|
||||
style={{
|
||||
maxHeight: isCollapsed ? "0px" : "10000px",
|
||||
overflow: "hidden",
|
||||
transition: "max-height 0.3s ease-out, opacity 0.3s ease-out",
|
||||
opacity: isCollapsed ? 0 : 1,
|
||||
}}
|
||||
className="overflow-hidden transition-[max-height,opacity] duration-300 ease-out"
|
||||
style={{ maxHeight: isCollapsed ? "0px" : "10000px", opacity: isCollapsed ? 0 : 1 }}
|
||||
>
|
||||
<div style={{ padding: "12px 16px" }}>
|
||||
<SimpleMessageBlock label="ASSISTANT" content={message.content} toolCalls={message.toolCalls} />
|
||||
<div className="px-4 py-3">
|
||||
{message ? (
|
||||
<SimpleMessageBlock label="ASSISTANT" content={message.content} toolCalls={message.toolCalls} />
|
||||
) : (
|
||||
<span className="text-[13px] text-muted-foreground italic">No response data available</span>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -3,12 +3,10 @@
|
|||
* Used for messages in tree view and last user message
|
||||
*/
|
||||
|
||||
import { Typography } from "antd";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { ToolCall } from "./prettyMessagesTypes";
|
||||
import { SimpleToolCallBlock } from "./SimpleToolCallBlock";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface SimpleMessageBlockProps {
|
||||
label: string;
|
||||
content?: string;
|
||||
|
|
@ -27,30 +25,15 @@ export function SimpleMessageBlock({ label, content, toolCalls, isCompact = fals
|
|||
}
|
||||
|
||||
return (
|
||||
<div style={{ marginBottom: isCompact ? 8 : 0 }}>
|
||||
<Text
|
||||
type="secondary"
|
||||
style={{
|
||||
fontSize: 10,
|
||||
letterSpacing: "0.5px",
|
||||
textTransform: "uppercase",
|
||||
display: "block",
|
||||
marginBottom: 3,
|
||||
}}
|
||||
>
|
||||
{label}
|
||||
</Text>
|
||||
<div className={cn(isCompact && "mb-2")}>
|
||||
<span className="mb-[3px] block text-[10px] uppercase tracking-[0.5px] text-muted-foreground">{label}</span>
|
||||
|
||||
{displayContent && (
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
lineHeight: 1.7,
|
||||
color: "#262626",
|
||||
whiteSpace: "pre-wrap",
|
||||
wordBreak: "break-word",
|
||||
marginBottom: hasToolCalls ? 6 : 0,
|
||||
}}
|
||||
className={cn(
|
||||
"whitespace-pre-wrap break-words text-[13px] leading-[1.7] text-foreground",
|
||||
hasToolCalls && "mb-1.5",
|
||||
)}
|
||||
>
|
||||
{displayContent}
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -3,11 +3,9 @@
|
|||
* Used in compact/tree views
|
||||
*/
|
||||
|
||||
import { Typography } from "antd";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { ToolCall } from "./prettyMessagesTypes";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface SimpleToolCallBlockProps {
|
||||
tool: ToolCall;
|
||||
compact?: boolean;
|
||||
|
|
@ -16,46 +14,24 @@ interface SimpleToolCallBlockProps {
|
|||
export function SimpleToolCallBlock({ tool, compact = false }: SimpleToolCallBlockProps) {
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
background: "#f8f9fa",
|
||||
border: "1px solid #e9ecef",
|
||||
borderRadius: 6,
|
||||
padding: compact ? "6px 10px" : "10px 14px",
|
||||
marginTop: 8,
|
||||
fontFamily: "monospace",
|
||||
fontSize: 12,
|
||||
position: "relative",
|
||||
}}
|
||||
className={cn(
|
||||
"relative mt-2 rounded-md border border-border bg-muted font-mono text-xs",
|
||||
compact ? "px-2.5 py-1.5" : "px-3.5 py-2.5",
|
||||
)}
|
||||
>
|
||||
{/* Function badge */}
|
||||
<div
|
||||
style={{
|
||||
position: "absolute",
|
||||
top: -8,
|
||||
left: 12,
|
||||
background: "#fff",
|
||||
padding: "0 6px",
|
||||
fontSize: 10,
|
||||
color: "#8c8c8c",
|
||||
border: "1px solid #e9ecef",
|
||||
borderRadius: 3,
|
||||
}}
|
||||
>
|
||||
<div className="absolute -top-2 left-3 rounded-[3px] border border-border bg-background px-1.5 text-[10px] text-muted-foreground">
|
||||
function
|
||||
</div>
|
||||
|
||||
<Text strong style={{ fontSize: 13, display: "block", marginBottom: 6 }}>
|
||||
{tool.name}
|
||||
</Text>
|
||||
<span className="mb-1.5 block text-[13px] font-semibold">{tool.name}</span>
|
||||
|
||||
{Object.keys(tool.arguments).length > 0 && (
|
||||
<div>
|
||||
{Object.entries(tool.arguments).map(([key, value]) => (
|
||||
<div key={key} style={{ marginBottom: 2 }}>
|
||||
<Text type="secondary" style={{ fontSize: 12 }}>
|
||||
{key}:{" "}
|
||||
</Text>
|
||||
<Text style={{ fontSize: 12 }}>{JSON.stringify(value)}</Text>
|
||||
<div key={key} className="mb-0.5">
|
||||
<span className="text-xs text-muted-foreground">{key}: </span>
|
||||
<span className="text-xs">{JSON.stringify(value)}</span>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -0,0 +1,29 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, expect, it } from "vitest";
|
||||
import { TokenFlow } from "./TokenFlow";
|
||||
|
||||
const localised = (count: number) => count.toLocaleString();
|
||||
|
||||
describe("TokenFlow", () => {
|
||||
it("should render the total followed by its prompt and completion breakdown", () => {
|
||||
render(<TokenFlow prompt={9} completion={3} total={12} />);
|
||||
|
||||
expect(screen.getByText("12 (9 prompt tokens + 3 completion tokens)")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should group large counts the way the reader's locale does", () => {
|
||||
render(<TokenFlow prompt={1234567} completion={89012} total={1323579} />);
|
||||
|
||||
expect(
|
||||
screen.getByText(
|
||||
`${localised(1323579)} (${localised(1234567)} prompt tokens + ${localised(89012)} completion tokens)`,
|
||||
),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should fall back to zero for counts the log entry does not carry", () => {
|
||||
render(<TokenFlow total={12} />);
|
||||
|
||||
expect(screen.getByText("12 (0 prompt tokens + 0 completion tokens)")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,7 +1,3 @@
|
|||
import { Typography } from "antd";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface TokenFlowProps {
|
||||
prompt?: number;
|
||||
completion?: number;
|
||||
|
|
@ -14,9 +10,9 @@ interface TokenFlowProps {
|
|||
*/
|
||||
export function TokenFlow({ prompt = 0, completion = 0, total = 0 }: TokenFlowProps) {
|
||||
return (
|
||||
<Text>
|
||||
<span>
|
||||
{total.toLocaleString()} ({prompt.toLocaleString()} prompt tokens + {completion.toLocaleString()} completion
|
||||
tokens)
|
||||
</Text>
|
||||
</span>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
import { Typography, Tooltip } from "antd";
|
||||
import { DEFAULT_MAX_WIDTH, FONT_FAMILY_MONO, FONT_SIZE_SMALL } from "./constants";
|
||||
|
||||
const { Text } = Typography;
|
||||
import CopyButton from "@/components/shared/CopyButton";
|
||||
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
|
||||
import { DEFAULT_MAX_WIDTH, FONT_FAMILY_MONO } from "./constants";
|
||||
|
||||
interface TruncatedValueProps {
|
||||
value?: string;
|
||||
|
|
@ -13,23 +12,23 @@ interface TruncatedValueProps {
|
|||
* Useful for displaying long IDs, URLs, or other text that may overflow.
|
||||
*/
|
||||
export function TruncatedValue({ value, maxWidth = DEFAULT_MAX_WIDTH }: TruncatedValueProps) {
|
||||
if (!value) return <Text type="secondary">-</Text>;
|
||||
if (!value) return <span className="text-muted-foreground">-</span>;
|
||||
|
||||
return (
|
||||
<Tooltip title={value}>
|
||||
<Text
|
||||
copyable={{ text: value, tooltips: ["Copy", "Copied!"] }}
|
||||
style={{
|
||||
maxWidth,
|
||||
display: "inline-block",
|
||||
verticalAlign: "bottom",
|
||||
fontFamily: FONT_FAMILY_MONO,
|
||||
fontSize: FONT_SIZE_SMALL,
|
||||
}}
|
||||
ellipsis
|
||||
>
|
||||
{value}
|
||||
</Text>
|
||||
</Tooltip>
|
||||
<TooltipProvider delay={300}>
|
||||
<Tooltip>
|
||||
<TooltipTrigger
|
||||
render={
|
||||
<span className="inline-flex items-center gap-1 align-bottom">
|
||||
<span className="truncate text-xs" style={{ maxWidth, fontFamily: FONT_FAMILY_MONO }}>
|
||||
{value}
|
||||
</span>
|
||||
<CopyButton value={value} label="Copy" className="size-4 shrink-0" iconClassName="size-3" />
|
||||
</span>
|
||||
}
|
||||
/>
|
||||
<TooltipContent>{value}</TooltipContent>
|
||||
</Tooltip>
|
||||
</TooltipProvider>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
359
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
359
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -807,6 +807,93 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/auto_router/shadow_eval": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* List Shadow Eval Jobs
|
||||
* @description List shadow eval jobs, newest first. Counts and results ride the detail endpoint only.
|
||||
*/
|
||||
get: operations["list_shadow_eval_jobs_auto_router_shadow_eval_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/auto_router/shadow_eval/start": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/**
|
||||
* Start Shadow Eval
|
||||
* @description Start a pre-adoption shadow eval: duplicate a sampled slice of a key's live traffic
|
||||
* through an auto-router, judge real vs. shadow responses blind, and stratify win rates
|
||||
* by the router's tier classification and by the incumbent model.
|
||||
*
|
||||
* Shadow responses are never served to users. The job samples until it has judged
|
||||
* max_turns turns, reaches the end of its window, or is stopped; sampling changes
|
||||
* propagate to pods within about 10 seconds. Shadow and judge calls bill to the
|
||||
* shadowed key but are excluded from request counts and auto-router adoption metrics.
|
||||
*/
|
||||
post: operations["start_shadow_eval_auto_router_shadow_eval_start_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/auto_router/shadow_eval/{job_id}": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* Get Shadow Eval Job
|
||||
* @description One job with derived counts, judge spend, latest error, and stratified results.
|
||||
*/
|
||||
get: operations["get_shadow_eval_job_auto_router_shadow_eval__job_id__get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/auto_router/shadow_eval/{job_id}/stop": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/**
|
||||
* Stop Shadow Eval Job
|
||||
* @description Stop an active shadow eval job. Attempts are kept; sampling halts within ~10s.
|
||||
*/
|
||||
post: operations["stop_shadow_eval_job_auto_router_shadow_eval__job_id__stop_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/auto_router/test_routing": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -32557,6 +32644,110 @@ export interface components {
|
|||
/** Timeout */
|
||||
timeout?: number | null;
|
||||
};
|
||||
/**
|
||||
* ShadowEvalJobResponse
|
||||
* @description A shadow-eval job. Validates directly from the prisma record (job_id reads the
|
||||
* row's id); status is derived from stopped_at and ends_at, never stored, so no writer
|
||||
* anywhere can produce an inconsistent one. Aggregate fields are populated by the
|
||||
* detail endpoint only and stay None on list responses.
|
||||
*/
|
||||
ShadowEvalJobResponse: {
|
||||
/**
|
||||
* Api Key Id
|
||||
* @description The hashed virtual key whose traffic this job evaluates, and only that key's
|
||||
*/
|
||||
api_key_id: string;
|
||||
/**
|
||||
* Created At
|
||||
* Format: date-time
|
||||
*/
|
||||
created_at: string;
|
||||
/**
|
||||
* Ends At
|
||||
* Format: date-time
|
||||
*/
|
||||
ends_at: string;
|
||||
/**
|
||||
* Error Count
|
||||
* @description Sampled attempts that errored; detail endpoint only
|
||||
*/
|
||||
error_count?: number | null;
|
||||
/** Job Id */
|
||||
job_id: string;
|
||||
/** Judge Model */
|
||||
judge_model: string;
|
||||
/**
|
||||
* Judge Spend
|
||||
* @description Judge cost so far; detail endpoint only
|
||||
*/
|
||||
judge_spend?: number | null;
|
||||
/**
|
||||
* Judged Count
|
||||
* @description Verdicts recorded; detail endpoint only
|
||||
*/
|
||||
judged_count?: number | null;
|
||||
/**
|
||||
* Last Error
|
||||
* @description Most recent attempt error; detail endpoint only
|
||||
*/
|
||||
last_error?: string | null;
|
||||
/** Max Turns */
|
||||
max_turns: number;
|
||||
/** @description Stratified verdicts; detail endpoint only */
|
||||
results?: components["schemas"]["ShadowEvalResult"] | null;
|
||||
/** Router Name */
|
||||
router_name: string;
|
||||
/** Shadow Percentage */
|
||||
shadow_percentage: number;
|
||||
/**
|
||||
* Status
|
||||
* @description A job whose window has passed reads completed even if a later sweep stamped
|
||||
* stopped_at; stopped means sampling ended before the window did.
|
||||
* @enum {string}
|
||||
*/
|
||||
readonly status: "running" | "completed" | "stopped";
|
||||
/** Stopped At */
|
||||
stopped_at?: string | null;
|
||||
};
|
||||
/**
|
||||
* ShadowEvalResult
|
||||
* @description Stratified results of a shadow-eval job's verdicts so far.
|
||||
*/
|
||||
ShadowEvalResult: {
|
||||
/** By Current Model */
|
||||
by_current_model: components["schemas"]["ShadowEvalSlice"][];
|
||||
/** By Tier */
|
||||
by_tier: components["schemas"]["ShadowEvalSlice"][];
|
||||
/** Overall Shadow Win Rate Pct */
|
||||
overall_shadow_win_rate_pct: number;
|
||||
/** Overall Tie Rate Pct */
|
||||
overall_tie_rate_pct: number;
|
||||
};
|
||||
/**
|
||||
* ShadowEvalSlice
|
||||
* @description Judge outcomes for one slice of a job's verdicts (a router tier, or one of the
|
||||
* models the shadowed key currently uses).
|
||||
*/
|
||||
ShadowEvalSlice: {
|
||||
/** Avg Judge Confidence */
|
||||
avg_judge_confidence: number;
|
||||
/** Group */
|
||||
group: string;
|
||||
/**
|
||||
* Real Win Rate Pct
|
||||
* @description Share of judged turns where the real (control) model won
|
||||
*/
|
||||
real_win_rate_pct: number;
|
||||
/**
|
||||
* Shadow Win Rate Pct
|
||||
* @description Share of judged turns where the shadowed router's pick won
|
||||
*/
|
||||
shadow_win_rate_pct: number;
|
||||
/** Tie Rate Pct */
|
||||
tie_rate_pct: number;
|
||||
/** Turn Count */
|
||||
turn_count: number;
|
||||
};
|
||||
/**
|
||||
* Skill
|
||||
* @description Represents a skill from the Anthropic Skills API
|
||||
|
|
@ -32730,6 +32921,45 @@ export interface components {
|
|||
/** Simple Medium */
|
||||
simple_medium: number;
|
||||
};
|
||||
/**
|
||||
* StartShadowEvalRequest
|
||||
* @description Start shadowing a key's traffic through an auto-router for blind comparison.
|
||||
*/
|
||||
StartShadowEvalRequest: {
|
||||
/**
|
||||
* Api Key Id
|
||||
* @description The hashed virtual key whose traffic will be shadowed. Shadow evaluation runs ONLY on this key's traffic; requests made with any other key are not sampled.
|
||||
*/
|
||||
api_key_id: string;
|
||||
/**
|
||||
* Duration Days
|
||||
* @description How many days the job samples traffic before completing on its own
|
||||
* @default 7
|
||||
*/
|
||||
duration_days: number;
|
||||
/**
|
||||
* Judge Model
|
||||
* @description Model used to blindly judge real vs. shadow responses. The judge only compares two answers, so a mid-tier model (Claude Sonnet or GPT-4o class) is the sweet spot: small/nano-class models produce unreliable or malformed verdicts, while frontier reasoning models add cost without changing outcomes.
|
||||
* @default anthropic/claude-sonnet-5
|
||||
*/
|
||||
judge_model: string;
|
||||
/**
|
||||
* Max Turns
|
||||
* @description Sample budget: the job judges at most this many turns, then completes. This is also the spend bound; expected judge cost is roughly max_turns times one judge call
|
||||
* @default 200
|
||||
*/
|
||||
max_turns: number;
|
||||
/**
|
||||
* Router Name
|
||||
* @description The auto-router config to shadow requests through
|
||||
*/
|
||||
router_name: string;
|
||||
/**
|
||||
* Shadow Percentage
|
||||
* @description Percentage of the key's requests to duplicate through the router
|
||||
*/
|
||||
shadow_percentage: number;
|
||||
};
|
||||
/**
|
||||
* SuccessfulKeyUpdate
|
||||
* @description Successfully updated key with its updated information
|
||||
|
|
@ -37006,6 +37236,135 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
list_shadow_eval_jobs_auto_router_shadow_eval_get: {
|
||||
parameters: {
|
||||
query?: {
|
||||
/** @description Filter to jobs shadowing this key */
|
||||
api_key_id?: string | null;
|
||||
/** @description Newest jobs to return */
|
||||
limit?: number;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["ShadowEvalJobResponse"][];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
start_shadow_eval_auto_router_shadow_eval_start_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody: {
|
||||
content: {
|
||||
"application/json": components["schemas"]["StartShadowEvalRequest"];
|
||||
};
|
||||
};
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
201: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["ShadowEvalJobResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
get_shadow_eval_job_auto_router_shadow_eval__job_id__get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
job_id: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["ShadowEvalJobResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
stop_shadow_eval_job_auto_router_shadow_eval__job_id__stop_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
job_id: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["ShadowEvalJobResponse"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
preview_auto_router_routing_auto_router_test_routing_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue