mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/internal-users-tags-usage-e58133
This commit is contained in:
commit
21bb011770
111 changed files with 9253 additions and 1251 deletions
|
|
@ -6,10 +6,10 @@
|
|||
"limit": 2564
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 320
|
||||
"limit": 319
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"limit": 483
|
||||
"limit": 480
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 113
|
||||
|
|
@ -30,7 +30,7 @@
|
|||
"limit": 7
|
||||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"limit": 154
|
||||
"limit": 105
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"limit": 56
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44528
|
||||
"limit": 44526
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38804
|
||||
"limit": 38782
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19829
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30355
|
||||
"limit": 30349
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 117
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 833
|
||||
"limit": 831
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -59,9 +59,11 @@ if TYPE_CHECKING:
|
|||
from litellm.types.llms.openai import (
|
||||
ALL_RESPONSES_API_TOOL_PARAMS,
|
||||
AllMessageValues,
|
||||
ChatCompletionFileObject,
|
||||
ChatCompletionImageObject,
|
||||
ChatCompletionRedactedThinkingBlock,
|
||||
ChatCompletionThinkingBlock,
|
||||
ChatCompletionToolReferenceObject,
|
||||
OpenAIMessageContentListBlock,
|
||||
)
|
||||
from litellm.types.utils import Choices
|
||||
|
|
@ -175,6 +177,16 @@ def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Li
|
|||
return "length"
|
||||
|
||||
|
||||
def _input_file_from_file_value(file_value: object) -> dict[str, object]:
|
||||
if not isinstance(file_value, dict):
|
||||
return {"type": "input_file"}
|
||||
file_dict: Final = cast("dict[str, object]", file_value) # cast-ok: runtime dict checked
|
||||
return {
|
||||
"type": "input_file",
|
||||
**{key: file_dict[key] for key in ("file_id", "file_data", "filename") if key in file_dict},
|
||||
}
|
||||
|
||||
|
||||
def _incomplete_reason_from_response_payload(response_payload: object) -> str | None:
|
||||
if not isinstance(response_payload, Mapping):
|
||||
return None
|
||||
|
|
@ -957,7 +969,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
content: str
|
||||
| list[object]
|
||||
| Iterable[
|
||||
Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]
|
||||
Union[
|
||||
"OpenAIMessageContentListBlock",
|
||||
"ChatCompletionThinkingBlock",
|
||||
"ChatCompletionRedactedThinkingBlock",
|
||||
"ChatCompletionToolReferenceObject",
|
||||
]
|
||||
]
|
||||
| None,
|
||||
role: str,
|
||||
|
|
@ -1006,17 +1023,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
result.append(converted)
|
||||
verbose_logger.debug("Chat provider: image -> %s", converted)
|
||||
elif item_type == "file":
|
||||
# Map Chat Completion file to Responses API input_file
|
||||
# {"type": "file", "file": {"file_data": "...", "filename": "..."}}
|
||||
# -> {"type": "input_file", "file_data": "...", "filename": "..."}
|
||||
file_data = item.get("file", {})
|
||||
converted = {"type": "input_file"}
|
||||
if isinstance(file_data, dict):
|
||||
for key in ["file_id", "file_data", "filename"]:
|
||||
if key in file_data:
|
||||
converted[key] = file_data[key]
|
||||
converted = _input_file_from_file_value(
|
||||
cast("ChatCompletionFileObject", item).get("file"), # cast-ok: type tag checked
|
||||
)
|
||||
result.append(converted)
|
||||
verbose_logger.debug("Chat provider: file -> %s", converted)
|
||||
elif item_type == "tool_reference":
|
||||
verbose_logger.debug(
|
||||
"Chat provider: tool_reference has no responses API equivalent; skipped"
|
||||
)
|
||||
elif item_type in [
|
||||
"input_text",
|
||||
"input_image",
|
||||
|
|
|
|||
|
|
@ -76,7 +76,10 @@ from litellm.llms.perplexity.cost_calculator import (
|
|||
from litellm.llms.tencent.cost_calculator import (
|
||||
cost_per_token as tencent_cost_per_token,
|
||||
)
|
||||
from litellm.llms.together_ai.cost_calculator import get_model_params_and_category
|
||||
from litellm.llms.together_ai.cost_calculator import (
|
||||
get_model_params_and_category,
|
||||
has_together_registry_pricing,
|
||||
)
|
||||
from litellm.llms.vertex_ai.cost_calculator import (
|
||||
cost_per_character as google_cost_per_character,
|
||||
)
|
||||
|
|
@ -1569,10 +1572,9 @@ def completion_cost(
|
|||
|
||||
return MCPCostCalculator.calculate_mcp_tool_call_cost(litellm_logging_obj=litellm_logging_obj)
|
||||
# Calculate cost based on prompt_tokens, completion_tokens
|
||||
if "togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai":
|
||||
# together ai prices based on size of llm
|
||||
# get_model_params_and_category takes a model name and returns the category of LLM size it is in model_prices_and_context_window.json
|
||||
|
||||
if (
|
||||
"togethercomputer" in model or "together_ai" in model or custom_llm_provider == "together_ai"
|
||||
) and not has_together_registry_pricing(model, litellm.model_cost):
|
||||
model = get_model_params_and_category(model, call_type=CallTypes(call_type))
|
||||
|
||||
# replicate llms are calculate based on time for request running
|
||||
|
|
|
|||
|
|
@ -56,6 +56,9 @@ from litellm.types.mcp import (
|
|||
MCPStdioConfig,
|
||||
MCPTransport,
|
||||
MCPTransportType,
|
||||
credential_redirect_hook,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -273,6 +276,7 @@ class MCPClient:
|
|||
transport_type: MCPTransportType = MCPTransport.http,
|
||||
auth_type: MCPAuthType = None,
|
||||
auth_value: str | dict[str, str] | None = None,
|
||||
auth_header_name: str | None = None,
|
||||
timeout: float | None = None,
|
||||
stdio_config: MCPStdioConfig | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -288,6 +292,11 @@ class MCPClient:
|
|||
self.auth_type: MCPAuthType = auth_type
|
||||
self.timeout: float = timeout if timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
self._mcp_auth_value: str | dict[str, str] | None = None
|
||||
# The one place this client decides which header its credential occupies: the operator's
|
||||
# configured slot on the v1 path, or the slot the v2 resolver's auth object already owns.
|
||||
# Every consumer reads this rather than re-deriving it, since each re-derivation so far
|
||||
# picked up a different bug.
|
||||
self._credential_slot: str | None = auth_header_name or getattr(resolved_auth, "header_name", None)
|
||||
self.stdio_config: MCPStdioConfig | None = stdio_config
|
||||
self.extra_headers: dict[str, str] | None = extra_headers
|
||||
self.ssl_verify: VerifyTypes | None = ssl_verify
|
||||
|
|
@ -501,26 +510,33 @@ class MCPClient:
|
|||
else:
|
||||
self._mcp_auth_value = mcp_auth_value
|
||||
|
||||
def _header_slot(self, default: str) -> str:
|
||||
return self._credential_slot or default
|
||||
|
||||
def _get_auth_headers(self) -> dict:
|
||||
"""Generate authentication headers based on auth type."""
|
||||
headers: Final = {}
|
||||
if self._mcp_auth_value:
|
||||
if isinstance(self._mcp_auth_value, str):
|
||||
if self.auth_type == MCPAuth.bearer_token:
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
|
||||
static_bearer: Final = strip_auth_scheme(self._mcp_auth_value, "Bearer")
|
||||
headers[self._header_slot("Authorization")] = f"Bearer {static_bearer}"
|
||||
elif self.auth_type == MCPAuth.basic:
|
||||
headers["Authorization"] = f"Basic {self._mcp_auth_value}"
|
||||
headers[self._header_slot("Authorization")] = f"Basic {self._mcp_auth_value}"
|
||||
elif self.auth_type == MCPAuth.api_key:
|
||||
headers["X-API-Key"] = self._mcp_auth_value
|
||||
headers[self._header_slot("X-API-Key")] = self._mcp_auth_value
|
||||
elif self.auth_type == MCPAuth.authorization:
|
||||
# This auth type means the caller owns the whole header value.
|
||||
headers["Authorization"] = self._mcp_auth_value
|
||||
headers[self._header_slot("Authorization")] = self._mcp_auth_value
|
||||
elif self.auth_type == MCPAuth.oauth2:
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
|
||||
oauth2_bearer: Final = strip_auth_scheme(self._mcp_auth_value, "Bearer")
|
||||
headers[self._header_slot("Authorization")] = f"Bearer {oauth2_bearer}"
|
||||
elif self.auth_type == MCPAuth.token:
|
||||
headers["Authorization"] = f"token {strip_auth_scheme(self._mcp_auth_value, 'token')}"
|
||||
scheme_token: Final = strip_auth_scheme(self._mcp_auth_value, "token")
|
||||
headers[self._header_slot("Authorization")] = f"token {scheme_token}"
|
||||
elif self.auth_type == MCPAuth.oauth2_token_exchange:
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}"
|
||||
exchanged_bearer: Final = strip_auth_scheme(self._mcp_auth_value, "Bearer")
|
||||
headers[self._header_slot("Authorization")] = f"Bearer {exchanged_bearer}"
|
||||
elif isinstance(self._mcp_auth_value, dict):
|
||||
headers.update(self._mcp_auth_value)
|
||||
# Note: aws_sigv4 auth is not handled here — SigV4 requires per-request
|
||||
|
|
@ -528,7 +544,14 @@ class MCPClient:
|
|||
# of static headers. See MCPSigV4Auth and _create_httpx_client_factory().
|
||||
# update the headers with the extra headers
|
||||
if self.extra_headers:
|
||||
headers.update(self.extra_headers)
|
||||
# Mirrors _resolve_v2_auth: when the operator named a slot for the credential the
|
||||
# gateway resolved, no injected header may shadow it, case-insensitively, since HTTP
|
||||
# header names are. Without a configured slot the old precedence stands unchanged.
|
||||
slot: Final = self._credential_slot
|
||||
injected: Final = (
|
||||
without_header(self.extra_headers, slot) if slot and has_header(headers, slot) else self.extra_headers
|
||||
)
|
||||
headers.update(injected or {})
|
||||
return _strip_header_whitespace(headers)
|
||||
|
||||
def _create_httpx_client_factory(self) -> Callable[..., httpx.AsyncClient]:
|
||||
|
|
@ -556,12 +579,14 @@ class MCPClient:
|
|||
# SigV4 aws_auth. Both are None for the common case — no behavior change.
|
||||
fallback_auth: Final = self._resolved_auth if self._resolved_auth is not None else self._aws_auth
|
||||
effective_auth: Final = auth if auth is not None else fallback_auth
|
||||
guard: Final = credential_redirect_hook(self.server_url, self._credential_slot)
|
||||
return httpx.AsyncClient(
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
auth=effective_auth,
|
||||
verify=ssl_config,
|
||||
follow_redirects=True,
|
||||
event_hooks={"request": [guard]} if guard else {},
|
||||
)
|
||||
|
||||
return factory
|
||||
|
|
|
|||
|
|
@ -376,7 +376,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
# 2. list of objects - only apply to last item per Anthropic spec
|
||||
elif isinstance(message_content, list):
|
||||
if len(message_content) > 0 and isinstance(message_content[-1], dict):
|
||||
message_content[-1]["cache_control"] = control
|
||||
message_content[-1]["cache_control"] = control # pyright: ignore[reportGeneralTypeIssues] # loose runtime dict
|
||||
return message
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -146,7 +146,7 @@ class SpanEmitter:
|
|||
For callers that own and manage their own span lifecycle. ``tracer``
|
||||
overrides the bound tracer for this span only, used for per-request
|
||||
multi-tenant credential routing. ``links`` records related-but-not-parent
|
||||
spans (e.g. the transport span of an MCP message, per MCP semconv).
|
||||
spans (e.g. the trace context an MCP client propagated in ``params._meta``).
|
||||
"""
|
||||
return (tracer or self._tracer).start_span(
|
||||
name,
|
||||
|
|
@ -196,8 +196,8 @@ class SpanEmitter:
|
|||
|
||||
Return the span, or ``None`` if it was deduplicated away. ``tracer``
|
||||
overrides the bound tracer for this span, used for per-request routing.
|
||||
``links`` records related-but-not-parent spans (the transport span of an
|
||||
MCP message).
|
||||
``links`` records related-but-not-parent spans (e.g. the trace context an
|
||||
MCP client propagated in ``params._meta``).
|
||||
"""
|
||||
# LLM-call and MCP tool-call spans carry a dedup key (their request's
|
||||
# call id), so a sync+async double-firing coalesces. ``isinstance`` narrows
|
||||
|
|
|
|||
|
|
@ -390,10 +390,10 @@ class OpenTelemetryV2(CustomLogger):
|
|||
|
||||
MCP tool calls reach the success/failure callbacks like any other request
|
||||
(with ``call_type`` ``call_mcp_tool``), but they are not LLM calls and have
|
||||
no ``pre_call`` carrier — so they get their own CLIENT span here. Per the MCP
|
||||
semconv it parents to the trace context the client propagated in
|
||||
``params._meta`` (or starts a new root) and links the transport span, rather
|
||||
than nesting under the HTTP/session span. Returns whether it handled the
|
||||
no ``pre_call`` carrier — so they get their own CLIENT span here. It nests
|
||||
under the transport span of the request carrying this message, and trace
|
||||
context the client propagated in ``params._meta`` is recorded as a span
|
||||
link (see ``resolve_mcp_span_context``). Returns whether it handled the
|
||||
event, so the caller skips the LLM-call path. The whole span is emitted at
|
||||
once (there is no boundary to open it at), deduped on the call id.
|
||||
"""
|
||||
|
|
@ -436,9 +436,9 @@ class OpenTelemetryV2(CustomLogger):
|
|||
|
||||
Like a tool call, listing reaches the success/failure callbacks (here with
|
||||
``call_type`` ``list_mcp_tools``) with no ``pre_call`` carrier, so it gets its
|
||||
own CLIENT span. Per the MCP semconv it parents to the ``params._meta`` trace
|
||||
context (or starts a new root) and links the transport span, rather than
|
||||
nesting under the HTTP/session span. Returns whether it handled the event so
|
||||
own CLIENT span, nested under the transport span of the request carrying
|
||||
this message with any ``params._meta`` trace context recorded as a span
|
||||
link (see ``resolve_mcp_span_context``). Returns whether it handled the event so
|
||||
the caller skips the LLM-call path.
|
||||
"""
|
||||
raw_payload: Final = kwargs.get("standard_logging_object")
|
||||
|
|
|
|||
|
|
@ -10,6 +10,8 @@ Canonical hierarchy::
|
|||
│ └── DB_CALL (CLIENT) # its key/user/team lookups nest here
|
||||
├── GUARDRAIL (INTERNAL) # request-lifecycle hook, sibling of LLM_CALL
|
||||
├── LLM_CALL (CLIENT)
|
||||
├── MCP_TOOL_CALL (CLIENT) # nests under the POST carrying the message
|
||||
├── MCP_LIST_TOOLS (CLIENT) # (client-propagated context is a span link)
|
||||
└── DB_CALL (CLIENT) # e.g. the spend-log write
|
||||
|
||||
Guardrails parent to PROXY_REQUEST, not LLM_CALL: pre/during/post-call guardrail
|
||||
|
|
@ -18,14 +20,14 @@ before the LLM call even starts), so a guardrail is a sibling of the LLM call,
|
|||
not a child of it. The emitter parents every span to the ambient OTel context
|
||||
(the active server span), which matches this.
|
||||
|
||||
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) have two shapes, chosen at emit
|
||||
time by :func:`resolve_mcp_span_context`. When the client propagates trace context
|
||||
in ``params._meta`` MCP and the HTTP transport are independent contexts per the
|
||||
OTel GenAI MCP semconv, so the span parents to that propagated context and records
|
||||
the ``PROXY_REQUEST`` transport span as a span *link*, never a parent — the shape
|
||||
this registry's ``parent=None, links=PROXY_REQUEST`` entry encodes. When nothing is
|
||||
propagated (the common case) the span nests under the transport span of the request
|
||||
carrying that message, so the tool call stays in one trace.
|
||||
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) are parented at emit time by
|
||||
:func:`resolve_mcp_span_context`: they nest under the ``PROXY_REQUEST`` transport
|
||||
span of the request carrying that message, so the tool call stays in one trace.
|
||||
Trace context the client propagated in ``params._meta`` (SEP-414) is recorded as
|
||||
a span *link*, never the parent — a remote parent would root the span in a trace
|
||||
whose root never reaches the gateway's tracing backend. Links always target that
|
||||
remote client context, never a registry role, so ``SpanSpec`` declares no link
|
||||
field; the concrete transport parent is resolved per message at emit time.
|
||||
|
||||
Not every service call becomes a span — :func:`span_role_for_service` decides:
|
||||
|
||||
|
|
@ -85,25 +87,19 @@ class SpanSpec:
|
|||
role: SpanRole
|
||||
kind: LiteLLMSpanKind
|
||||
parent: SpanRole | None
|
||||
links: SpanRole | None = None
|
||||
|
||||
|
||||
SPAN_REGISTRY: Final[dict[SpanRole, SpanSpec]] = {
|
||||
SpanRole.PROXY_REQUEST: SpanSpec(SpanRole.PROXY_REQUEST, LiteLLMSpanKind.SERVER, parent=None),
|
||||
SpanRole.LLM_CALL: SpanSpec(SpanRole.LLM_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
|
||||
# The proxy is an MCP client to the upstream server, so MCP spans are CLIENT
|
||||
# spans. With trace context propagated in ``params._meta``, MCP and the HTTP
|
||||
# transport are independent contexts (OTel GenAI MCP semconv): the span parents
|
||||
# to the propagated context and records the PROXY_REQUEST transport span as a
|
||||
# span *link*, never a parent — the shape ``parent=None, links=PROXY_REQUEST``
|
||||
# encodes. With nothing propagated, ``resolve_mcp_span_context`` nests the span
|
||||
# under that message's transport span instead, keeping the call in one trace.
|
||||
SpanRole.MCP_TOOL_CALL: SpanSpec(
|
||||
SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=None, links=SpanRole.PROXY_REQUEST
|
||||
),
|
||||
SpanRole.MCP_LIST_TOOLS: SpanSpec(
|
||||
SpanRole.MCP_LIST_TOOLS, LiteLLMSpanKind.CLIENT, parent=None, links=SpanRole.PROXY_REQUEST
|
||||
),
|
||||
# spans. ``resolve_mcp_span_context`` nests them under the PROXY_REQUEST
|
||||
# transport span of the request carrying that message (resolved per message at
|
||||
# emit time), keeping the call in one trace. Trace context the client
|
||||
# propagated in ``params._meta`` becomes a span *link* to that remote context,
|
||||
# which is not a registry role, so ``SpanSpec`` has no link field.
|
||||
SpanRole.MCP_TOOL_CALL: SpanSpec(SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
|
||||
SpanRole.MCP_LIST_TOOLS: SpanSpec(SpanRole.MCP_LIST_TOOLS, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
|
||||
SpanRole.GUARDRAIL: SpanSpec(SpanRole.GUARDRAIL, LiteLLMSpanKind.INTERNAL, parent=SpanRole.PROXY_REQUEST),
|
||||
SpanRole.DB_CALL: SpanSpec(SpanRole.DB_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
|
||||
SpanRole.SERVICE: SpanSpec(SpanRole.SERVICE, LiteLLMSpanKind.INTERNAL, parent=SpanRole.PROXY_REQUEST),
|
||||
|
|
@ -209,8 +205,8 @@ def service_span_name(data: "ServiceSpanData") -> str:
|
|||
|
||||
|
||||
def root_roles() -> list[SpanRole]:
|
||||
"""Roles with no in-process parent. They start a new trace unless they adopt a
|
||||
remote parent (e.g. an MCP span joining the client's propagated context)."""
|
||||
"""Roles with no in-process parent, i.e. they start a new trace (only the
|
||||
instrumentor-owned ``PROXY_REQUEST`` server span today)."""
|
||||
return [role for role, spec in SPAN_REGISTRY.items() if spec.parent is None]
|
||||
|
||||
|
||||
|
|
@ -227,8 +223,6 @@ def validate_registry(
|
|||
raise ValueError(f"SPAN_REGISTRY[{role}] has mismatched role {spec.role}")
|
||||
if spec.parent is not None and spec.parent not in reg:
|
||||
raise ValueError(f"span role {role} declares unknown parent {spec.parent}")
|
||||
if spec.links is not None and spec.links not in reg:
|
||||
raise ValueError(f"span role {role} declares unknown link target {spec.links}")
|
||||
missing: Final = [role for role in SpanRole if role not in reg]
|
||||
if missing:
|
||||
raise ValueError(f"SPAN_REGISTRY is missing roles: {missing}")
|
||||
|
|
|
|||
|
|
@ -57,8 +57,8 @@ def request_root_span() -> "Span | None":
|
|||
|
||||
# The W3C trace-context carrier (``traceparent``/``tracestate``/``baggage``) the
|
||||
# MCP client propagated in the current request's ``params._meta``. The MCP gateway
|
||||
# sets it per message so the MCP span can parent to the client's span rather than
|
||||
# to the transport. A ``ContextVar`` because, like the root-span anchor, it must
|
||||
# sets it per message so the MCP span can record the client's span as a span
|
||||
# link. A ``ContextVar`` because, like the root-span anchor, it must
|
||||
# ride the request task and be readable by the inline success-logging callback.
|
||||
_mcp_message_trace_carrier: Final["ContextVar[Mapping[str, str] | None]"] = ContextVar(
|
||||
"litellm_otel_mcp_message_trace_carrier", default=None
|
||||
|
|
@ -148,10 +148,10 @@ def _mcp_transport_span_context() -> "SpanContext | None":
|
|||
|
||||
Prefers the transport the gateway published for this specific message; falls
|
||||
back to the ambient request anchor for paths that emit an MCP span on the
|
||||
request task itself (the REST MCP endpoints, the SDK). Parenting and linking
|
||||
only need the immutable context, and unlike ``mcp_message_transport_span`` they
|
||||
stay correct against a transport that has already finished, so this does not
|
||||
require the span to still be recording.
|
||||
request task itself (the REST MCP endpoints). Parenting needs only the
|
||||
immutable context, and unlike ``mcp_message_transport_span`` it stays correct
|
||||
against a transport that has already finished, so this does not require the
|
||||
span to still be recording.
|
||||
"""
|
||||
published: Final = _mcp_message_transport_span.get()
|
||||
if published is not None:
|
||||
|
|
@ -222,25 +222,31 @@ def resolve_mcp_span_context(
|
|||
) -> "tuple[Context, tuple[Link, ...]]":
|
||||
"""Parent context + links for an MCP message span.
|
||||
|
||||
The span always nests under the transport span of the request carrying this
|
||||
message, so a tool call and the ``POST`` that carried it stay in one trace.
|
||||
The transport comes from :func:`_mcp_transport_span_context`, which is the
|
||||
*current message's* POST rather than whatever request happened to open the
|
||||
session, so a long-lived session does not glue every message under its first
|
||||
request.
|
||||
|
||||
When the client propagates W3C trace context in the request's ``params._meta``
|
||||
(SEP-414), MCP and the underlying transport are independent lifecycles — one
|
||||
streamable-HTTP session multiplexes many messages, and the client's own span is
|
||||
the truthful parent. So, per the OTel GenAI MCP semconv:
|
||||
(SEP-414), that remote context is recorded as a span *link*, never the parent.
|
||||
The OTel GenAI MCP semconv prefers the inverse (remote parent, transport link),
|
||||
but the gateway's tracing backend only ever receives the gateway's half of such
|
||||
a trace: parenting into the client's trace id roots the span in a trace whose
|
||||
root span never reaches the backend, so the span is unreachable from the trace
|
||||
view and the transport transaction shows a dangling link (observed with
|
||||
clients that propagate synthetic trace ids). Anchoring to the gateway's own
|
||||
request and linking the client's context keeps every trace renderable while
|
||||
preserving the client-side correlation.
|
||||
|
||||
* parent to the trace context the client propagated (a *remote* parent), and
|
||||
* record the transport span as a *link*, never the parent.
|
||||
|
||||
Almost no client implements SEP-414 yet, so in practice nothing is propagated.
|
||||
Rooting the span there splits a single tool call into two disconnected traces
|
||||
joined only by a link, which is how it surfaces in APM: the ``POST`` transaction
|
||||
and the ``tools/call`` span share no trace. With no remote parent to honor,
|
||||
parent to the transport span of the request carrying this message instead, so
|
||||
the call stays in one trace; no link is added since the transport is now the
|
||||
real parent. The transport comes from :func:`_mcp_transport_span_context`, which
|
||||
is the *current message's* POST rather than whatever request happened to open
|
||||
the session, so a long-lived session does not glue every message under its
|
||||
first request. With neither a remote parent nor a transport the returned context
|
||||
carries no span and the span legitimately starts its own root trace.
|
||||
With no transport at all the span starts its own root trace, still carrying
|
||||
the link — the client context is only ever a link, so this event keeps one
|
||||
shape everywhere. Both returned contexts are built on an explicitly empty
|
||||
base, so ambient (stale session) state can never leak in, and the span
|
||||
inherits the transport's sampling decision exactly like every other
|
||||
request-level span — a client's sampled flag neither forces nor suppresses
|
||||
recording.
|
||||
|
||||
Only trace context (``traceparent``/``tracestate``) is extracted, never the
|
||||
client's W3C Baggage: ``params._meta`` is caller-controlled, and the otel
|
||||
|
|
@ -251,13 +257,12 @@ def resolve_mcp_span_context(
|
|||
never fall through to the ambient (stale session) span.
|
||||
"""
|
||||
source: Final = carrier if carrier is not None else _mcp_message_trace_carrier.get()
|
||||
parent: Final = _PROPAGATOR.extract(dict(source or {}), context=Context())
|
||||
propagated: Final = get_current_span(_PROPAGATOR.extract(dict(source or {}), context=Context()))
|
||||
links: Final = (Link(propagated.get_span_context()),) if is_recordable_span(propagated) else ()
|
||||
transport: Final = _mcp_transport_span_context()
|
||||
if is_recordable_span(get_current_span(parent)):
|
||||
return parent, (Link(transport),) if transport is not None else ()
|
||||
if transport is not None:
|
||||
return context_from_span(NonRecordingSpan(transport)), ()
|
||||
return parent, ()
|
||||
if transport is None:
|
||||
return Context(), links
|
||||
return context_from_span(NonRecordingSpan(transport), context=Context()), links
|
||||
|
||||
|
||||
def is_recordable_span(obj: object) -> bool:
|
||||
|
|
|
|||
193
litellm/litellm_core_utils/audio_utils/subtitle_utils.py
Normal file
193
litellm/litellm_core_utils/audio_utils/subtitle_utils.py
Normal file
|
|
@ -0,0 +1,193 @@
|
|||
"""Provider-agnostic SRT/WebVTT subtitle synthesis from timestamped transcription tokens."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import accumulate, chain
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
CUE_MAX_TOKENS: Final = 15
|
||||
CUE_MAX_DURATION_MS: Final = 5000
|
||||
|
||||
SRT_RESPONSE_FORMAT: Final = "srt"
|
||||
VTT_RESPONSE_FORMAT: Final = "vtt"
|
||||
SUBTITLE_RESPONSE_FORMATS: Final = frozenset((SRT_RESPONSE_FORMAT, VTT_RESPONSE_FORMAT))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubtitleToken:
|
||||
text: str
|
||||
start_ms: int | None = None
|
||||
end_ms: int | None = None
|
||||
speaker: str | int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubtitleCue:
|
||||
start_ms: int
|
||||
end_ms: int
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CueAccumulator:
|
||||
texts: tuple[str, ...] = ()
|
||||
start_ms: int | None = None
|
||||
end_ms: int | None = None
|
||||
speaker: str | int | None = None
|
||||
|
||||
|
||||
def _completed_cue(accumulator: _CueAccumulator) -> tuple[SubtitleCue, ...]:
|
||||
if not accumulator.texts or accumulator.start_ms is None:
|
||||
return ()
|
||||
text: Final = "".join(accumulator.texts).strip()
|
||||
if not text:
|
||||
return ()
|
||||
end_ms: Final = accumulator.end_ms if accumulator.end_ms is not None else accumulator.start_ms
|
||||
return (SubtitleCue(start_ms=accumulator.start_ms, end_ms=end_ms, text=text),)
|
||||
|
||||
|
||||
def _cue_break_reached(accumulator: _CueAccumulator, token: SubtitleToken) -> bool:
|
||||
if len(accumulator.texts) >= CUE_MAX_TOKENS:
|
||||
return True
|
||||
return (
|
||||
accumulator.start_ms is not None
|
||||
and token.start_ms is not None
|
||||
and token.start_ms - accumulator.start_ms >= CUE_MAX_DURATION_MS
|
||||
)
|
||||
|
||||
|
||||
_AbsorbStep = tuple[tuple[SubtitleCue, ...], _CueAccumulator]
|
||||
|
||||
|
||||
def _absorb_token(accumulator: _CueAccumulator, token: SubtitleToken) -> _AbsorbStep:
|
||||
if token.start_ms is None and accumulator.start_ms is None:
|
||||
return (), accumulator
|
||||
if token.speaker is not None and token.speaker != accumulator.speaker:
|
||||
return _completed_cue(accumulator), _CueAccumulator(
|
||||
texts=(token.text,),
|
||||
start_ms=token.start_ms,
|
||||
end_ms=token.end_ms,
|
||||
speaker=token.speaker,
|
||||
)
|
||||
if _cue_break_reached(accumulator, token):
|
||||
return _completed_cue(accumulator), _CueAccumulator(
|
||||
texts=(token.text,),
|
||||
start_ms=token.start_ms,
|
||||
end_ms=token.end_ms,
|
||||
speaker=accumulator.speaker,
|
||||
)
|
||||
return (), _CueAccumulator(
|
||||
texts=(*accumulator.texts, token.text),
|
||||
start_ms=accumulator.start_ms if accumulator.start_ms is not None else token.start_ms,
|
||||
end_ms=token.end_ms if token.end_ms is not None else accumulator.end_ms,
|
||||
speaker=accumulator.speaker,
|
||||
)
|
||||
|
||||
|
||||
def _absorb_step(carry: _AbsorbStep, token: SubtitleToken) -> _AbsorbStep:
|
||||
return _absorb_token(carry[1], token)
|
||||
|
||||
|
||||
def group_subtitle_tokens_into_cues(tokens: Sequence[SubtitleToken]) -> tuple[SubtitleCue, ...]:
|
||||
steps: Final = tuple(accumulate(tokens, _absorb_step, initial=((), _CueAccumulator())))
|
||||
completed: Final = chain.from_iterable(emitted for emitted, _ in steps)
|
||||
return (*completed, *_completed_cue(steps[-1][1]))
|
||||
|
||||
|
||||
def _format_timestamp(total_ms: int, millis_separator: str) -> str:
|
||||
clamped: Final = max(total_ms, 0)
|
||||
hours, hour_remainder = divmod(clamped, 3_600_000)
|
||||
minutes, minute_remainder = divmod(hour_remainder, 60_000)
|
||||
seconds, millis = divmod(minute_remainder, 1_000)
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d}{millis_separator}{millis:03d}"
|
||||
|
||||
|
||||
def _render_srt(cues: Sequence[SubtitleCue]) -> str:
|
||||
lines: Final = tuple(
|
||||
line
|
||||
for index, cue in enumerate(cues, start=1)
|
||||
for line in (
|
||||
str(index),
|
||||
f"{_format_timestamp(cue.start_ms, ',')} --> {_format_timestamp(cue.end_ms, ',')}",
|
||||
cue.text,
|
||||
"",
|
||||
)
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _render_vtt(cues: Sequence[SubtitleCue]) -> str:
|
||||
cue_lines: Final = tuple(
|
||||
line
|
||||
for cue in cues
|
||||
for line in (
|
||||
f"{_format_timestamp(cue.start_ms, '.')} --> {_format_timestamp(cue.end_ms, '.')}",
|
||||
cue.text,
|
||||
"",
|
||||
)
|
||||
)
|
||||
return "\n".join(("WEBVTT", "", *cue_lines))
|
||||
|
||||
|
||||
def render_subtitle_tokens_as_srt(tokens: Sequence[SubtitleToken]) -> str:
|
||||
"""Render tokens as an SRT document; empty string when no token has timestamp data."""
|
||||
cues: Final = group_subtitle_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return ""
|
||||
return _render_srt(cues)
|
||||
|
||||
|
||||
def render_subtitle_tokens_as_vtt(tokens: Sequence[SubtitleToken]) -> str:
|
||||
"""Render tokens as a WebVTT document; the WEBVTT header is emitted even without cues."""
|
||||
return _render_vtt(group_subtitle_tokens_into_cues(tokens))
|
||||
|
||||
|
||||
class TranscriptionWordTiming(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
word: str = ""
|
||||
start: float | None = None
|
||||
end: float | None = None
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
_WORD_TIMINGS_ADAPTER: Final = TypeAdapter(tuple[TranscriptionWordTiming, ...])
|
||||
|
||||
|
||||
def _seconds_to_ms(seconds: float | None) -> int | None:
|
||||
if seconds is None:
|
||||
return None
|
||||
return round(seconds * 1000)
|
||||
|
||||
|
||||
def _word_to_subtitle_token(word: TranscriptionWordTiming) -> SubtitleToken:
|
||||
return SubtitleToken(
|
||||
text=f"{word.word} ",
|
||||
start_ms=_seconds_to_ms(word.start),
|
||||
end_ms=_seconds_to_ms(word.end),
|
||||
speaker=word.speaker,
|
||||
)
|
||||
|
||||
|
||||
def _parse_word_timings(words: object) -> tuple[TranscriptionWordTiming, ...]:
|
||||
try:
|
||||
return _WORD_TIMINGS_ADAPTER.validate_python(words)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def synthesize_subtitle_document(words: object, response_format: str) -> str | None:
|
||||
"""
|
||||
Build an SRT/VTT document from OpenAI verbose_json-style word dicts
|
||||
(word/start/end in float seconds, optional speaker). Returns None when the
|
||||
format is not a subtitle format or the words carry no usable timestamps.
|
||||
"""
|
||||
if response_format not in SUBTITLE_RESPONSE_FORMATS:
|
||||
return None
|
||||
tokens: Final = tuple(_word_to_subtitle_token(word) for word in _parse_word_timings(words))
|
||||
cues: Final = group_subtitle_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return None
|
||||
return _render_srt(cues) if response_format == SRT_RESPONSE_FORMAT else _render_vtt(cues)
|
||||
|
|
@ -1747,6 +1747,46 @@ def hoist_images_from_tool_messages(
|
|||
]
|
||||
|
||||
|
||||
def _is_tool_reference_part(part: object) -> bool:
|
||||
return isinstance(part, dict) and part.get("type") == "tool_reference"
|
||||
|
||||
|
||||
def _tool_message_carries_tool_reference(message: AllMessageValues) -> bool:
|
||||
if message.get("role") != "tool":
|
||||
return False
|
||||
content = message.get("content")
|
||||
return isinstance(content, list) and any(_is_tool_reference_part(part) for part in content)
|
||||
|
||||
|
||||
def _drop_tool_reference_parts(message: AllMessageValues) -> AllMessageValues:
|
||||
if not _tool_message_carries_tool_reference(message):
|
||||
return message
|
||||
content = cast(list, message.get("content")) # cast-ok: shape checked by _tool_message_carries_tool_reference
|
||||
remaining_parts = [ # mutable-ok: tool message content must stay a json list
|
||||
part for part in content if not _is_tool_reference_part(part)
|
||||
]
|
||||
new_content = remaining_parts if remaining_parts else ""
|
||||
rewritten = {**message, "content": new_content} # mutable-ok: chat messages are plain json dicts
|
||||
return cast(AllMessageValues, rewritten) # cast-ok: dict spread keeps keys like cache_control
|
||||
|
||||
|
||||
def drop_tool_reference_parts_from_tool_messages(
|
||||
messages: list[AllMessageValues], # mutable-ok: message pipelines type messages as mutable lists
|
||||
) -> list[AllMessageValues]: # mutable-ok: message pipelines type messages as mutable lists
|
||||
"""
|
||||
Remove tool_reference content parts from role:"tool" messages.
|
||||
|
||||
The OpenAI chat spec only accepts text in tool messages, so a tool_reference
|
||||
part carried through the Anthropic adapter makes strict providers reject the
|
||||
request. The reference names an already-declared tool rather than carrying
|
||||
content, so it is dropped; a reference-only result keeps its tool message with
|
||||
empty text so the preceding tool_call stays answered.
|
||||
"""
|
||||
if not any(_tool_message_carries_tool_reference(message) for message in messages):
|
||||
return messages
|
||||
return [_drop_tool_reference_parts(message) for message in messages] # mutable-ok: pipelines mutate message lists
|
||||
|
||||
|
||||
def _attempt_json_repair(s: str) -> Any | None:
|
||||
"""
|
||||
Attempt to repair truncated JSON produced by LLM tool calls.
|
||||
|
|
|
|||
|
|
@ -1412,7 +1412,7 @@ def convert_to_gemini_tool_call_result(
|
|||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to process image in tool response: %s", e)
|
||||
elif content_type in ("file", "input_file"):
|
||||
elif content_type in ("file", "input_file"): # pyright: ignore[reportUnnecessaryContains] # loose runtime dict
|
||||
# Extract file for inline_data (for tool results with PDF, audio, video, etc.)
|
||||
file_data = content.get("file_data", "")
|
||||
if not file_data:
|
||||
|
|
@ -1564,14 +1564,23 @@ def convert_to_anthropic_tool_result(
|
|||
}
|
||||
"""
|
||||
anthropic_content: (
|
||||
str | list[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam]
|
||||
str
|
||||
| list[
|
||||
AnthropicMessagesToolResultContent
|
||||
| AnthropicMessagesImageParam
|
||||
| AnthropicMessagesDocumentParam
|
||||
| ToolReference
|
||||
]
|
||||
) = ""
|
||||
if isinstance(message["content"], str):
|
||||
anthropic_content = message["content"]
|
||||
elif isinstance(message["content"], list):
|
||||
content_list: Final = message["content"]
|
||||
anthropic_content_list: list[
|
||||
AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam
|
||||
AnthropicMessagesToolResultContent
|
||||
| AnthropicMessagesImageParam
|
||||
| AnthropicMessagesDocumentParam
|
||||
| ToolReference
|
||||
] = []
|
||||
for content in content_list:
|
||||
if content["type"] == "text":
|
||||
|
|
@ -1614,6 +1623,8 @@ def convert_to_anthropic_tool_result(
|
|||
original_content_element=content,
|
||||
)
|
||||
anthropic_content_list.append(cast(AnthropicMessagesImageParam, _anthropic_image_param))
|
||||
elif content["type"] == "tool_reference":
|
||||
anthropic_content_list.append(ToolReference(type="tool_reference", tool_name=content["tool_name"]))
|
||||
elif content["type"] == "file":
|
||||
file_content = cast(ChatCompletionFileObject, content)
|
||||
_file_block = anthropic_process_openai_file_message(file_content)
|
||||
|
|
|
|||
|
|
@ -8,11 +8,12 @@ import time
|
|||
import traceback
|
||||
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, NoReturn, Protocol, TypeVar, cast
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
import litellm
|
||||
|
|
@ -182,6 +183,23 @@ class _VertexChunkLike(Protocol):
|
|||
candidates: Sequence[_VertexCandidateLike]
|
||||
|
||||
|
||||
class _ParsedChunkHiddenParams(BaseModel):
|
||||
provider_specific_fields: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
def _provider_hidden_params(chunk: object) -> Mapping[str, object] | None:
|
||||
hidden: Final[object] = getattr(chunk, "_hidden_params", None)
|
||||
if not isinstance(hidden, dict):
|
||||
return None
|
||||
try:
|
||||
parsed: Final = _ParsedChunkHiddenParams.model_validate(hidden)
|
||||
except ValidationError:
|
||||
return None
|
||||
if not parsed.provider_specific_fields:
|
||||
return None
|
||||
return MappingProxyType({"provider_specific_fields": dict(parsed.provider_specific_fields)})
|
||||
|
||||
|
||||
class CustomStreamWrapper:
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -801,7 +819,7 @@ class CustomStreamWrapper:
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
def model_response_creator(self, chunk: dict | None = None, hidden_params: dict | None = None):
|
||||
def model_response_creator(self, chunk: dict | None = None, hidden_params: Mapping[str, object] | None = None):
|
||||
_model: Final = self._cached_model_name
|
||||
_logging_obj_llm_provider: Final = self._cached_logging_llm_provider
|
||||
|
||||
|
|
@ -1504,7 +1522,7 @@ class CustomStreamWrapper:
|
|||
def chunk_creator(self, chunk: Any):
|
||||
if hasattr(chunk, "id"):
|
||||
self.response_id = chunk.id
|
||||
model_response = self.model_response_creator()
|
||||
model_response = self.model_response_creator(hidden_params=_provider_hidden_params(chunk))
|
||||
response_obj: dict[str, Any] = {}
|
||||
try:
|
||||
# return this for all models
|
||||
|
|
|
|||
|
|
@ -24,10 +24,12 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
is_provider_native_tool_dict,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
anthropic_tool_name,
|
||||
anthropic_tool_names,
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
|
|
@ -360,7 +362,13 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
|
||||
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = (
|
||||
[] if scan_only_tool_results else chat_completion_compatible_request.get("tools", [])
|
||||
[]
|
||||
if scan_only_tool_results
|
||||
else [
|
||||
tool
|
||||
for tool in chat_completion_compatible_request.get("tools", [])
|
||||
if not is_provider_native_tool_dict(tool)
|
||||
]
|
||||
)
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
|
|
@ -419,7 +427,10 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
tool_name=anthropic_tool_name,
|
||||
)
|
||||
if scan_only_tool_results
|
||||
else anthropic_tools
|
||||
else [
|
||||
*(tool for tool in data.get("tools") or [] if is_provider_native_tool_dict(tool)),
|
||||
*anthropic_tools,
|
||||
]
|
||||
)
|
||||
|
||||
guardrailed_structured_messages: Final = guardrailed_inputs.get("structured_messages")
|
||||
|
|
@ -677,12 +688,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
)
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
"""Extract tool names from Anthropic messages request (tools[].name)."""
|
||||
names: Final[list[str]] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if isinstance(tool, dict) and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
return names
|
||||
"""Extract every tool name in an Anthropic messages request: tools[].name, plus
|
||||
tools[].function.name for OpenAI-format tools the bridge forwards verbatim."""
|
||||
return [name for tool in data.get("tools") or [] for name in anthropic_tool_names(tool)]
|
||||
|
||||
@classmethod
|
||||
def _extract_input_text_and_images(
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypeVar, cast
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
|
|
@ -18,6 +18,22 @@ TOOL_NAME_PREFIX_LENGTH: Final = OPENAI_MAX_TOOL_NAME_LENGTH - TOOL_NAME_HASH_LE
|
|||
PROVIDERS_PROXYING_AN_UNKNOWN_BACKEND: Final = frozenset({"litellm_proxy"})
|
||||
|
||||
|
||||
_ANTHROPIC_TOOL_SCHEMA_KEYS: Final = frozenset(
|
||||
{"name", "type", "input_schema", "description", "cache_control", "strict"}
|
||||
)
|
||||
|
||||
|
||||
def _is_openai_function_tool(tool: Mapping[str, object]) -> bool:
|
||||
return tool.get("type") == "function" and "function" in tool
|
||||
|
||||
|
||||
def is_provider_native_tool_dict(tool: Mapping[str, object]) -> bool:
|
||||
if len(tool) != 1:
|
||||
return False
|
||||
key, value = next(iter(tool.items()))
|
||||
return key not in _ANTHROPIC_TOOL_SCHEMA_KEYS and isinstance(value, dict)
|
||||
|
||||
|
||||
def truncate_tool_name(name: str) -> str:
|
||||
"""
|
||||
Truncate tool names that exceed OpenAI's 64-character limit.
|
||||
|
|
@ -126,7 +142,9 @@ from litellm.types.llms.openai import (
|
|||
ChatCompletionToolMessage,
|
||||
ChatCompletionToolParam,
|
||||
ChatCompletionToolParamFunctionChunk,
|
||||
ChatCompletionToolReferenceObject,
|
||||
ChatCompletionUserMessage,
|
||||
ToolMessageContentPart,
|
||||
)
|
||||
from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage
|
||||
|
||||
|
|
@ -135,6 +153,8 @@ from .streaming_iterator import AnthropicStreamWrapper
|
|||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.anthropic import ContentBlockContentBlockDict
|
||||
|
||||
ToolResultContent: TypeAlias = str | list[ToolMessageContentPart]
|
||||
|
||||
|
||||
class AnthropicAdapter:
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -412,90 +432,13 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
self._add_cache_control_if_applicable(content, doc_obj, model)
|
||||
new_user_content_list.append(doc_obj)
|
||||
elif content.get("type") == "tool_result":
|
||||
if "content" not in content:
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content="",
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif isinstance(content.get("content"), str):
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=str(content.get("content", "")),
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif isinstance(content.get("content"), list):
|
||||
# Combine all content items into a single tool message
|
||||
# to avoid creating multiple tool_result blocks with the same ID
|
||||
# (each tool_use must have exactly one tool_result)
|
||||
content_items = list(content.get("content", []))
|
||||
|
||||
# Single-item text keeps the backward-compatible string format; a single
|
||||
# image or document becomes a structured image_url part
|
||||
if len(content_items) == 1:
|
||||
c = content_items[0]
|
||||
if isinstance(c, str):
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=c,
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif isinstance(c, dict):
|
||||
if c.get("type") == "text":
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=c.get("text", ""),
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
elif c.get("type") in ("image", "document"):
|
||||
image_part = self._tool_result_image_part(c.get("source"))
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=[image_part] # mutable-ok: content must be a json list
|
||||
if image_part
|
||||
else "",
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
else:
|
||||
# For multiple content items, combine into a single tool message
|
||||
# with list content to preserve all items while having one tool_use_id
|
||||
combined_content_parts: list[
|
||||
ChatCompletionTextObject | ChatCompletionImageObject
|
||||
] = []
|
||||
for c in content_items:
|
||||
if isinstance(c, str):
|
||||
combined_content_parts.append(ChatCompletionTextObject(type="text", text=c))
|
||||
elif isinstance(c, dict):
|
||||
if c.get("type") == "text":
|
||||
combined_content_parts.append(
|
||||
ChatCompletionTextObject(
|
||||
type="text",
|
||||
text=c.get("text", ""),
|
||||
)
|
||||
)
|
||||
elif c.get("type") in ("image", "document"):
|
||||
image_part = self._tool_result_image_part(c.get("source"))
|
||||
if image_part:
|
||||
combined_content_parts.append(image_part)
|
||||
# Create a single tool message with combined content
|
||||
if combined_content_parts:
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=combined_content_parts,
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
tool_result = ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id=content.get("tool_use_id", ""),
|
||||
content=self._tool_result_content(content.get("content")),
|
||||
)
|
||||
self._add_cache_control_if_applicable(content, tool_result, model)
|
||||
tool_message_list.append(tool_result)
|
||||
|
||||
if len(tool_message_list) > 0:
|
||||
new_messages.extend(tool_message_list)
|
||||
|
|
@ -771,6 +714,10 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
new_tools.append(tool)
|
||||
continue
|
||||
|
||||
if _is_openai_function_tool(tool) or is_provider_native_tool_dict(tool):
|
||||
new_tools.append(cast(ChatCompletionToolParam, tool)) # cast-ok: passed through verbatim to provider
|
||||
continue
|
||||
|
||||
raw_name = tool.get("name")
|
||||
if raw_name is None or (isinstance(raw_name, str) and not str(raw_name).strip()):
|
||||
original_name = f"litellm_unnamed_tool_{idx}"
|
||||
|
|
@ -1210,6 +1157,39 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
|
||||
return None
|
||||
|
||||
def _tool_result_content(self, raw_content: object) -> ToolResultContent:
|
||||
if isinstance(raw_content, str):
|
||||
return raw_content
|
||||
if not isinstance(raw_content, list):
|
||||
return ""
|
||||
items: Final = cast(Sequence[object], raw_content) # cast-ok: untrusted client payload
|
||||
parts: Final = tuple(part for part in (self._tool_result_part(item) for item in items) if part is not None)
|
||||
match parts:
|
||||
case ():
|
||||
return ""
|
||||
case ({"type": "text", "text": str(text)},):
|
||||
return text
|
||||
case _:
|
||||
return list(parts) # mutable-ok: content must be a json list
|
||||
|
||||
def _tool_result_part(self, item: object) -> ToolMessageContentPart | None:
|
||||
if isinstance(item, str):
|
||||
return ChatCompletionTextObject(type="text", text=item)
|
||||
if not isinstance(item, dict):
|
||||
return None
|
||||
block: Final = cast(Mapping[str, object], item) # cast-ok: untrusted client payload
|
||||
match block.get("type"):
|
||||
case "text":
|
||||
return ChatCompletionTextObject(type="text", text=str(block.get("text") or ""))
|
||||
case "image" | "document":
|
||||
return self._tool_result_image_part(block.get("source"))
|
||||
case "tool_reference":
|
||||
return ChatCompletionToolReferenceObject(
|
||||
type="tool_reference", tool_name=str(block.get("tool_name") or "")
|
||||
)
|
||||
case _:
|
||||
return None
|
||||
|
||||
def _tool_result_image_part(self, image_source: object) -> ChatCompletionImageObject | None:
|
||||
if not isinstance(image_source, dict):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from httpx._models import Headers, Response
|
|||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
hoist_images_from_tool_messages,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
|
|
@ -252,7 +253,8 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
azure_messages: Final = convert_to_azure_openai_messages(hoist_images_from_tool_messages(messages))
|
||||
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
azure_messages: Final = convert_to_azure_openai_messages(hoist_images_from_tool_messages(stripped_messages))
|
||||
return {
|
||||
"model": model,
|
||||
"messages": azure_messages,
|
||||
|
|
|
|||
|
|
@ -40,6 +40,16 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC):
|
|||
def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]:
|
||||
pass
|
||||
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
"""
|
||||
Opt-in for providers without a native srt/vtt response body: when True
|
||||
and the user asked for response_format srt/vtt, the http handler
|
||||
synthesizes the subtitle document from the word timestamps the
|
||||
provider's TranscriptionResponse carries in `words`.
|
||||
"""
|
||||
return False
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
|
|
|
|||
|
|
@ -209,9 +209,20 @@ def openai_tool_name(tool: object) -> str | None:
|
|||
return flat_name if isinstance(flat_name, str) else None
|
||||
|
||||
|
||||
def anthropic_tool_names(tool: object) -> tuple[str, ...]:
|
||||
"""Every name a /v1/messages tool dict can act under: the flat Anthropic ``name`` plus
|
||||
``function.name`` for OpenAI-format tools the bridge forwards verbatim. Allowlist checks
|
||||
must see both, or a decoy flat name could smuggle a disallowed ``function.name`` through."""
|
||||
if not isinstance(tool, dict):
|
||||
return ()
|
||||
function: Final = tool.get("function") if tool.get("type") == "function" else None
|
||||
function_name: Final = function.get("name") if isinstance(function, dict) else None
|
||||
return tuple(name for name in (tool.get("name"), function_name) if isinstance(name, str) and name)
|
||||
|
||||
|
||||
def anthropic_tool_name(tool: object) -> str | None:
|
||||
name: Final = tool.get("name") if isinstance(tool, dict) else None
|
||||
return name if isinstance(name, str) else None
|
||||
names: Final = anthropic_tool_names(tool)
|
||||
return names[0] if names else None
|
||||
|
||||
|
||||
def merge_returned_tools_into_request_tools(
|
||||
|
|
|
|||
|
|
@ -25,6 +25,10 @@ from litellm.litellm_core_utils.agentic_loop_settings import (
|
|||
validated_max_agentic_loops,
|
||||
)
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SUBTITLE_RESPONSE_FORMATS,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
|
|
@ -1296,9 +1300,23 @@ class BaseLLMHTTPHandler:
|
|||
api_key: str | None,
|
||||
) -> TranscriptionResponse:
|
||||
"""Shared logic for transforming audio transcription responses."""
|
||||
return provider_config.transform_audio_transcription_response(
|
||||
transformed: Final = provider_config.transform_audio_transcription_response(
|
||||
raw_response=response,
|
||||
)
|
||||
if not provider_config.supports_subtitle_synthesis:
|
||||
return transformed
|
||||
requested_format: Final = optional_params.get("response_format")
|
||||
if not isinstance(requested_format, str) or requested_format not in SUBTITLE_RESPONSE_FORMATS:
|
||||
return transformed
|
||||
document: Final = synthesize_subtitle_document(
|
||||
words=transformed.get("words"),
|
||||
response_format=requested_format,
|
||||
)
|
||||
if document is not None:
|
||||
transformed.text = document
|
||||
if "words" in transformed:
|
||||
delattr(transformed, "words")
|
||||
return transformed
|
||||
|
||||
def audio_transcriptions(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ Request format:
|
|||
"input": {
|
||||
"messages": [{"role": "user", "content": [{"text": "<prompt>"}]}]
|
||||
},
|
||||
"parameters": {"size": "1024*1024", ...}
|
||||
"parameters": {"size": "1024*1024", "n": 1, ...}
|
||||
}
|
||||
|
||||
Response format:
|
||||
|
|
@ -19,7 +19,7 @@ Response format:
|
|||
"output": {
|
||||
"choices": [{"message": {"content": [{"image": "<url>"}]}}]
|
||||
},
|
||||
"usage": {"input_tokens": 0, "output_tokens": 0, "width": 1024, "height": 1024, "image_count": 1}
|
||||
"usage": {"output_width": 1024, "output_height": 1024, "output_image_count": 1}
|
||||
}
|
||||
"""
|
||||
|
||||
|
|
@ -46,6 +46,8 @@ else:
|
|||
|
||||
DEFAULT_API_BASE: Final = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"
|
||||
|
||||
CHAT_COMPATIBLE_MODE_PATH: Final = "/compatible-mode/v1"
|
||||
|
||||
# Maps OpenAI size strings (WxH) to DashScope size strings (W*H)
|
||||
OPENAI_TO_DASHSCOPE_SIZE: Final[dict] = {
|
||||
"256x256": "256*256",
|
||||
|
|
@ -59,7 +61,8 @@ OPENAI_TO_DASHSCOPE_SIZE: Final[dict] = {
|
|||
|
||||
class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""
|
||||
Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro).
|
||||
Configuration for DashScope image generation (qwen-image-2.0, qwen-image-2.0-pro,
|
||||
qwen-image-3.0, qwen-image-3.0-pro).
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[OpenAIImageGenerationOptionalParams]:
|
||||
|
|
@ -82,8 +85,8 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
|||
if k == "size":
|
||||
# Convert "WxH" → "W*H"
|
||||
mapped["size"] = OPENAI_TO_DASHSCOPE_SIZE.get(v, v.replace("x", "*"))
|
||||
elif k == "n":
|
||||
mapped["image_count"] = v
|
||||
else:
|
||||
mapped[k] = v
|
||||
return mapped
|
||||
|
||||
def get_complete_url(
|
||||
|
|
@ -95,7 +98,10 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
|||
litellm_params: dict,
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
return api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE
|
||||
image_api_base: Final = (
|
||||
api_base if api_base and not api_base.rstrip("/").endswith(CHAT_COMPATIBLE_MODE_PATH) else None
|
||||
)
|
||||
return image_api_base or get_secret_str("DASHSCOPE_API_BASE_IMAGE") or DEFAULT_API_BASE
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Final
|
|||
|
||||
from httpx import Headers, Response
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import SUBTITLE_RESPONSE_FORMATS
|
||||
from litellm.litellm_core_utils.audio_utils.utils import (
|
||||
normalize_transcription_language_to_bcp47,
|
||||
process_audio_file,
|
||||
|
|
@ -48,6 +49,10 @@ class GeminiAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
|
|||
) -> list[OpenAIAudioTranscriptionOptionalParams]: # mutable-ok: BaseAudioTranscriptionConfig signature
|
||||
return ["language", "response_format", "timestamp_granularities"] # mutable-ok: base contract returns a list
|
||||
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
return True
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
|
|
@ -215,16 +220,17 @@ def _language_config(language: object) -> GeminiTranscriptionConfig:
|
|||
return language_config
|
||||
|
||||
|
||||
def _timestamp_config(timestamp_granularities: object) -> GeminiTranscriptionConfig:
|
||||
if isinstance(timestamp_granularities, list) and "word" in timestamp_granularities:
|
||||
return _WORD_TIMESTAMP_CONFIG
|
||||
return _EMPTY_TRANSCRIPTION_CONFIG
|
||||
def _timestamp_config(timestamp_granularities: object, response_format: object) -> GeminiTranscriptionConfig:
|
||||
wants_word_timestamps: Final = (
|
||||
isinstance(timestamp_granularities, list) and "word" in timestamp_granularities
|
||||
) or (isinstance(response_format, str) and response_format in SUBTITLE_RESPONSE_FORMATS)
|
||||
return _WORD_TIMESTAMP_CONFIG if wants_word_timestamps else _EMPTY_TRANSCRIPTION_CONFIG
|
||||
|
||||
|
||||
def _build_transcription_config(optional_params: Mapping[str, object]) -> GeminiTranscriptionConfig:
|
||||
transcription_config: Final[GeminiTranscriptionConfig] = {
|
||||
**_language_config(optional_params.get("language")),
|
||||
**_timestamp_config(optional_params.get("timestamp_granularities")),
|
||||
**_timestamp_config(optional_params.get("timestamp_granularities"), optional_params.get("response_format")),
|
||||
}
|
||||
return transcription_config
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject, ChatCompletionImageObject
|
||||
from litellm.types.llms.vertex_ai import ContentType, PartType
|
||||
from litellm.utils import supports_reasoning
|
||||
|
||||
|
|
@ -16,6 +16,13 @@ from ...vertex_ai.gemini.transformation import _gemini_convert_messages_with_his
|
|||
from ...vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
|
||||
|
||||
def _image_url_fields(img_element: ChatCompletionImageObject) -> tuple[str | None, str | None, str | None]:
|
||||
image_value: Final = img_element.get("image_url")
|
||||
if isinstance(image_value, dict):
|
||||
return image_value.get("url"), image_value.get("format"), image_value.get("detail")
|
||||
return image_value, None, None
|
||||
|
||||
|
||||
class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
|
||||
"""
|
||||
Reference: https://ai.google.dev/api/rest/v1beta/GenerationConfig
|
||||
|
|
@ -118,16 +125,8 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
|
|||
_parts: list[PartType] = []
|
||||
for element in _message_content:
|
||||
if element.get("type") == "image_url":
|
||||
img_element = element
|
||||
_image_url: str | None = None
|
||||
format: str | None = None
|
||||
detail: str | None = None
|
||||
if isinstance(img_element.get("image_url"), dict):
|
||||
_image_url = img_element["image_url"].get("url")
|
||||
format = img_element["image_url"].get("format")
|
||||
detail = img_element["image_url"].get("detail")
|
||||
else:
|
||||
_image_url = img_element.get("image_url")
|
||||
img_element = cast(ChatCompletionImageObject, element) # cast-ok: runtime type tag checked
|
||||
_image_url, format, detail = _image_url_fields(img_element)
|
||||
if _image_url and "https://" in _image_url:
|
||||
image_obj = convert_to_anthropic_image_obj(_image_url, format=format)
|
||||
converted_image_url = convert_generic_image_chunk_to_openai_image_obj(image_obj)
|
||||
|
|
|
|||
|
|
@ -292,7 +292,7 @@ class MistralConfig(OpenAIGPTConfig):
|
|||
file_id = file_content.get("file", {}).get("file_id")
|
||||
if file_id:
|
||||
# Replace 'file' with 'file_id'
|
||||
file_content["file_id"] = file_id
|
||||
file_content["file_id"] = file_id # pyright: ignore[reportGeneralTypeIssues] # legacy in-place rewrite of the block shape
|
||||
file_content.pop("file", None)
|
||||
return messages
|
||||
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
_should_convert_tool_call_to_json_mode,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
get_tool_call_names,
|
||||
hoist_images_from_tool_messages,
|
||||
)
|
||||
|
|
@ -336,7 +337,8 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
self, messages: list[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]:
|
||||
"""OpenAI no longer supports image_url as a string, so we need to convert it to a dict"""
|
||||
hoisted_messages: Final = hoist_images_from_tool_messages(messages)
|
||||
stripped_messages: Final = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
hoisted_messages: Final = hoist_images_from_tool_messages(stripped_messages)
|
||||
|
||||
async def _async_transform():
|
||||
for message in hoisted_messages:
|
||||
|
|
|
|||
|
|
@ -4,6 +4,11 @@ Shared utilities for the Soniox provider (https://soniox.com).
|
|||
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SubtitleToken,
|
||||
render_subtitle_tokens_as_srt,
|
||||
render_subtitle_tokens_as_vtt,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
# Soniox API base URL.
|
||||
|
|
@ -109,121 +114,13 @@ def render_soniox_tokens(tokens: list[dict[str, Any]]) -> str:
|
|||
return "".join(text_parts)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SRT / VTT subtitle rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Maximum number of tokens to group into a single subtitle cue.
|
||||
_CUE_MAX_TOKENS: Final[int] = 15
|
||||
|
||||
# Maximum duration (in ms) for a single cue before forcing a break.
|
||||
_CUE_MAX_DURATION_MS: Final[int] = 5000
|
||||
|
||||
|
||||
def _format_timestamp_srt(ms: int) -> str:
|
||||
"""Format milliseconds as SRT timestamp: HH:MM:SS,mmm"""
|
||||
ms = max(ms, 0)
|
||||
hours: Final = ms // 3_600_000
|
||||
ms %= 3_600_000
|
||||
minutes: Final = ms // 60_000
|
||||
ms %= 60_000
|
||||
seconds: Final = ms // 1_000
|
||||
millis: Final = ms % 1_000
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d},{millis:03d}"
|
||||
|
||||
|
||||
def _format_timestamp_vtt(ms: int) -> str:
|
||||
"""Format milliseconds as VTT timestamp: HH:MM:SS.mmm"""
|
||||
ms = max(ms, 0)
|
||||
hours: Final = ms // 3_600_000
|
||||
ms %= 3_600_000
|
||||
minutes: Final = ms // 60_000
|
||||
ms %= 60_000
|
||||
seconds: Final = ms // 1_000
|
||||
millis: Final = ms % 1_000
|
||||
return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{millis:03d}"
|
||||
|
||||
|
||||
def _group_tokens_into_cues(
|
||||
tokens: list[dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Group Soniox tokens into subtitle cues.
|
||||
|
||||
Each cue has:
|
||||
- start_ms: int
|
||||
- end_ms: int
|
||||
- text: str
|
||||
|
||||
Grouping heuristics:
|
||||
- A new cue starts when token count exceeds _CUE_MAX_TOKENS.
|
||||
- A new cue starts when duration exceeds _CUE_MAX_DURATION_MS.
|
||||
- A new cue starts when the speaker changes (if diarization is on).
|
||||
- Tokens without timestamps are appended to the current cue.
|
||||
"""
|
||||
cues: Final[list[dict[str, Any]]] = []
|
||||
current_tokens: list[str] = []
|
||||
current_start: int | None = None
|
||||
current_end: int | None = None
|
||||
current_speaker: Any | None = None
|
||||
|
||||
def _flush() -> None:
|
||||
if current_tokens and current_start is not None:
|
||||
text: Final = "".join(current_tokens).strip()
|
||||
if text:
|
||||
cues.append(
|
||||
{
|
||||
"start_ms": current_start,
|
||||
"end_ms": (current_end if current_end is not None else current_start),
|
||||
"text": text,
|
||||
}
|
||||
)
|
||||
|
||||
for token in tokens:
|
||||
start_ms = token.get("start_ms")
|
||||
end_ms = token.get("end_ms")
|
||||
text = token.get("text", "")
|
||||
speaker = token.get("speaker")
|
||||
|
||||
# Skip tokens with no timestamp data entirely if we have no cue started
|
||||
if start_ms is None and current_start is None:
|
||||
continue
|
||||
|
||||
# Speaker change forces a new cue
|
||||
if speaker is not None and speaker != current_speaker:
|
||||
_flush()
|
||||
current_tokens = []
|
||||
current_start = start_ms
|
||||
current_end = end_ms
|
||||
current_speaker = speaker
|
||||
current_tokens.append(text)
|
||||
continue
|
||||
|
||||
# Duration or token count exceeded -> flush
|
||||
should_break = False
|
||||
if (
|
||||
len(current_tokens) >= _CUE_MAX_TOKENS
|
||||
or current_start is not None
|
||||
and start_ms is not None
|
||||
and (start_ms - current_start) >= _CUE_MAX_DURATION_MS
|
||||
):
|
||||
should_break = True
|
||||
|
||||
if should_break:
|
||||
_flush()
|
||||
current_tokens = []
|
||||
current_start = start_ms
|
||||
current_end = end_ms
|
||||
current_tokens.append(text)
|
||||
else:
|
||||
if current_start is None:
|
||||
current_start = start_ms
|
||||
if end_ms is not None:
|
||||
current_end = end_ms
|
||||
current_tokens.append(text)
|
||||
|
||||
_flush()
|
||||
return cues
|
||||
def _soniox_token_to_subtitle_token(token: dict[str, Any]) -> SubtitleToken:
|
||||
return SubtitleToken(
|
||||
text=token.get("text", ""),
|
||||
start_ms=token.get("start_ms"),
|
||||
end_ms=token.get("end_ms"),
|
||||
speaker=token.get("speaker"),
|
||||
)
|
||||
|
||||
|
||||
def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
|
||||
|
|
@ -232,20 +129,7 @@ def render_soniox_tokens_as_srt(tokens: list[dict[str, Any]]) -> str:
|
|||
|
||||
Returns an empty string if no tokens have timestamp data.
|
||||
"""
|
||||
cues: Final = _group_tokens_into_cues(tokens)
|
||||
if not cues:
|
||||
return ""
|
||||
|
||||
lines: Final[list[str]] = []
|
||||
for idx, cue in enumerate(cues, start=1):
|
||||
start = _format_timestamp_srt(cue["start_ms"])
|
||||
end = _format_timestamp_srt(cue["end_ms"])
|
||||
lines.append(str(idx))
|
||||
lines.append(f"{start} --> {end}")
|
||||
lines.append(cue["text"])
|
||||
lines.append("") # blank line between cues
|
||||
|
||||
return "\n".join(lines)
|
||||
return render_subtitle_tokens_as_srt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
|
||||
|
||||
|
||||
def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
|
||||
|
|
@ -254,14 +138,4 @@ def render_soniox_tokens_as_vtt(tokens: list[dict[str, Any]]) -> str:
|
|||
|
||||
Returns the VTT header even if no cues are present.
|
||||
"""
|
||||
cues: Final = _group_tokens_into_cues(tokens)
|
||||
|
||||
lines: Final[list[str]] = ["WEBVTT", ""]
|
||||
for cue in cues:
|
||||
start = _format_timestamp_vtt(cue["start_ms"])
|
||||
end = _format_timestamp_vtt(cue["end_ms"])
|
||||
lines.append(f"{start} --> {end}")
|
||||
lines.append(cue["text"])
|
||||
lines.append("") # blank line between cues
|
||||
|
||||
return "\n".join(lines)
|
||||
return render_subtitle_tokens_as_vtt(tuple(_soniox_token_to_subtitle_token(token) for token in tokens))
|
||||
|
|
|
|||
|
|
@ -4,7 +4,8 @@ Translates from OpenAI's `/v1/chat/completions` to Together AI's `/v1/chat/compl
|
|||
Docs: https://docs.together.ai/docs/chat-overview
|
||||
"""
|
||||
|
||||
from collections.abc import Callable, Container, Coroutine
|
||||
from collections.abc import Callable, Container, Coroutine, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
Final,
|
||||
Literal,
|
||||
|
|
@ -12,11 +13,13 @@ from typing import (
|
|||
overload,
|
||||
)
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.exceptions import UnsupportedParamsError
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.utils import supports_function_calling, supports_response_schema
|
||||
from litellm.utils import supports_function_calling, supports_reasoning, supports_response_schema
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
|
@ -38,6 +41,34 @@ def _registry_verdict(model: str, flag: str, check: Callable[[str], bool]) -> bo
|
|||
return None
|
||||
|
||||
|
||||
ADJUSTABLE_EFFORT_REASONING_MODELS: Final = frozenset(
|
||||
{
|
||||
"openai/gpt-oss-120b",
|
||||
"openai/gpt-oss-20b",
|
||||
}
|
||||
)
|
||||
HYBRID_REASONING_MODELS: Final = frozenset(
|
||||
{
|
||||
"MiniMaxAI/MiniMax-M3",
|
||||
"Qwen/Qwen3.5-9B",
|
||||
"Qwen/Qwen3.6-Plus",
|
||||
"deepseek-ai/DeepSeek-V4-Pro",
|
||||
"moonshotai/Kimi-K3",
|
||||
"nvidia/nemotron-3-ultra-550b-a55b",
|
||||
"zai-org/GLM-5.2",
|
||||
}
|
||||
)
|
||||
HIGH_MAX_EFFORT_MODEL_PREFIX: Final = "deepseek-ai/DeepSeek-V4-Pro"
|
||||
EFFORT_TRANSLATION: Final = MappingProxyType({"minimal": "low", "xhigh": "high", "max": "high"})
|
||||
HIGH_MAX_EFFORT_TRANSLATION: Final = MappingProxyType(
|
||||
{"minimal": "high", "low": "high", "medium": "high", "xhigh": "max"}
|
||||
)
|
||||
|
||||
|
||||
class TogetherReasoningToggle(TypedDict):
|
||||
enabled: ReadOnly[bool]
|
||||
|
||||
|
||||
def _function_calling_verdict(model: str) -> bool | None:
|
||||
return _registry_verdict(
|
||||
model,
|
||||
|
|
@ -83,6 +114,36 @@ def _tool_params_to_drop(passed_params: Container[str], model: str, drop_params:
|
|||
)
|
||||
|
||||
|
||||
def _supports_together_reasoning(model: str) -> bool:
|
||||
if model in ADJUSTABLE_EFFORT_REASONING_MODELS or model in HYBRID_REASONING_MODELS:
|
||||
return True
|
||||
if model.startswith(HIGH_MAX_EFFORT_MODEL_PREFIX):
|
||||
return True
|
||||
return supports_reasoning(model, custom_llm_provider="together_ai")
|
||||
|
||||
|
||||
def _adjustable_effort(effort: str, model: str) -> str:
|
||||
if effort == "none":
|
||||
verbose_logger.debug(
|
||||
"together_ai model %s cannot disable reasoning; mapping reasoning_effort=none to low", model
|
||||
)
|
||||
return "low"
|
||||
return EFFORT_TRANSLATION.get(effort, effort)
|
||||
|
||||
|
||||
def _reasoning_effort_payload(effort: str, model: str) -> Mapping[str, object]:
|
||||
if effort == "default":
|
||||
return MappingProxyType({})
|
||||
if model in ADJUSTABLE_EFFORT_REASONING_MODELS:
|
||||
return MappingProxyType({"reasoning_effort": _adjustable_effort(effort, model)})
|
||||
if effort == "none":
|
||||
disable_reasoning: Final[TogetherReasoningToggle] = {"enabled": False}
|
||||
return MappingProxyType({"reasoning": disable_reasoning})
|
||||
if model.startswith(HIGH_MAX_EFFORT_MODEL_PREFIX):
|
||||
return MappingProxyType({"reasoning_effort": HIGH_MAX_EFFORT_TRANSLATION.get(effort, effort)})
|
||||
return MappingProxyType({"reasoning_effort": EFFORT_TRANSLATION.get(effort, effort)})
|
||||
|
||||
|
||||
def _drop_response_format(passed_params: Container[str], model: str, drop_params: bool) -> bool:
|
||||
if "response_format" not in passed_params:
|
||||
return False
|
||||
|
|
@ -153,6 +214,15 @@ class TogetherAIChatConfig(OpenAIGPTConfig):
|
|||
return super()._transform_messages(stripped, model, is_async=True)
|
||||
return super()._transform_messages(stripped, model, is_async=False)
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: inherited contract
|
||||
supported_params: Final = super().get_supported_openai_params(model)
|
||||
if not _supports_together_reasoning(model):
|
||||
return supported_params
|
||||
return [ # mutable-ok: the inherited contract returns a plain list; building fresh avoids mutating the base class's value
|
||||
*supported_params,
|
||||
"reasoning_effort",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
@ -165,4 +235,10 @@ class TogetherAIChatConfig(OpenAIGPTConfig):
|
|||
mapped_openai_params.pop(param)
|
||||
if _drop_response_format(mapped_openai_params, model, drop_params):
|
||||
mapped_openai_params.pop("response_format")
|
||||
effort: Final = mapped_openai_params.get("reasoning_effort")
|
||||
if not isinstance(effort, str):
|
||||
return mapped_openai_params
|
||||
mapped_openai_params.pop("reasoning_effort")
|
||||
for key, value in _reasoning_effort_payload(effort, model).items():
|
||||
mapped_openai_params.setdefault(key, value)
|
||||
return mapped_openai_params
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Handles calculating cost for together ai models
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import (
|
||||
|
|
@ -18,6 +19,12 @@ from litellm.constants import (
|
|||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
def has_together_registry_pricing(model: str, cost_map: Mapping[str, object]) -> bool:
|
||||
stripped: Final = model.removeprefix("together_ai/")
|
||||
entry: Final = cost_map.get(f"together_ai/{stripped}")
|
||||
return isinstance(entry, Mapping) and "input_cost_per_token" in entry
|
||||
|
||||
|
||||
# Extract the number of billion parameters from the model name
|
||||
# only used for together_computer LLMs
|
||||
def get_model_params_and_category(model_name, call_type: CallTypes) -> str:
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -34,6 +34,7 @@ from mcp.types import (
|
|||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl, BaseModel
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -72,6 +73,7 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
|||
MCPPerUserTokenCache,
|
||||
mcp_per_user_token_cache,
|
||||
resolve_mcp_auth,
|
||||
resolved_token_header,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
|
|
@ -99,6 +101,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
|
|||
build_token_exchanger,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
AuthorizationCodeConfig,
|
||||
ClientCredentialsConfig,
|
||||
CredError,
|
||||
|
|
@ -153,6 +156,8 @@ from litellm.types.mcp import (
|
|||
MCPAuth,
|
||||
MCPStdioConfig,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import (
|
||||
MCPInfo,
|
||||
|
|
@ -349,6 +354,7 @@ class MCPServerConfig(TypedDict, total=False):
|
|||
audience: str
|
||||
subject_token_type: str
|
||||
upstream_resource: str
|
||||
upstream_token_header: ReadOnly[str]
|
||||
id_jag_resource_token_endpoint: str
|
||||
id_jag_resource: str
|
||||
client_private_key: str
|
||||
|
|
@ -828,18 +834,6 @@ def _should_strip_caller_authorization(
|
|||
)
|
||||
|
||||
|
||||
def _without_authorization(
|
||||
headers: dict[str, str] | None,
|
||||
) -> dict[str, str] | None:
|
||||
"""A copy of ``headers`` with any ``Authorization`` key removed (case-insensitive), or
|
||||
None if nothing remains. Drops only the credential, keeping other forwarded headers.
|
||||
"""
|
||||
if not headers:
|
||||
return None
|
||||
filtered: Final = {k: v for k, v in headers.items() if k.lower() != "authorization"}
|
||||
return filtered or None
|
||||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
|
||||
|
||||
|
|
@ -914,7 +908,9 @@ def _resolve_openapi_tool_auth(
|
|||
|
||||
if isinstance(per_server, dict):
|
||||
authorization: Final = next((v for k, v in per_server.items() if k.lower() == "authorization"), None)
|
||||
merged: Final = merge_mcp_headers(extra_headers=forwarded, static_headers=_without_authorization(per_server))
|
||||
merged: Final = merge_mcp_headers(
|
||||
extra_headers=forwarded, static_headers=without_header(per_server, DEFAULT_CREDENTIAL_HEADER)
|
||||
)
|
||||
if authorization is None:
|
||||
byok: Final = _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
|
||||
return byok, merged, mcp_auth_header
|
||||
|
|
@ -981,7 +977,7 @@ def _client_forwarded_authorization_headers(
|
|||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
return _without_authorization(extra_headers)
|
||||
return without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
|
||||
return extra_headers
|
||||
|
||||
|
||||
|
|
@ -994,7 +990,7 @@ def _take_forwarded_authorization(
|
|||
if not headers:
|
||||
return None, headers
|
||||
value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
|
||||
return value, _without_authorization(headers)
|
||||
return value, without_header(headers, DEFAULT_CREDENTIAL_HEADER)
|
||||
|
||||
|
||||
def _passthrough_token_from_mcp_auth_header(
|
||||
|
|
@ -2166,6 +2162,7 @@ class MCPServerManager:
|
|||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
),
|
||||
upstream_resource=server_config.get("upstream_resource", None),
|
||||
upstream_token_header=server_config.get("upstream_token_header", None),
|
||||
# ID-JAG fields
|
||||
id_jag_resource_token_endpoint=server_config.get("id_jag_resource_token_endpoint", None),
|
||||
id_jag_resource=server_config.get("id_jag_resource", None),
|
||||
|
|
@ -2698,6 +2695,7 @@ class MCPServerManager:
|
|||
or (credentials_dict.get("subject_token_type") if credentials_dict else None)
|
||||
or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
upstream_resource=(credentials_dict.get("upstream_resource") if credentials_dict else None),
|
||||
upstream_token_header=(credentials_dict.get("upstream_token_header") if credentials_dict else None),
|
||||
# ID-JAG fields — read from credentials JSON blob
|
||||
id_jag_resource_token_endpoint=(
|
||||
credentials_dict.get("id_jag_resource_token_endpoint") if credentials_dict else None
|
||||
|
|
@ -3525,10 +3523,9 @@ class MCPServerManager:
|
|||
case Ok(auth):
|
||||
# NoOpAuth has no header_name and so never conflicts.
|
||||
header_name: Final[str | None] = getattr(auth, "header_name", None)
|
||||
conflicts: Final = bool(
|
||||
header_name and extra_headers and any(key.lower() == header_name.lower() for key in extra_headers)
|
||||
)
|
||||
if not conflicts:
|
||||
if header_name is None or not extra_headers:
|
||||
return auth, extra_headers
|
||||
if not has_header(extra_headers, header_name):
|
||||
return auth, extra_headers
|
||||
if isinstance(
|
||||
spec.config,
|
||||
|
|
@ -3540,9 +3537,10 @@ class MCPServerManager:
|
|||
# guardrail such as MCPJWTSigner, static_headers, or any other injected
|
||||
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
|
||||
# signer's JWT instead of the minted token and rejects it, and for M2M the
|
||||
# one-shot 401 refetch is lost with it). Drop the conflicting header so the
|
||||
# resolved token reaches upstream.
|
||||
return auth, _without_authorization(extra_headers)
|
||||
# one-shot 401 refetch is lost with it). Drop only the header the resolved
|
||||
# credential is about to occupy, so a static credential the operator aimed at a
|
||||
# DIFFERENT header still reaches upstream.
|
||||
return auth, without_header(extra_headers, header_name)
|
||||
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
|
||||
# header or static_headers) is intentional and wins; v1 applies those last.
|
||||
return None, extra_headers
|
||||
|
|
@ -3650,6 +3648,7 @@ class MCPServerManager:
|
|||
):
|
||||
spec = None
|
||||
auth_value: Final = await resolve_mcp_auth(resolved_server, mcp_auth_header) if spec is None else None
|
||||
auth_header_name: Final = resolved_token_header(resolved_server, mcp_auth_header) if spec is None else None
|
||||
|
||||
# Create sampling and elicitation callbacks for this client
|
||||
sampling_cb = (
|
||||
|
|
@ -3758,6 +3757,7 @@ class MCPServerManager:
|
|||
transport_type=transport,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
|
|
@ -5306,7 +5306,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
|
||||
elif mcp_server.is_client_forwarded_token:
|
||||
extra_headers = _client_forwarded_authorization_headers(
|
||||
mcp_server=mcp_server,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ with ``client_id``, ``client_secret``, and ``token_url``.
|
|||
|
||||
import asyncio
|
||||
import hashlib
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -313,9 +314,26 @@ async def resolve_mcp_auth(
|
|||
1. ``mcp_auth_header`` — per-request/per-user override
|
||||
2. OAuth2 client_credentials token — auto-fetched and cached
|
||||
3. ``server.authentication_token`` — static token from config/DB
|
||||
|
||||
``resolved_token_header`` answers, for the same two inputs, which header the value belongs in.
|
||||
"""
|
||||
if mcp_auth_header:
|
||||
return mcp_auth_header
|
||||
if server.has_client_credentials:
|
||||
return await mcp_oauth2_token_cache.async_get_token(server)
|
||||
return server.authentication_token
|
||||
|
||||
|
||||
def resolved_token_header(
|
||||
server: "MCPServer",
|
||||
mcp_auth_header: str | Mapping[str, str] | None = None,
|
||||
) -> str | None:
|
||||
"""Which upstream header the value ``resolve_mcp_auth`` just returned belongs in.
|
||||
|
||||
``None`` means keep the auth_type default. A caller-supplied ``mcp_auth_header`` is the caller's
|
||||
own credential aimed at the slot the upstream normally uses, so it never moves; only the values
|
||||
the gateway resolved from its own config (the minted M2M token, the static token) follow
|
||||
``upstream_token_header``. Same inputs and same branch order as ``resolve_mcp_auth``, so the two
|
||||
cannot disagree about which case they are in.
|
||||
"""
|
||||
return None if mcp_auth_header else server.upstream_token_header
|
||||
|
|
|
|||
|
|
@ -47,12 +47,14 @@ def sanitize_openapi_tool_name(raw_name: str) -> str:
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
from litellm.types.mcp import credential_redirect_hook, custom_credential_slot
|
||||
|
||||
|
||||
class _OpenAPIJSONSchema(TypedDict, total=False):
|
||||
|
|
@ -119,6 +121,10 @@ _request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | No
|
|||
"_request_resolved_auth_headers", default=None
|
||||
)
|
||||
|
||||
_request_upstream_url: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar(
|
||||
"_request_upstream_url", default=None
|
||||
)
|
||||
|
||||
|
||||
def _sanitize_path_parameter_value(param_value: object, param_name: str) -> str:
|
||||
"""Ensure path params cannot introduce directory traversal."""
|
||||
|
|
@ -349,6 +355,35 @@ def build_input_schema(operation: _OpenAPIOperation) -> dict[str, object]:
|
|||
}
|
||||
|
||||
|
||||
async def _drop_credential_across_origin(request: httpx.Request) -> None:
|
||||
"""Apply this request's cross-origin credential guard, if it needs one.
|
||||
|
||||
Reads the per-request context rather than closing over it so the hook is one stable object, which
|
||||
keeps the guarded client cacheable. A closure would key a new entry per call, and the handler it
|
||||
built would never be closed.
|
||||
"""
|
||||
guard: Final = credential_redirect_hook(
|
||||
_request_upstream_url.get() or "", custom_credential_slot(_request_resolved_auth_headers.get())
|
||||
)
|
||||
if guard is not None:
|
||||
await guard(request)
|
||||
|
||||
|
||||
def _upstream_client() -> AsyncHTTPHandler:
|
||||
"""The HTTP client for one upstream call, guarded when a credential rides a custom slot.
|
||||
|
||||
A resolved credential outside ``Authorization`` is not stripped across origins by the client
|
||||
itself, so this arm installs the same hook the MCP client uses. Both variants come from the
|
||||
shared cache, so a guarded call reuses its connection pool like any other.
|
||||
"""
|
||||
if custom_credential_slot(_request_resolved_auth_headers.get()) is None:
|
||||
return get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
return get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.MCP,
|
||||
params={"event_hooks": {"request": [_drop_credential_across_origin]}},
|
||||
)
|
||||
|
||||
|
||||
def _merge_openapi_tool_request_headers(
|
||||
static_headers: dict[str, str],
|
||||
) -> dict[str, str]:
|
||||
|
|
@ -510,8 +545,9 @@ def create_tool_function(
|
|||
except (json.JSONDecodeError, TypeError):
|
||||
json_body = {"data": body_value}
|
||||
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
client: Final = _upstream_client()
|
||||
upstream: Final = server_label or f"{original_method.upper()} {path}"
|
||||
url_token: Final = _request_upstream_url.set(url)
|
||||
|
||||
try:
|
||||
if original_method == "get":
|
||||
|
|
@ -529,6 +565,8 @@ def create_tool_function(
|
|||
except MaskedHTTPStatusError as e:
|
||||
_raise_for_upstream_failure(e.response, upstream, relays_upstream_auth)
|
||||
raise
|
||||
finally:
|
||||
_request_upstream_url.reset(url_token)
|
||||
|
||||
_raise_for_upstream_failure(response, upstream, relays_upstream_auth)
|
||||
return response.text
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
|||
Result,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
Ambient,
|
||||
ApiKeyConfig,
|
||||
ApiKeySource,
|
||||
|
|
@ -35,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
ClientCredentialsConfig,
|
||||
ClientSecretAuth,
|
||||
CredError,
|
||||
HeaderCarrier,
|
||||
IdJagConfig,
|
||||
NoneConfig,
|
||||
PassthroughConfig,
|
||||
|
|
@ -45,9 +47,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
Subject,
|
||||
TokenExchangeConfig,
|
||||
parse_auth_spec_kind,
|
||||
validate_header_name,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_CREDENTIAL_HEADER",
|
||||
"Ambient",
|
||||
"ApiKeyConfig",
|
||||
"ApiKeySource",
|
||||
|
|
@ -63,6 +67,7 @@ __all__ = [
|
|||
"ClientSecretAuth",
|
||||
"CredError",
|
||||
"Error",
|
||||
"HeaderCarrier",
|
||||
"IdJagConfig",
|
||||
"NoOpAuth",
|
||||
"NoneConfig",
|
||||
|
|
@ -78,4 +83,5 @@ __all__ = [
|
|||
"TokenExchangeConfig",
|
||||
"UpstreamCredentialProvider",
|
||||
"parse_auth_spec_kind",
|
||||
"validate_header_name",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from typing_extensions import assert_never
|
|||
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
ApiKeyConfig,
|
||||
AuthorizationCodeConfig,
|
||||
ClientAuth,
|
||||
|
|
@ -45,6 +46,15 @@ _TOKEN_EXCHANGE_SUBJECT_TOKEN_DEFAULT: Final = "urn:ietf:params:oauth:token-type
|
|||
_ID_JAG_SUBJECT_TOKEN_DEFAULT: Final = "urn:ietf:params:oauth:token-type:id_token"
|
||||
|
||||
|
||||
def token_header(server: MCPServer) -> str:
|
||||
"""The upstream header this server's resolved credential occupies.
|
||||
|
||||
One owner for every arm, so no spec builder spells the default itself and a server can never
|
||||
hand two arms different answers.
|
||||
"""
|
||||
return server.upstream_token_header or DEFAULT_CREDENTIAL_HEADER
|
||||
|
||||
|
||||
def to_subject(user_api_key_auth: UserAPIKeyAuth | None, subject_token: str | None) -> Subject:
|
||||
"""Map v1's authenticated principal onto the resolver's Subject.
|
||||
|
||||
|
|
@ -122,7 +132,7 @@ def _oauth2_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
|||
return ServerSpec(
|
||||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=AuthorizationCodeConfig(),
|
||||
config=AuthorizationCodeConfig(header_name=token_header(server)),
|
||||
)
|
||||
return None
|
||||
|
||||
|
|
@ -140,6 +150,7 @@ def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
|
|||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=ClientCredentialsConfig(
|
||||
header_name=token_header(server),
|
||||
client_id=server.client_id,
|
||||
client_secret=SecretStr(server.client_secret) if server.client_secret else None,
|
||||
token_url=server.effective_token_url,
|
||||
|
|
@ -173,6 +184,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
|||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=TokenExchangeConfig(
|
||||
header_name=token_header(server),
|
||||
profile=profile,
|
||||
subject_token_type=server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
token_exchange_endpoint=endpoint,
|
||||
|
|
@ -206,7 +218,7 @@ def _shared_key_spec(
|
|||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=ApiKeyConfig(
|
||||
header_name=header_name,
|
||||
header_name=server.upstream_token_header or header_name,
|
||||
value_prefix=value_prefix,
|
||||
key_source=SharedKey(value=SecretStr(value)),
|
||||
),
|
||||
|
|
@ -231,6 +243,7 @@ def _id_jag_spec(server: MCPServer, resource: str) -> ServerSpec | None:
|
|||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
config=IdJagConfig(
|
||||
header_name=token_header(server),
|
||||
org_token_endpoint=org_token_endpoint,
|
||||
resource_token_endpoint=resource_token_endpoint,
|
||||
client_id=client_id,
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
|||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ClientCredentialsConfig,
|
||||
CredError,
|
||||
HeaderCarrier,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -328,14 +329,21 @@ class ClientCredentialsBearerAuth(httpx.Auth):
|
|||
refetch fails, or the retried request 401s again, the upstream's response stands.
|
||||
"""
|
||||
|
||||
def __init__(self, access_token: str, refetch: Callable[[str], Awaitable[str | None]]) -> None:
|
||||
self.header_name = "Authorization"
|
||||
def __init__(
|
||||
self,
|
||||
access_token: str,
|
||||
refetch: Callable[[str], Awaitable[str | None]],
|
||||
carrier: HeaderCarrier,
|
||||
) -> None:
|
||||
self._carrier = carrier
|
||||
self.header_name = carrier.header_name
|
||||
self._access_token = SecretStr(access_token)
|
||||
self._refetch = refetch
|
||||
|
||||
async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]:
|
||||
token: Final = self._access_token.get_secret_value()
|
||||
request.headers[self.header_name] = f"Bearer {token}"
|
||||
name, value = self._carrier.header(token)
|
||||
request.headers[name] = value
|
||||
response: Final = yield request
|
||||
if response.status_code != 401:
|
||||
return
|
||||
|
|
@ -343,7 +351,8 @@ class ClientCredentialsBearerAuth(httpx.Auth):
|
|||
if fresh is None:
|
||||
return
|
||||
self._access_token = SecretStr(fresh)
|
||||
request.headers[self.header_name] = f"Bearer {fresh}"
|
||||
fresh_name, fresh_value = self._carrier.header(fresh)
|
||||
request.headers[fresh_name] = fresh_value
|
||||
yield request
|
||||
|
||||
def sync_auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]:
|
||||
|
|
|
|||
|
|
@ -145,8 +145,8 @@ class UpstreamCredentialProvider:
|
|||
return await self._token_exchange(subject, server, config)
|
||||
case IdJagConfig() as config:
|
||||
return await self._id_jag(subject, server, config)
|
||||
case AuthorizationCodeConfig():
|
||||
return await self._authorization_code(subject, server)
|
||||
case AuthorizationCodeConfig() as config:
|
||||
return await self._authorization_code(subject, server, config)
|
||||
case AwsSigV4Config():
|
||||
return _not_implemented(AuthSpecKind.aws_sigv4)
|
||||
assert_never(server.config)
|
||||
|
|
@ -284,15 +284,19 @@ class UpstreamCredentialProvider:
|
|||
|
||||
match await self._exchanged_tokens.get_or_compute(slot, _exchange, fingerprint=fingerprint):
|
||||
case Ok(access_token):
|
||||
return Ok(StaticHeaderAuth(f"Bearer {access_token}"))
|
||||
header_name, header_value = config.header(access_token)
|
||||
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
|
||||
case Error(err):
|
||||
return Error(err)
|
||||
|
||||
async def _authorization_code(self, subject: Subject, server: ServerSpec) -> Result[StaticHeaderAuth, CredError]:
|
||||
async def _authorization_code(
|
||||
self, subject: Subject, server: ServerSpec, config: AuthorizationCodeConfig
|
||||
) -> Result[StaticHeaderAuth, CredError]:
|
||||
token: Final = await self._authz_token(subject, server)
|
||||
if token is None:
|
||||
return Error(CredError.of_unauthorized("Authorization required: complete the OAuth flow for this server."))
|
||||
return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization"))
|
||||
header_name, header_value = config.header(token.access_token)
|
||||
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
|
||||
|
||||
async def _client_credentials(
|
||||
self, server_id: str, config: ClientCredentialsConfig
|
||||
|
|
@ -307,7 +311,7 @@ class UpstreamCredentialProvider:
|
|||
match await self._client_credentials_source.get(server_id, config):
|
||||
case Ok(token):
|
||||
refetch: Final = partial(self._client_credentials_source.refetch, server_id, config)
|
||||
return Ok(ClientCredentialsBearerAuth(token.access_token, refetch))
|
||||
return Ok(ClientCredentialsBearerAuth(token.access_token, refetch, config))
|
||||
case Error(err):
|
||||
return Error(err)
|
||||
|
||||
|
|
@ -332,7 +336,8 @@ class UpstreamCredentialProvider:
|
|||
inbound.get_secret_value(), server, config, tenant_id=subject.tenant_id
|
||||
):
|
||||
case Ok(token):
|
||||
return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization"))
|
||||
header_name, header_value = config.header(token.access_token)
|
||||
return Ok(StaticHeaderAuth(header_value, header_name=header_name))
|
||||
case Error(err):
|
||||
return Error(err)
|
||||
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ from enum import Enum
|
|||
from typing import Annotated, Final, Literal
|
||||
|
||||
from expression import case, tag, tagged_union
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
||||
|
|
@ -39,7 +39,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
|||
Ok,
|
||||
Result,
|
||||
)
|
||||
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
|
||||
from litellm.types.mcp import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
normalize_upstream_header_name,
|
||||
)
|
||||
|
||||
|
||||
class AuthSpecKind(str, Enum):
|
||||
|
|
@ -161,7 +165,52 @@ class CredError:
|
|||
assert_never(self.tag)
|
||||
|
||||
|
||||
class AuthorizationCodeConfig(BaseModel):
|
||||
def validate_header_name(raw: str) -> Result[str, CredError]:
|
||||
"""``normalize_upstream_header_name`` with this package's error-as-value policy.
|
||||
|
||||
The grammar itself lives in ``litellm.types.mcp`` so the v1 model, the management endpoint and
|
||||
this vocabulary all judge a header name the same way while each keeps its own failure shape.
|
||||
"""
|
||||
normalized: Final = normalize_upstream_header_name(raw)
|
||||
if normalized is None:
|
||||
return Error(CredError.of_misconfigured(f"invalid upstream header name: {raw!r}"))
|
||||
return Ok(normalized)
|
||||
|
||||
|
||||
class HeaderCarrier(BaseModel):
|
||||
"""Where a resolved credential is written upstream, and how its value is formatted.
|
||||
|
||||
``Authorization: Bearer`` is only OAuth's *default* conveyance (RFC 6750 section 2.1), not its
|
||||
only one: an ESB or API gateway commonly terminates its own credential in a private header while
|
||||
a second credential passes through to the origin, so a credential has to be able to say which
|
||||
slot it owns. Modeled like OpenAPI's apiKey scheme, so any upstream convention is expressible
|
||||
(Authorization + Bearer, a raw value on X-API-Key, Ocp-Apim-Subscription-Key, esb-oauth, ...).
|
||||
|
||||
Every config whose credential the gateway mints or holds inherits this, so no resolver arm names
|
||||
a header itself and the conflict rule in ``_resolve_v2_auth`` can always ask the auth object
|
||||
which slot it is about to occupy. ``passthrough`` deliberately does not: it forwards the
|
||||
caller's own credential into the slot the caller used, and mints nothing to place.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
header_name: str = DEFAULT_CREDENTIAL_HEADER
|
||||
value_prefix: str = "Bearer"
|
||||
|
||||
@field_validator("header_name")
|
||||
@classmethod
|
||||
def _check_header_name(cls, value: str) -> str:
|
||||
match validate_header_name(value):
|
||||
case Ok(name):
|
||||
return name
|
||||
case Error(err):
|
||||
raise ValueError(err.summary)
|
||||
|
||||
def header(self, value: str) -> tuple[str, str]:
|
||||
formatted: Final = f"{self.value_prefix} {value}" if self.value_prefix else value
|
||||
return self.header_name, formatted
|
||||
|
||||
|
||||
class AuthorizationCodeConfig(HeaderCarrier):
|
||||
"""Per-user 3LO; the gateway is the OAuth client and stores the user's token.
|
||||
|
||||
Endpoints are discovered (RFC 9728 -> RFC 8414) and the client is registered via DCR
|
||||
|
|
@ -179,7 +228,7 @@ class AuthorizationCodeConfig(BaseModel):
|
|||
token_url: str | None = None
|
||||
|
||||
|
||||
class ClientCredentialsConfig(BaseModel):
|
||||
class ClientCredentialsConfig(HeaderCarrier):
|
||||
"""M2M service account; one upstream identity for every user.
|
||||
|
||||
Fields are optional so the config can be built incomplete: a value may be supplied at
|
||||
|
|
@ -203,7 +252,7 @@ class ClientCredentialsConfig(BaseModel):
|
|||
token_endpoint_auth_method: Literal["client_secret_post", "client_secret_basic"] | None = None
|
||||
|
||||
|
||||
class TokenExchangeConfig(BaseModel):
|
||||
class TokenExchangeConfig(HeaderCarrier):
|
||||
"""OBO: swap the caller's live inbound token for a token bound to the upstream's audience. The
|
||||
gateway authenticates to the exchange endpoint as an OAuth client (`client_id`/`client_secret`);
|
||||
the inbound token is sent only to that endpoint, never to the upstream.
|
||||
|
|
@ -255,7 +304,7 @@ class ClientSecretAuth(BaseModel):
|
|||
ClientAuth = Annotated[PrivateKeyJwtAuth | ClientSecretAuth, Field(discriminator="source")]
|
||||
|
||||
|
||||
class IdJagConfig(BaseModel):
|
||||
class IdJagConfig(HeaderCarrier):
|
||||
"""draft-ietf-oauth-identity-assertion-authz-grant (Okta "AI agent token exchange").
|
||||
|
||||
Two legs: leg 1 is an RFC 8693 token exchange at the IdP org AS (`org_token_endpoint`) that
|
||||
|
|
@ -297,23 +346,16 @@ class Byok(BaseModel):
|
|||
ApiKeySource = Annotated[SharedKey | Byok, Field(discriminator="source")]
|
||||
|
||||
|
||||
class ApiKeyConfig(BaseModel):
|
||||
class ApiKeyConfig(HeaderCarrier):
|
||||
"""A fixed credential injected as a header. The value is shared (in config) or seeded
|
||||
per-user (pulled from the store); `header_name` and `value_prefix` say where and how it is
|
||||
written, modeled like OpenAPI's apiKey scheme so any upstream convention is expressible
|
||||
(Authorization + Bearer, a raw value on X-API-Key, Ocp-Apim-Subscription-Key, etc.).
|
||||
per-user (pulled from the store); the inherited `header_name` and `value_prefix` say where
|
||||
and how it is written.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal[AuthSpecKind.api_key] = AuthSpecKind.api_key
|
||||
header_name: str = "Authorization"
|
||||
value_prefix: str = "Bearer"
|
||||
key_source: ApiKeySource
|
||||
|
||||
def header(self, value: str) -> tuple[str, str]:
|
||||
formatted: Final = f"{self.value_prefix} {value}" if self.value_prefix else value
|
||||
return self.header_name, formatted
|
||||
|
||||
|
||||
class PassthroughConfig(BaseModel):
|
||||
"""Client-driven upstream OAuth; the gateway forwards the client's upstream token."""
|
||||
|
|
|
|||
|
|
@ -246,11 +246,12 @@ def _mcp_meta_trace_carrier(req_ctx: object) -> dict[str, str] | None:
|
|||
"""The W3C trace context (``traceparent``/``tracestate``) the MCP client
|
||||
propagated in the request's ``params._meta`` (SEP-414), or ``None``.
|
||||
|
||||
When present, per the OTel MCP semconv the MCP span parents to this propagated
|
||||
context rather than to the HTTP transport (which is recorded as a link instead).
|
||||
When absent, the span nests under the transport span of the request carrying
|
||||
this specific message, so a streamable-HTTP session that multiplexes many
|
||||
messages still does not glue every message under the session's first request;
|
||||
When present, the MCP span records this propagated context as a span *link*,
|
||||
never the parent — a remote parent would root the span in a trace whose root
|
||||
never reaches the gateway's tracing backend. The span itself nests under the
|
||||
transport span of the request carrying this specific message, so a
|
||||
streamable-HTTP session that multiplexes many messages still does not glue
|
||||
every message under the session's first request;
|
||||
see ``resolve_mcp_span_context``. The client's W3C Baggage is
|
||||
deliberately excluded: it is caller-controlled, and the otel baggage processor
|
||||
stamps allowlisted baggage keys (``litellm.team.id``, ``litellm.metadata.*``,
|
||||
|
|
@ -432,7 +433,6 @@ if MCP_AVAILABLE:
|
|||
_client_forwarded_authorization_headers,
|
||||
_resolve_openapi_tool_auth,
|
||||
_should_strip_caller_authorization,
|
||||
_without_authorization,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
|
|
@ -451,6 +451,7 @@ if MCP_AVAILABLE:
|
|||
split_server_prefix_from_name,
|
||||
strip_known_server_prefix,
|
||||
)
|
||||
from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, without_header
|
||||
|
||||
######################################################
|
||||
############ MCP Tools List REST API Response Object #
|
||||
|
|
@ -1732,7 +1733,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
):
|
||||
extra_headers = _without_authorization(extra_headers)
|
||||
extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
|
||||
elif is_client_forwarded_mode:
|
||||
if not withhold_forwarded_authorization:
|
||||
extra_headers = _client_forwarded_authorization_headers(
|
||||
|
|
|
|||
|
|
@ -15038,6 +15038,17 @@
|
|||
}
|
||||
],
|
||||
"title": "Upstream Resource"
|
||||
},
|
||||
"upstream_token_header": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Upstream Token Header"
|
||||
}
|
||||
},
|
||||
"title": "MCPCredentials",
|
||||
|
|
@ -17518,6 +17529,17 @@
|
|||
}
|
||||
],
|
||||
"title": "Upstream Resource"
|
||||
},
|
||||
"upstream_token_header": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Upstream Token Header"
|
||||
}
|
||||
},
|
||||
"title": "MCPCredentials",
|
||||
|
|
@ -20352,6 +20374,17 @@
|
|||
}
|
||||
],
|
||||
"title": "Upstream Resource"
|
||||
},
|
||||
"upstream_token_header": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Upstream Token Header"
|
||||
}
|
||||
},
|
||||
"title": "MCPCredentials",
|
||||
|
|
@ -23699,6 +23732,17 @@
|
|||
}
|
||||
],
|
||||
"title": "Upstream Resource"
|
||||
},
|
||||
"upstream_token_header": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Upstream Token Header"
|
||||
}
|
||||
},
|
||||
"title": "MCPCredentials",
|
||||
|
|
|
|||
|
|
@ -686,6 +686,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
aws_profile_name: Final = self.optional_params.get("aws_profile_name", None)
|
||||
aws_web_identity_token: Final = self.optional_params.get("aws_web_identity_token", None)
|
||||
aws_sts_endpoint: Final = self.optional_params.get("aws_sts_endpoint", None)
|
||||
aws_external_id: Final = self.optional_params.get("aws_external_id", None)
|
||||
|
||||
### SET REGION NAME ###
|
||||
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
||||
|
|
@ -702,6 +703,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
aws_role_name=aws_role_name,
|
||||
aws_web_identity_token=aws_web_identity_token,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
)
|
||||
return credentials, aws_region_name
|
||||
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
aws_role_name=litellm_params.aws_role_name,
|
||||
aws_web_identity_token=litellm_params.aws_web_identity_token,
|
||||
aws_sts_endpoint=litellm_params.aws_sts_endpoint,
|
||||
aws_external_id=litellm_params.aws_external_id,
|
||||
aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint,
|
||||
experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only,
|
||||
only_scan_new_messages=litellm_params.only_scan_new_messages or False,
|
||||
|
|
|
|||
|
|
@ -204,6 +204,7 @@ if MCP_AVAILABLE:
|
|||
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS,
|
||||
MCPAuth,
|
||||
MCPCredentials,
|
||||
normalize_upstream_header_name,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
|
@ -239,9 +240,26 @@ if MCP_AVAILABLE:
|
|||
detail={"error": error_messages_text},
|
||||
)
|
||||
|
||||
def _validate_upstream_token_header(payload: McpServerPayloadLike) -> None:
|
||||
credentials: Final = getattr(payload, "credentials", None)
|
||||
raw: Final = credentials.get("upstream_token_header") if isinstance(credentials, dict) else None
|
||||
if not isinstance(raw, str) or raw == "":
|
||||
return
|
||||
if normalize_upstream_header_name(raw) is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail={
|
||||
"error": (
|
||||
f"Invalid upstream_token_header {raw!r}: must be a valid HTTP header name "
|
||||
"(RFC 7230 token, e.g. 'esb-oauth')"
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
def validate_and_normalize_mcp_server_payload(payload: McpServerPayloadLike) -> None:
|
||||
_base_validate_and_normalize_mcp_server_payload(payload)
|
||||
_validate_mcp_server_name_fields(payload)
|
||||
_validate_upstream_token_header(payload)
|
||||
|
||||
def stamp_omitted_oauth2_flow(payload: NewMCPServerRequest) -> None:
|
||||
"""Fallback only: fill in oauth2_flow when an oauth2 create omits it.
|
||||
|
|
@ -739,6 +757,7 @@ if MCP_AVAILABLE:
|
|||
("aws_region_name", "aws_region_name"),
|
||||
("aws_service_name", "aws_service_name"),
|
||||
("upstream_resource", "upstream_resource"),
|
||||
("upstream_token_header", "upstream_token_header"),
|
||||
)
|
||||
|
||||
def _has_non_admin_config_credentials(credentials: "MCPCredentials | None") -> bool:
|
||||
|
|
|
|||
|
|
@ -9253,6 +9253,7 @@ class ProxyStartupEvent:
|
|||
prisma_client,
|
||||
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
|
||||
alert=_alert_ptu_rollup_failure,
|
||||
router=llm_router,
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
|
|
|
|||
|
|
@ -14,7 +14,6 @@ and share the existing unique constraint.
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, time, timedelta, timezone
|
||||
|
|
@ -326,16 +325,6 @@ class _LoadedDeployments:
|
|||
scanned_ids: frozenset[str]
|
||||
|
||||
|
||||
def _running_router() -> object | None:
|
||||
"""The proxy's router, or None outside a running proxy.
|
||||
|
||||
Read out of ``sys.modules`` rather than imported, so a rollup driven from a test or a
|
||||
script does not pull the whole proxy server in behind it.
|
||||
"""
|
||||
proxy_server: Final = sys.modules.get("litellm.proxy.proxy_server")
|
||||
return getattr(proxy_server, "llm_router", None) if proxy_server is not None else None
|
||||
|
||||
|
||||
def _config_deployments(router: object | None, *, owned_by_db: frozenset[str]) -> tuple[_PTUDeployment, ...]:
|
||||
"""Deployments the router holds that no ``LiteLLM_ProxyModelTable`` row owns.
|
||||
|
||||
|
|
@ -356,15 +345,17 @@ def _config_deployments(router: object | None, *, owned_by_db: frozenset[str]) -
|
|||
)
|
||||
|
||||
|
||||
async def _load_ptu_models(prisma_client: "PrismaClient") -> _LoadedDeployments:
|
||||
async def _load_ptu_models(prisma_client: "PrismaClient", *, router: object | None) -> _LoadedDeployments:
|
||||
"""Every deployment carrying valid manual PTU config, and every id the scan saw.
|
||||
|
||||
Reserved capacity is billed by the provider whichever file declared it, so a
|
||||
deployment the proxy only knows from config.yaml accrues alongside the stored ones.
|
||||
The router is handed in rather than read off the proxy module, so a run prices exactly
|
||||
the deployments its caller declares and nothing a co-resident process left behind.
|
||||
"""
|
||||
rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many()
|
||||
db_ids: Final = frozenset(model_id for row in rows if (model_id := str(getattr(row, "model_id", "") or "")))
|
||||
config_records: Final = _config_deployments(_running_router(), owned_by_db=db_ids)
|
||||
config_records: Final = _config_deployments(router, owned_by_db=db_ids)
|
||||
models: Final = tuple(
|
||||
parsed for parsed in (_parse_ptu_model(row) for row in (*rows, *config_records)) if parsed is not None
|
||||
)
|
||||
|
|
@ -380,6 +371,7 @@ async def run_ptu_flat_cost_rollup(
|
|||
prisma_client: "PrismaClient",
|
||||
target_date: date | None = None,
|
||||
may_prune: bool = True,
|
||||
router: object | None = None,
|
||||
) -> RollupResult:
|
||||
"""Rollup one UTC day of flat PTU cost across all PTU-configured model deployments.
|
||||
|
||||
|
|
@ -406,7 +398,7 @@ async def run_ptu_flat_cost_rollup(
|
|||
date_str: Final = day.isoformat()
|
||||
run_started: Final = datetime.now(timezone.utc)
|
||||
|
||||
loaded: Final = await _load_ptu_models(prisma_client)
|
||||
loaded: Final = await _load_ptu_models(prisma_client, router=router)
|
||||
ptu_models: Final = loaded.models
|
||||
charges: Final = _aggregate_charges(ptu_models, day)
|
||||
|
||||
|
|
@ -527,6 +519,7 @@ async def _existing_sentinel_keys(
|
|||
async def run_ptu_flat_cost_backfill(
|
||||
prisma_client: "PrismaClient",
|
||||
today: date | None = None,
|
||||
router: object | None = None,
|
||||
) -> BackfillResult:
|
||||
"""Price the elapsed days of every PTU window that carry no sentinel row yet.
|
||||
|
||||
|
|
@ -546,7 +539,7 @@ async def run_ptu_flat_cost_backfill(
|
|||
verbose_proxy_logger.warning("PTU backfill: prisma_client is None, skipping")
|
||||
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
|
||||
|
||||
ptu_models: Final = (await _load_ptu_models(prisma_client)).models
|
||||
ptu_models: Final = (await _load_ptu_models(prisma_client, router=router)).models
|
||||
days: Final = _backfill_window(ptu_models, end)
|
||||
|
||||
if not days:
|
||||
|
|
@ -591,6 +584,7 @@ async def run_scheduled_ptu_rollup(
|
|||
pod_lock_manager: "PodLockManager | None" = None,
|
||||
target_date: date | None = None,
|
||||
alert: Callable[[str], Awaitable[None]] | None = None,
|
||||
router: object | None = None,
|
||||
) -> RollupResult | None:
|
||||
"""Run the daily rollup under a cross-pod lock so only one proxy reconciles a day.
|
||||
|
||||
|
|
@ -615,7 +609,7 @@ async def run_scheduled_ptu_rollup(
|
|||
return None
|
||||
|
||||
if pod_lock_manager is None or pod_lock_manager.redis_cache is None:
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router)
|
||||
|
||||
if not await pod_lock_manager.acquire_lock(cronjob_id=PTU_ROLLUP_JOB_ID, ttl=PTU_ROLLUP_LOCK_TTL_SECONDS):
|
||||
if await _lock_is_held(pod_lock_manager):
|
||||
|
|
@ -629,10 +623,10 @@ async def run_scheduled_ptu_rollup(
|
|||
"PTU rollup: could not take the rollup lock and no other pod holds it, "
|
||||
"running unguarded rather than skipping the day"
|
||||
)
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router)
|
||||
|
||||
try:
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True)
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True, router=router)
|
||||
finally:
|
||||
await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID)
|
||||
|
||||
|
|
@ -657,6 +651,7 @@ async def _run_and_alert(
|
|||
target_date: date | None,
|
||||
alert: "Callable[[str], Awaitable[None]] | None",
|
||||
may_prune: bool = True,
|
||||
router: object | None = None,
|
||||
) -> RollupResult:
|
||||
"""Reconcile the day, catch up any days left unpriced, and alert on charges that did not land.
|
||||
|
||||
|
|
@ -669,7 +664,9 @@ async def _run_and_alert(
|
|||
explicit date means reconcile exactly that day, so it stays a single-day operation.
|
||||
Its failure is contained: the day's own result is returned either way.
|
||||
"""
|
||||
result: Final = await run_ptu_flat_cost_rollup(prisma_client, target_date=target_date, may_prune=may_prune)
|
||||
result: Final = await run_ptu_flat_cost_rollup(
|
||||
prisma_client, target_date=target_date, may_prune=may_prune, router=router
|
||||
)
|
||||
if result.rows_failed:
|
||||
await _deliver_alert(
|
||||
alert,
|
||||
|
|
@ -686,7 +683,7 @@ async def _run_and_alert(
|
|||
"by the provider with nothing attributing it here. Extend the window, or retire the deployment.",
|
||||
)
|
||||
if target_date is None:
|
||||
await _backfill_and_alert(prisma_client, alert=alert)
|
||||
await _backfill_and_alert(prisma_client, alert=alert, router=router)
|
||||
return result
|
||||
|
||||
|
||||
|
|
@ -694,6 +691,7 @@ async def _backfill_and_alert(
|
|||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
alert: "Callable[[str], Awaitable[None]] | None",
|
||||
router: object | None = None,
|
||||
) -> None:
|
||||
"""Catch up unpriced PTU days, alerting on charges that did not land.
|
||||
|
||||
|
|
@ -701,7 +699,7 @@ async def _backfill_and_alert(
|
|||
caller whatever the catch-up pass does.
|
||||
"""
|
||||
try:
|
||||
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client)
|
||||
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client, router=router)
|
||||
except Exception as exc: # noqa: BLE001 # the catch-up pass must not fail the day's rollup
|
||||
verbose_proxy_logger.error("PTU backfill: catch-up pass failed, the day's rollup still stands: %s", exc)
|
||||
return
|
||||
|
|
|
|||
|
|
@ -2661,6 +2661,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
optional_output_details: Final[dict[str, int]] = {
|
||||
field: value
|
||||
for field, value in (
|
||||
("audio_tokens", getattr(completion_details, "audio_tokens", None)),
|
||||
("text_tokens", getattr(completion_details, "text_tokens", None)),
|
||||
("image_tokens", getattr(completion_details, "image_tokens", None)),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -496,6 +496,9 @@ class BedrockGuardrailConfigModel(BaseModel):
|
|||
aws_role_name: str | None = Field(default=None, description="AWS role name for assuming roles")
|
||||
aws_web_identity_token: str | None = Field(default=None, description="Web identity token for AWS role assumption")
|
||||
aws_sts_endpoint: str | None = Field(default=None, description="AWS STS endpoint URL")
|
||||
aws_external_id: str | None = Field(
|
||||
default=None, description="External ID required by the target role's trust policy on sts:AssumeRole"
|
||||
)
|
||||
aws_bedrock_runtime_endpoint: str | None = Field(default=None, description="AWS Bedrock runtime endpoint URL")
|
||||
checks: BedrockChecksConfigModel | None = Field(
|
||||
default=None,
|
||||
|
|
|
|||
|
|
@ -324,7 +324,12 @@ class AnthropicMessagesToolResultParam(TypedDict, total=False):
|
|||
is_error: bool
|
||||
content: (
|
||||
str
|
||||
| Iterable[AnthropicMessagesToolResultContent | AnthropicMessagesImageParam | AnthropicMessagesDocumentParam]
|
||||
| Iterable[
|
||||
AnthropicMessagesToolResultContent
|
||||
| AnthropicMessagesImageParam
|
||||
| AnthropicMessagesDocumentParam
|
||||
| ToolReference
|
||||
]
|
||||
)
|
||||
cache_control: dict | ChatCompletionCachedContent | None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from collections.abc import Iterable, Mapping
|
||||
from enum import Enum
|
||||
from os import PathLike
|
||||
from typing import IO, Any, Final, Literal, Optional, Union
|
||||
from typing import IO, Any, Final, Literal, Optional, TypeAlias, Union
|
||||
|
||||
import httpx
|
||||
from openai import Omit
|
||||
|
|
@ -820,9 +820,21 @@ class ChatCompletionAssistantMessage(OpenAIChatCompletionAssistantMessage, total
|
|||
reasoning_items: list[ChatCompletionReasoningItem] | None
|
||||
|
||||
|
||||
class ChatCompletionToolReferenceObject(TypedDict):
|
||||
"""Anthropic tool-search result block, carried through untouched so it survives a round trip."""
|
||||
|
||||
type: Literal["tool_reference"] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
tool_name: str # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
|
||||
|
||||
ToolMessageContentPart: TypeAlias = (
|
||||
ChatCompletionTextObject | ChatCompletionImageObject | ChatCompletionToolReferenceObject
|
||||
)
|
||||
|
||||
|
||||
class ChatCompletionToolMessage(TypedDict):
|
||||
role: Literal["tool"]
|
||||
content: str | Iterable[ChatCompletionTextObject | ChatCompletionImageObject]
|
||||
content: str | Iterable[ToolMessageContentPart] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
|
||||
tool_call_id: str
|
||||
|
||||
|
||||
|
|
@ -1258,6 +1270,8 @@ class ResponsesAPIRequestParams(ResponsesAPIOptionalRequestParams, total=False):
|
|||
|
||||
|
||||
class OutputTokensDetails(BaseLiteLLMOpenAIResponseObject):
|
||||
audio_tokens: int | None = None
|
||||
|
||||
reasoning_tokens: int | None = None
|
||||
|
||||
text_tokens: int | None = None
|
||||
|
|
|
|||
|
|
@ -1,6 +1,11 @@
|
|||
import enum
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
|
@ -181,6 +186,15 @@ class MCPCredentials(TypedDict, total=False):
|
|||
``audience``, which is the RFC 8693 token-exchange parameter.
|
||||
"""
|
||||
|
||||
upstream_token_header: str | None # writable-ok: pydantic warns it cannot honour ReadOnly here
|
||||
"""
|
||||
Which upstream header carries the credential LiteLLM resolves for this server. Omitted when
|
||||
unset, which keeps RFC 6750's default of ``Authorization``. Set it when the upstream expects the
|
||||
gateway's token somewhere else (an ESB terminating its own credential on e.g. ``esb-oauth``), so
|
||||
a separate operator-configured ``Authorization`` reaches the origin untouched. Non-secret, so it
|
||||
is stored in plaintext and returned on admin reads.
|
||||
"""
|
||||
|
||||
client_private_key: str | None
|
||||
"""
|
||||
PEM private key used to sign the private-key-JWT client_assertion (RFC 7523)
|
||||
|
|
@ -223,7 +237,92 @@ class MCPCredentials(TypedDict, total=False):
|
|||
"""
|
||||
|
||||
|
||||
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: Final[tuple[str, ...]] = ("upstream_resource",)
|
||||
DEFAULT_CREDENTIAL_HEADER: Final = "Authorization"
|
||||
|
||||
_HEADER_NAME_TOKEN: Final = re.compile(r"^[!#$%&'*+\-.^_`|~0-9A-Za-z]+$")
|
||||
|
||||
|
||||
def normalize_upstream_header_name(raw: str) -> str | None:
|
||||
"""The trimmed header name if it is a usable RFC 7230 ``token``, else None.
|
||||
|
||||
One owner for the grammar; each caller picks its own failure shape (a config-load raise, an
|
||||
API 400, a typed CredError). An operator-supplied name reaches egress verbatim, so a value
|
||||
carrying CR/LF, spaces or separators must never get that far.
|
||||
"""
|
||||
stripped: Final = raw.strip()
|
||||
return stripped if stripped and _HEADER_NAME_TOKEN.match(stripped) else None
|
||||
|
||||
|
||||
def same_header(name: str, other: str) -> bool:
|
||||
"""Whether two HTTP header names are the same one. They are case-insensitive (RFC 7230 3.2)."""
|
||||
return name.lower() == other.lower()
|
||||
|
||||
|
||||
def has_header(headers: Mapping[str, str] | None, name: str) -> bool:
|
||||
"""Whether ``headers`` carries ``name`` under any casing."""
|
||||
return bool(headers) and any(same_header(key, name) for key in headers or {})
|
||||
|
||||
|
||||
def without_header(headers: Mapping[str, str] | None, name: str) -> dict[str, str] | None:
|
||||
"""A copy of ``headers`` with every casing of ``name`` removed, or None if nothing remains.
|
||||
|
||||
The one owner of "drop this credential's header". Both MCP stacks and the upstream-credential
|
||||
resolver share it so a slot can never be dropped case-sensitively in one place and
|
||||
case-insensitively in another, which is how an injected header came to shadow a resolved
|
||||
credential on the v1 path.
|
||||
"""
|
||||
if not headers:
|
||||
return None
|
||||
filtered: Final = {key: value for key, value in headers.items() if not same_header(key, name)}
|
||||
return filtered or None
|
||||
|
||||
|
||||
_DEFAULT_PORTS: Final[Mapping[str, int]] = MappingProxyType({"http": 80, "https": 443})
|
||||
|
||||
|
||||
def crosses_origin(configured: str, target: str) -> bool:
|
||||
"""Whether ``target`` leaves ``configured``'s origin, by the rule HTTP clients use.
|
||||
|
||||
Origin is scheme, host and port, not host alone, so a same-host HTTPS downgrade or a port change
|
||||
counts as crossing it. A plain http -> https upgrade of the same host is exempt, matching what
|
||||
httpx exempts when it decides whether to keep ``Authorization`` across a redirect.
|
||||
"""
|
||||
a: Final = urlsplit(configured)
|
||||
b: Final = urlsplit(target)
|
||||
port_a: Final = a.port or _DEFAULT_PORTS.get(a.scheme)
|
||||
port_b: Final = b.port or _DEFAULT_PORTS.get(b.scheme)
|
||||
if a.scheme == b.scheme and a.hostname == b.hostname and port_a == port_b:
|
||||
return False
|
||||
return not (
|
||||
a.hostname == b.hostname and a.scheme == "http" and port_a == 80 and b.scheme == "https" and port_b == 443
|
||||
)
|
||||
|
||||
|
||||
def custom_credential_slot(headers: Mapping[str, str] | None) -> str | None:
|
||||
"""The first header carrying a credential somewhere other than ``Authorization``, if any."""
|
||||
return next((name for name in headers or {} if not same_header(name, DEFAULT_CREDENTIAL_HEADER)), None)
|
||||
|
||||
|
||||
def credential_redirect_hook(
|
||||
configured_url: str, slot: str | None
|
||||
) -> Callable[[httpx.Request], Awaitable[None]] | None:
|
||||
"""An httpx request hook dropping ``slot`` once a redirect leaves ``configured_url``'s origin.
|
||||
|
||||
None when no guard is needed, so callers do not each repeat the exemption: HTTP clients already
|
||||
strip ``Authorization`` across origins, but forward every other header, so only a credential an
|
||||
operator moved to its own slot can be replayed to whatever host the upstream redirects to.
|
||||
"""
|
||||
if not configured_url or not slot or same_header(slot, DEFAULT_CREDENTIAL_HEADER):
|
||||
return None
|
||||
|
||||
async def guard(request: httpx.Request) -> None:
|
||||
if slot in request.headers and crosses_origin(configured_url, str(request.url)):
|
||||
del request.headers[slot]
|
||||
|
||||
return guard
|
||||
|
||||
|
||||
MCP_ADMIN_CONFIG_CREDENTIAL_KEYS: Final[tuple[str, ...]] = ("upstream_resource", "upstream_token_header")
|
||||
"""Non-secret credential keys returned on read so the admin form can show and clear them. Mirrors
|
||||
``ADMIN_CONFIG_CREDENTIAL_KEYS`` in ``ui/litellm-dashboard/src/components/mcp_tools/types.tsx``."""
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from pydantic import BaseModel, ConfigDict, field_validator
|
||||
|
||||
from litellm.types.mcp import (
|
||||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
|
|
@ -9,6 +9,7 @@ from litellm.types.mcp import (
|
|||
MCPAuthType,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPTransportType,
|
||||
normalize_upstream_header_name,
|
||||
)
|
||||
|
||||
# MCPInfo now allows arbitrary additional fields for custom metadata
|
||||
|
|
@ -86,6 +87,22 @@ class MCPServer(BaseModel):
|
|||
# today's behavior; "auto" derives the canonical URI from ``url``; any other value is sent
|
||||
# verbatim. Resolved by ``oauth_utils.resolve_upstream_resource``.
|
||||
upstream_resource: str | None = None
|
||||
# Which upstream header carries the credential LiteLLM resolves for this server (the minted
|
||||
# OAuth token, or the static key). None keeps RFC 6750's default, ``Authorization``. An ESB or
|
||||
# API gateway that terminates its own credential in a private header needs this so a second,
|
||||
# operator-configured ``Authorization`` can pass through to the origin untouched.
|
||||
upstream_token_header: str | None = None
|
||||
|
||||
@field_validator("upstream_token_header")
|
||||
@classmethod
|
||||
def _check_upstream_token_header(cls, value: str | None) -> str | None:
|
||||
if value is None or not value.strip():
|
||||
return None
|
||||
normalized: Final = normalize_upstream_header_name(value)
|
||||
if normalized is None:
|
||||
raise ValueError(f"upstream_token_header must be a valid HTTP header name (RFC 7230 token), got {value!r}")
|
||||
return normalized
|
||||
|
||||
# AWS SigV4 fields
|
||||
aws_access_key_id: str | None = None
|
||||
aws_secret_access_key: str | None = None
|
||||
|
|
|
|||
|
|
@ -3025,8 +3025,7 @@ def register_model(
|
|||
and value.get("cache_read_input_token_cost") is None
|
||||
and value.get("tiered_pricing") is None
|
||||
and (
|
||||
value.get("input_cost_per_token") is not None
|
||||
or value.get("output_cost_per_token") is not None
|
||||
value.get("input_cost_per_token") is not None or value.get("output_cost_per_token") is not None
|
||||
)
|
||||
):
|
||||
verbose_logger.warning(
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -5,8 +5,9 @@ import json
|
|||
from pathlib import Path
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
import tomllib
|
||||
from typing import Dict, List, Optional, Set, Tuple
|
||||
from typing import Callable, Dict, Final, List, Optional, Protocol, Set, Tuple
|
||||
|
||||
from packaging.requirements import Requirement
|
||||
import requests
|
||||
|
|
@ -37,6 +38,13 @@ DEFAULT_TRANSITIVE_PIN_PACKAGES = (
|
|||
# of the identifier, not an operator.
|
||||
_SPDX_OPERATOR_SPLIT = re.compile(r"\s+(?:OR|AND)\s+")
|
||||
_SPDX_WITH_SUFFIX = re.compile(r"\s+WITH\s+.*", re.DOTALL)
|
||||
_PYPI_FETCH_ATTEMPTS: Final[int] = 3
|
||||
_PYPI_FETCH_BACKOFF_SECONDS: Final[float] = 0.5
|
||||
|
||||
|
||||
class _HttpGet(Protocol):
|
||||
def __call__(self, url: str, *, timeout: float) -> requests.Response:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -50,7 +58,10 @@ class PackageLicense:
|
|||
|
||||
class LicenseChecker:
|
||||
def __init__(
|
||||
self, config_file: Path = Path("./tests/code_coverage_tests/liccheck.ini")
|
||||
self,
|
||||
config_file: Path = Path("./tests/code_coverage_tests/liccheck.ini"),
|
||||
http_get: Optional[_HttpGet] = None,
|
||||
sleep: Optional[Callable[[float], None]] = None,
|
||||
):
|
||||
if not config_file.exists():
|
||||
print(f"Error: Config file {config_file} not found")
|
||||
|
|
@ -79,6 +90,8 @@ class LicenseChecker:
|
|||
|
||||
# Track package results
|
||||
self.package_results: List[PackageLicense] = []
|
||||
self._http_get = http_get
|
||||
self._sleep = sleep
|
||||
|
||||
@staticmethod
|
||||
def _normalize_package_name(package_name: str) -> str:
|
||||
|
|
@ -123,21 +136,38 @@ class LicenseChecker:
|
|||
last resort derives the license from the ``License :: OSI Approved ::
|
||||
...`` trove classifiers.
|
||||
"""
|
||||
try:
|
||||
url = f"https://pypi.org/pypi/{package_name}/{version}/json"
|
||||
response = requests.get(url, timeout=10)
|
||||
response.raise_for_status()
|
||||
info = response.json().get("info", {}) or {}
|
||||
return (
|
||||
info.get("license_expression")
|
||||
or info.get("license")
|
||||
or self._license_from_classifiers(info.get("classifiers") or [])
|
||||
)
|
||||
except Exception as e:
|
||||
print(
|
||||
f"Warning: Failed to fetch license for {package_name} {version}: {str(e)}"
|
||||
)
|
||||
return None
|
||||
url = f"https://pypi.org/pypi/{package_name}/{version}/json"
|
||||
http_get = self._http_get if self._http_get is not None else requests.get
|
||||
sleep = self._sleep if self._sleep is not None else time.sleep
|
||||
|
||||
for attempt in range(_PYPI_FETCH_ATTEMPTS):
|
||||
try:
|
||||
response = http_get(url, timeout=10)
|
||||
response.raise_for_status()
|
||||
info = response.json().get("info", {}) or {}
|
||||
return (
|
||||
info.get("license_expression")
|
||||
or info.get("license")
|
||||
or self._license_from_classifiers(info.get("classifiers") or [])
|
||||
)
|
||||
except Exception as error:
|
||||
if self._is_retryable_pypi_error(error) and attempt < _PYPI_FETCH_ATTEMPTS - 1:
|
||||
sleep(_PYPI_FETCH_BACKOFF_SECONDS)
|
||||
continue
|
||||
print(
|
||||
f"Warning: Failed to fetch license for {package_name} {version}: {str(error)}"
|
||||
)
|
||||
return None
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _is_retryable_pypi_error(error: Exception) -> bool:
|
||||
if isinstance(error, (requests.ConnectionError, requests.Timeout)):
|
||||
return True
|
||||
if not isinstance(error, requests.HTTPError) or error.response is None:
|
||||
return False
|
||||
status_code = error.response.status_code
|
||||
return status_code == 429 or status_code >= 500
|
||||
|
||||
@staticmethod
|
||||
def _license_from_classifiers(classifiers: List[str]) -> Optional[str]:
|
||||
|
|
|
|||
BIN
tests/e2e/llm_translation/fixtures/cat.jpg
Normal file
BIN
tests/e2e/llm_translation/fixtures/cat.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 30 KiB |
|
|
@ -16,7 +16,10 @@ via /model/new (Cohere, Gemini, hosted_vllm), each deleted on teardown.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -79,18 +82,22 @@ def _streamed_tool_call(events: list[str]) -> tuple[str, str]:
|
|||
return name, arguments
|
||||
|
||||
|
||||
CAT_IMAGE_URL = "https://upload.wikimedia.org/wikipedia/commons/3/3a/Cat03.jpg"
|
||||
_FIXTURES_DIR: Final = Path(__file__).parent / "fixtures"
|
||||
CAT_IMAGE: Final = _FIXTURES_DIR / "cat.jpg"
|
||||
OPENAI_VISION_BACKEND = "openai/gpt-4o"
|
||||
|
||||
# OpenAI caches a shared prompt prefix once it exceeds ~1024 tokens; this is well
|
||||
# past that, so a repeat call reports cached prompt tokens.
|
||||
|
||||
def _cat_image_data_url() -> str:
|
||||
return "data:image/jpeg;base64," + base64.b64encode(CAT_IMAGE.read_bytes()).decode()
|
||||
|
||||
|
||||
def _vision_messages() -> list[ChatMessage]:
|
||||
return [
|
||||
ChatMessage(
|
||||
role="user",
|
||||
content=[
|
||||
TextContentPart(text="What animal is in this image? Answer in one word."),
|
||||
ImageContentPart(image_url=ImageUrl(url=CAT_IMAGE_URL)),
|
||||
ImageContentPart(image_url=ImageUrl(url=_cat_image_data_url())),
|
||||
],
|
||||
)
|
||||
]
|
||||
|
|
|
|||
|
|
@ -210,16 +210,24 @@ def _deltas(result: StreamingResponse) -> list[_StreamDelta]:
|
|||
]
|
||||
|
||||
|
||||
def _single_weather_call(message: OutMessage) -> ToolCall:
|
||||
assert message.tool_calls, f"Together dropped the tool call: {message}"
|
||||
assert len(message.tool_calls) == 1, f"expected one tool call, got {message.tool_calls}"
|
||||
call = message.tool_calls[0]
|
||||
def _validated_weather_call_id(call: ToolCall) -> str:
|
||||
assert call.id, f"tool call carries no id, so a tool result cannot answer it: {call}"
|
||||
assert call.function.name == "get_weather", f"wrong tool called: {call}"
|
||||
assert call.function.arguments, f"tool call carries no arguments: {call}"
|
||||
args = _WeatherArgs.model_validate_json(call.function.arguments)
|
||||
assert "paris" in args.location.lower(), f"tool arguments lost the location: {args}"
|
||||
return call
|
||||
return call.id
|
||||
|
||||
|
||||
def _weather_call_ids(message: OutMessage) -> tuple[str, ...]:
|
||||
"""The id of every tool call the model made, each one checked for the fields a
|
||||
caller needs to answer it. The backend is whichever together_ai row is cheapest
|
||||
with tools and reasoning, and those rows carry supports_parallel_function_calling,
|
||||
so one weather prompt can legitimately come back as several get_weather calls.
|
||||
What the gateway owes us is that each call survives translation intact; how many
|
||||
the model chose to make is the model's business."""
|
||||
assert message.tool_calls, f"Together dropped the tool call: {message}"
|
||||
return tuple(_validated_weather_call_id(call) for call in message.tool_calls)
|
||||
|
||||
|
||||
def _weather_call(client: PassthroughClient, key: str, model: str) -> OutMessage:
|
||||
|
|
@ -289,7 +297,7 @@ class TestTogetherChatCompletions:
|
|||
self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str
|
||||
) -> None:
|
||||
model, key = _register(client, resources, reasoning_tool_backend)
|
||||
_single_weather_call(_weather_call(client, key, model))
|
||||
_ = _weather_call_ids(_weather_call(client, key, model))
|
||||
|
||||
@pytest.mark.covers("llm.chat_completions.together_ai.tool_use.stream.works")
|
||||
def test_tool_call_is_streamed(
|
||||
|
|
@ -328,8 +336,7 @@ class TestTogetherChatCompletions:
|
|||
) -> None:
|
||||
model, key = _register(client, resources, reasoning_tool_backend)
|
||||
first = _weather_call(client, key, model)
|
||||
call = _single_weather_call(first)
|
||||
assert call.id is not None
|
||||
call_ids = _weather_call_ids(first)
|
||||
|
||||
answer = _message(
|
||||
unwrap(
|
||||
|
|
@ -344,7 +351,10 @@ class TestTogetherChatCompletions:
|
|||
reasoning_content=first.reasoning_content,
|
||||
tool_calls=first.tool_calls,
|
||||
),
|
||||
ChatToolResultTurn(tool_call_id=call.id, content=WEATHER_REPORT),
|
||||
*(
|
||||
ChatToolResultTurn(tool_call_id=call_id, content=WEATHER_REPORT)
|
||||
for call_id in call_ids
|
||||
),
|
||||
],
|
||||
tools=[WEATHER_TOOL],
|
||||
max_tokens=512,
|
||||
|
|
@ -470,9 +480,21 @@ def _tool_use_blocks(content: list[AnthropicContentBlock] | None) -> list[Anthro
|
|||
return [block for block in content if block.type == "tool_use"]
|
||||
|
||||
|
||||
def _validated_tool_use_id(block: AnthropicContentBlock) -> str:
|
||||
assert block.name == "get_weather", f"wrong tool called: {block}"
|
||||
assert block.id, f"tool_use block carries no id, so a tool_result cannot answer it: {block}"
|
||||
assert block.input is not None, f"tool_use block carries no input: {block}"
|
||||
args = _WeatherArgs.model_validate(block.input)
|
||||
assert "paris" in args.location.lower(), f"tool input lost the location: {args}"
|
||||
return block.id
|
||||
|
||||
|
||||
def _messages_weather_call(
|
||||
client: PassthroughClient, key: str, model: str
|
||||
) -> tuple[list[AnthropicContentBlock], AnthropicContentBlock]:
|
||||
) -> tuple[list[AnthropicContentBlock], tuple[str, ...]]:
|
||||
"""The blocks /v1/messages returned and the id of every tool_use among them. The
|
||||
count is the model's choice (see _weather_call_ids); what this surface owes us is
|
||||
that each tool_use arrives named and addressable."""
|
||||
response = unwrap(
|
||||
client.proxy.messages(
|
||||
key,
|
||||
|
|
@ -485,12 +507,9 @@ def _messages_weather_call(
|
|||
)
|
||||
)
|
||||
tool_uses = _tool_use_blocks(response.content)
|
||||
assert len(tool_uses) == 1, f"expected one tool_use block, got {response.content}"
|
||||
block = tool_uses[0]
|
||||
assert block.name == "get_weather", f"wrong tool called: {block}"
|
||||
assert block.id, f"tool_use block carries no id, so a tool_result cannot answer it: {block}"
|
||||
assert tool_uses, f"/v1/messages carried no tool_use block: {response.content}"
|
||||
assert response.content is not None
|
||||
return response.content, block
|
||||
return response.content, tuple(_validated_tool_use_id(block) for block in tool_uses)
|
||||
|
||||
|
||||
class TestTogetherMessages:
|
||||
|
|
@ -506,8 +525,7 @@ class TestTogetherMessages:
|
|||
self, client: PassthroughClient, resources: ResourceManager, reasoning_tool_backend: str
|
||||
) -> None:
|
||||
model, key = _register(client, resources, reasoning_tool_backend)
|
||||
first_content, block = _messages_weather_call(client, key, model)
|
||||
assert block.id is not None
|
||||
first_content, tool_use_ids = _messages_weather_call(client, key, model)
|
||||
|
||||
response = unwrap(
|
||||
client.proxy.messages(
|
||||
|
|
@ -520,7 +538,10 @@ class TestTogetherMessages:
|
|||
ChatMessage(role="user", content=WEATHER_PROMPT),
|
||||
AnthropicAssistantTurn(content=first_content),
|
||||
AnthropicToolResultTurn(
|
||||
content=[AnthropicToolResultBlock(tool_use_id=block.id, content=WEATHER_REPORT)]
|
||||
content=[
|
||||
AnthropicToolResultBlock(tool_use_id=tool_use_id, content=WEATHER_REPORT)
|
||||
for tool_use_id in tool_use_ids
|
||||
]
|
||||
),
|
||||
],
|
||||
),
|
||||
|
|
|
|||
|
|
@ -421,6 +421,7 @@ class AnthropicContentBlock(BaseModel):
|
|||
text: str | None = None
|
||||
id: str | None = None
|
||||
name: str | None = None
|
||||
input: dict[str, object] | None = None
|
||||
|
||||
|
||||
class AnthropicToolResultBlock(BaseModel):
|
||||
|
|
|
|||
|
|
@ -23,26 +23,18 @@ class TestTogetherAI(BaseLLMChatTest):
|
|||
pass
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, expected_bool",
|
||||
"model",
|
||||
[
|
||||
("meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo", True),
|
||||
("nvidia/Llama-3.1-Nemotron-70B-Instruct-HF", False),
|
||||
"meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo",
|
||||
"nvidia/Llama-3.1-Nemotron-70B-Instruct-HF",
|
||||
],
|
||||
)
|
||||
def test_get_supported_response_format_together_ai(
|
||||
self, model: str, expected_bool: bool
|
||||
) -> None:
|
||||
def test_get_supported_response_format_together_ai(self, model: str) -> None:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
optional_params = litellm.get_supported_openai_params(
|
||||
model, custom_llm_provider="together_ai"
|
||||
)
|
||||
# Mapped provider
|
||||
assert isinstance(optional_params, list)
|
||||
|
||||
if expected_bool:
|
||||
assert "response_format" in optional_params
|
||||
assert "tools" in optional_params
|
||||
else:
|
||||
assert "response_format" not in optional_params
|
||||
assert "tools" not in optional_params
|
||||
assert "response_format" in optional_params
|
||||
assert "tools" in optional_params
|
||||
|
|
|
|||
|
|
@ -45,6 +45,22 @@ def setup_and_teardown():
|
|||
asyncio.set_event_loop(None) # Remove the reference to the loop
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
async def drain_logging_worker():
|
||||
"""
|
||||
The logging queue is bound to the running loop, so anything left queued when a test's loop
|
||||
goes away is carried onto the next test's loop and fires against its callbacks.
|
||||
"""
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
yield
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.clear_queue(), timeout=10)
|
||||
except asyncio.TimeoutError:
|
||||
pass
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
# Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests
|
||||
custom_logger_tests = [
|
||||
|
|
|
|||
|
|
@ -3903,3 +3903,39 @@ def test_stored_reasoning_items_win_over_thinking_blocks():
|
|||
reasoning_items = [item for item in input_items if item.get("type") == "reasoning"]
|
||||
assert len(reasoning_items) == 1
|
||||
assert reasoning_items[0]["id"] == "rs_real"
|
||||
|
||||
|
||||
def test_convert_chat_completion_messages_to_responses_api_tool_result_with_tool_reference():
|
||||
"""Tool-search tool_reference blocks have no Responses API equivalent: skip them, never stringify them."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {"name": "ToolSearch", "arguments": '{"query": "web"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_abc123",
|
||||
"content": [
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
{"type": "text", "text": "1 tool found"},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
response, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
function_call_output = next(item for item in response if item.get("type") == "function_call_output")
|
||||
assert function_call_output["output"] == [{"type": "input_text", "text": "1 tool found"}]
|
||||
|
|
|
|||
|
|
@ -375,6 +375,9 @@ def isolate_litellm_state():
|
|||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
image_handling_module.in_memory_cache.flush_cache()
|
||||
_reset_module_level_aws_auth_caches()
|
||||
# litellm.get_model_info() memoizes ModelInfo built from litellm.model_cost, so a
|
||||
# test that rebinds the cost map leaves later tests pricing against the old map.
|
||||
litellm_utils_module._invalidate_model_cost_lowercase_map()
|
||||
|
||||
# Clear all callback lists to prevent cross-test contamination
|
||||
if hasattr(litellm, "callbacks"):
|
||||
|
|
@ -418,6 +421,7 @@ def isolate_litellm_state():
|
|||
|
||||
litellm_utils_module._runtime_registered_model_cost.clear()
|
||||
litellm_utils_module._runtime_registered_model_cost.update(original_runtime_registered_model_cost)
|
||||
litellm_utils_module._invalidate_model_cost_lowercase_map()
|
||||
|
||||
for _router in tuple(litellm_router_module._live_routers):
|
||||
litellm_router_module._live_routers.discard(_router)
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import anyio
|
||||
import httpx
|
||||
import pytest
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
|
||||
from mcp import McpError
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.types import (
|
||||
|
|
@ -1095,3 +1096,188 @@ def test_mcp_extra_matches_proxy_extra_and_supports_streamable_http():
|
|||
specifier = Requirement(mcp_extra[0]).specifier
|
||||
assert not specifier.contains("1.23.0")
|
||||
assert specifier.contains("1.28.1")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, default_header",
|
||||
[
|
||||
(MCPAuth.oauth2, "Authorization"),
|
||||
(MCPAuth.bearer_token, "Authorization"),
|
||||
(MCPAuth.api_key, "X-API-Key"),
|
||||
],
|
||||
)
|
||||
def test_v1_auth_headers_default_to_the_auth_type_slot(auth_type: MCPAuth, default_header: str) -> None:
|
||||
client = MCPClient(server_url="http://up.example.com/mcp", auth_type=auth_type)
|
||||
client.update_auth_value("tok")
|
||||
assert default_header in client._get_auth_headers()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2, MCPAuth.bearer_token, MCPAuth.api_key])
|
||||
def test_v1_auth_headers_honor_the_configured_slot(auth_type: MCPAuth) -> None:
|
||||
"""The v1 stack mints its own client_credentials token (oauth2_token_cache) and writes it here,
|
||||
so leaving this table hardcoded makes the knob a silent no-op for every server that resolves
|
||||
through v1 rather than the v2 resolver."""
|
||||
client = MCPClient(
|
||||
server_url="http://up.example.com/mcp",
|
||||
auth_type=auth_type,
|
||||
auth_header_name="esb-oauth",
|
||||
)
|
||||
client.update_auth_value("tok")
|
||||
headers = client._get_auth_headers()
|
||||
assert "esb-oauth" in headers
|
||||
assert "Authorization" not in headers
|
||||
assert "X-API-Key" not in headers
|
||||
|
||||
|
||||
def test_v1_static_headers_still_win_their_own_slot():
|
||||
# extra_headers (which carries static_headers) is applied last on the v1 path, so a static
|
||||
# Authorization survives untouched while the resolved credential sits on its own header.
|
||||
client = MCPClient(
|
||||
server_url="http://up.example.com/mcp",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
auth_header_name="esb-oauth",
|
||||
extra_headers={"Authorization": "Bearer static-upstream-mcp-token"},
|
||||
)
|
||||
client.update_auth_value("minted")
|
||||
headers = client._get_auth_headers()
|
||||
assert headers["esb-oauth"] == "Bearer minted"
|
||||
assert headers["Authorization"] == "Bearer static-upstream-mcp-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_custom_credential_header_is_stripped_when_a_redirect_crosses_origin():
|
||||
"""httpx drops Authorization across origins but keeps every other header, so a credential the
|
||||
operator moved to its own slot would be replayed to whatever host the upstream redirects to.
|
||||
Verified against real httpx redirect handling, not a hand-built request.
|
||||
"""
|
||||
seen: "list[tuple[str, str]]" = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen.append((request.url.host, request.headers.get("esb-oauth", "<stripped>")))
|
||||
if request.url.host == "upstream.example.com":
|
||||
return httpx.Response(302, headers={"Location": "https://attacker.example.com/collect"})
|
||||
return httpx.Response(200)
|
||||
|
||||
client = MCPClient(
|
||||
server_url="https://upstream.example.com/mcp",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
auth_header_name="esb-oauth",
|
||||
)
|
||||
client.update_auth_value("minted-token")
|
||||
factory = client._create_httpx_client_factory()
|
||||
async with factory(headers=client._get_auth_headers(), timeout=None) as http_client:
|
||||
http_client._transport = httpx.MockTransport(handler)
|
||||
await http_client.get("https://upstream.example.com/mcp")
|
||||
|
||||
assert seen[0] == ("upstream.example.com", "Bearer minted-token")
|
||||
assert seen[1] == ("attacker.example.com", "<stripped>")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorization_is_left_to_httpx_and_needs_no_guard():
|
||||
# The default slot is already protected by httpx, so the client must not install a guard for it
|
||||
# and must not interfere with the ordinary Authorization path.
|
||||
url = "https://upstream.example.com/mcp"
|
||||
from litellm.types.mcp import credential_redirect_hook
|
||||
|
||||
def guard_for(client: MCPClient):
|
||||
return credential_redirect_hook(client.server_url, client._credential_slot)
|
||||
|
||||
assert guard_for(MCPClient(server_url=url, auth_type=MCPAuth.oauth2)) is None
|
||||
assert guard_for(MCPClient(server_url=url, resolved_auth=StaticHeaderAuth("Bearer x"))) is None
|
||||
# a v2 resolver slot is discovered from the auth object, without the caller naming it again
|
||||
custom = MCPClient(server_url=url, resolved_auth=StaticHeaderAuth("Bearer x", header_name="esb-oauth"))
|
||||
assert guard_for(custom) is not None
|
||||
# and the same answer arrives via the v1 configured slot
|
||||
assert guard_for(MCPClient(server_url=url, auth_header_name="ESB-OAuth")) is not None
|
||||
|
||||
|
||||
def test_an_injected_header_cannot_shadow_the_configured_credential_slot():
|
||||
"""The v2 path drops a colliding injected header so the resolved credential wins its slot. The
|
||||
v1 path applies extra_headers last, so without this it silently sends the injected value and the
|
||||
upstream rejects a credential the gateway thought it had sent.
|
||||
"""
|
||||
client = MCPClient(
|
||||
server_url="https://upstream.example.com/mcp",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
auth_header_name="esb-oauth",
|
||||
extra_headers={"esb-oauth": "Bearer injected", "X-Trace": "keep"},
|
||||
)
|
||||
client.update_auth_value("minted-token")
|
||||
headers = client._get_auth_headers()
|
||||
assert headers["esb-oauth"] == "Bearer minted-token"
|
||||
assert headers["X-Trace"] == "keep"
|
||||
|
||||
|
||||
def test_without_a_configured_slot_the_existing_precedence_is_unchanged():
|
||||
# extra_headers winning over authentication_token is long-standing v1 behavior; the fix above
|
||||
# must apply only to the slot the operator explicitly named.
|
||||
client = MCPClient(
|
||||
server_url="https://upstream.example.com/mcp",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
extra_headers={"Authorization": "Bearer injected"},
|
||||
)
|
||||
client.update_auth_value("minted-token")
|
||||
assert client._get_auth_headers()["Authorization"] == "Bearer injected"
|
||||
|
||||
|
||||
_REDIRECT_CASES = [
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com/other"), # same origin
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com:443/other"), # explicit default port
|
||||
("https://upstream.example.com/mcp", "https://attacker.example.com/collect"), # different host
|
||||
("https://upstream.example.com/mcp", "http://upstream.example.com/collect"), # scheme downgrade
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com:8443/other"), # different port
|
||||
("https://upstream.example.com/mcp", "https://sub.upstream.example.com/x"), # different host
|
||||
("http://upstream.example.com/mcp", "https://upstream.example.com/other"), # http -> https upgrade
|
||||
("http://upstream.example.com/mcp", "http://upstream.example.com/other"), # same origin, plain http
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("start,target", _REDIRECT_CASES)
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_guard_agrees_with_httpx_about_authorization(start: str, target: str) -> None:
|
||||
"""Our custom slot must be dropped on exactly the redirects where httpx drops Authorization.
|
||||
|
||||
The rule is mirrored rather than imported, so this drives real httpx and compares the two
|
||||
outcomes. A future httpx that changes its redirect rule reds here instead of silently leaving
|
||||
the custom slot forwarded where Authorization is not (or stripped where it is not needed).
|
||||
"""
|
||||
seen: "list[tuple[str, str, str]]" = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen.append(
|
||||
(
|
||||
str(request.url),
|
||||
request.headers.get("authorization", "<stripped>"),
|
||||
request.headers.get("esb-oauth", "<stripped>"),
|
||||
)
|
||||
)
|
||||
if str(request.url) == start:
|
||||
return httpx.Response(302, headers={"Location": target})
|
||||
return httpx.Response(200)
|
||||
|
||||
client = MCPClient(server_url=start, auth_type=MCPAuth.oauth2, auth_header_name="esb-oauth")
|
||||
factory = client._create_httpx_client_factory()
|
||||
async with factory(headers={"Authorization": "Bearer AUTH", "esb-oauth": "Bearer ESB"}, timeout=None) as http:
|
||||
http._transport = httpx.MockTransport(handler)
|
||||
await http.get(start)
|
||||
|
||||
_url, authorization, esb = seen[-1]
|
||||
assert (authorization == "<stripped>") == (esb == "<stripped>"), (
|
||||
f"httpx and the guard disagree for {target}: authorization={authorization!r} esb-oauth={esb!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_a_differently_cased_injected_header_cannot_shadow_the_slot() -> None:
|
||||
# HTTP header names are case-insensitive and v2 drops the collision case-insensitively, so an
|
||||
# exact-key check here would leave both spellings in the dict and let the injected value win.
|
||||
client = MCPClient(
|
||||
server_url="https://upstream.example.com/mcp",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
auth_header_name="esb-oauth",
|
||||
extra_headers={"ESB-OAuth": "Bearer injected", "X-Trace": "keep"},
|
||||
)
|
||||
client.update_auth_value("minted-token")
|
||||
headers = client._get_auth_headers()
|
||||
assert [v for k, v in headers.items() if k.lower() == "esb-oauth"] == ["Bearer minted-token"]
|
||||
assert headers["X-Trace"] == "keep"
|
||||
|
|
|
|||
|
|
@ -778,11 +778,15 @@ def test_mcp_span_roots_without_transport_or_propagated_context(
|
|||
|
||||
|
||||
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
|
||||
def test_mcp_span_parents_to_propagated_meta_trace_context(make_payload, span_name):
|
||||
def test_mcp_span_links_propagated_meta_trace_context_and_nests_under_transport(
|
||||
make_payload, span_name
|
||||
):
|
||||
"""When the client propagates W3C trace context in the request's
|
||||
``params._meta`` (SEP-414), the MCP span parents to it (one distributed trace)
|
||||
and still links the transport span — never falling through to the
|
||||
ambient/session span."""
|
||||
``params._meta`` (SEP-414), the MCP span still nests under the gateway's own
|
||||
transport span — one renderable trace — and records the client's context as a
|
||||
span *link*. Parenting to the remote context instead would root the span in a
|
||||
trace whose root span never reaches the gateway's tracing backend, leaving the
|
||||
span unreachable from the trace view."""
|
||||
logger, exporter = _logger()
|
||||
transport = logger._emitter.start_span(
|
||||
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
|
|
@ -801,12 +805,65 @@ def test_mcp_span_parents_to_propagated_meta_trace_context(make_payload, span_na
|
|||
reset_mcp_message_trace_carrier(token)
|
||||
transport.end()
|
||||
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
|
||||
assert span.context.trace_id == 0x11111111111111111111111111111111
|
||||
assert span.parent is not None
|
||||
assert span.parent.span_id == 0x2222222222222222
|
||||
assert [link.context.span_id for link in span.links] == [
|
||||
transport.get_span_context().span_id
|
||||
assert span.parent.span_id == transport.get_span_context().span_id
|
||||
assert span.context.trace_id == transport.get_span_context().trace_id
|
||||
assert [link.context.trace_id for link in span.links] == [
|
||||
0x11111111111111111111111111111111
|
||||
]
|
||||
assert [link.context.span_id for link in span.links] == [0x2222222222222222]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
|
||||
def test_mcp_span_without_transport_roots_and_links_propagated_context(
|
||||
make_payload, span_name
|
||||
):
|
||||
"""With no transport span at all there is nothing of the gateway's to anchor
|
||||
to, so the span starts its own root trace — and the client context stays a
|
||||
span link there too, so the event keeps one shape everywhere."""
|
||||
logger, exporter = _logger()
|
||||
token = set_mcp_message_trace_carrier(
|
||||
{"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"}
|
||||
)
|
||||
try:
|
||||
asyncio.run(
|
||||
logger.async_log_success_event(
|
||||
{"standard_logging_object": make_payload()}, None, None, None
|
||||
)
|
||||
)
|
||||
finally:
|
||||
reset_mcp_message_trace_carrier(token)
|
||||
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
|
||||
assert span.parent is None
|
||||
assert span.context.trace_id != 0x11111111111111111111111111111111
|
||||
assert [link.context.span_id for link in span.links] == [0x2222222222222222]
|
||||
|
||||
|
||||
def test_mcp_span_links_unsampled_client_traceparent():
|
||||
"""A client traceparent with the sampled flag off ('-00') still yields a valid
|
||||
remote context, so the link is recorded; the span's own recording follows the
|
||||
transport's sampling decision, never the client's flag."""
|
||||
logger, exporter = _logger()
|
||||
transport = logger._emitter.start_span(
|
||||
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
)
|
||||
set_request_root_span(transport)
|
||||
token = set_mcp_message_trace_carrier(
|
||||
{"traceparent": "00-11111111111111111111111111111111-2222222222222222-00"}
|
||||
)
|
||||
try:
|
||||
asyncio.run(
|
||||
logger.async_log_success_event(
|
||||
{"standard_logging_object": _mcp_list_payload()}, None, None, None
|
||||
)
|
||||
)
|
||||
finally:
|
||||
reset_mcp_message_trace_carrier(token)
|
||||
transport.end()
|
||||
span = next(s for s in exporter.get_finished_spans() if s.name == "tools/list")
|
||||
assert span.parent is not None
|
||||
assert span.parent.span_id == transport.get_span_context().span_id
|
||||
assert [link.context.span_id for link in span.links] == [0x2222222222222222]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
|
||||
|
|
@ -839,8 +896,11 @@ def test_mcp_span_ignores_client_supplied_baggage(make_payload, span_name):
|
|||
reset_mcp_message_trace_carrier(token)
|
||||
transport.end()
|
||||
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
|
||||
# Trace context still honored: proves the carrier was processed, not dropped wholesale.
|
||||
assert span.parent is not None and span.parent.span_id == 0x2222222222222222
|
||||
# Trace context still honored (as a link): proves the carrier was processed,
|
||||
# not dropped wholesale.
|
||||
assert [link.context.span_id for link in span.links] == [0x2222222222222222]
|
||||
assert span.parent is not None
|
||||
assert span.parent.span_id == transport.get_span_context().span_id
|
||||
# Identity is the authenticated payload's team, never the client's spoofed value.
|
||||
assert span.attributes[LiteLLM.TEAM_ID] == "t1"
|
||||
assert "litellm.metadata.user_api_key_user_id" not in span.attributes
|
||||
|
|
@ -888,10 +948,10 @@ def test_mcp_span_malformed_traceparent_nests_under_transport():
|
|||
assert span.links == ()
|
||||
|
||||
|
||||
def test_mcp_span_links_this_messages_transport_when_context_is_propagated():
|
||||
"""On the semconv path the transport is recorded as a link, and that link must
|
||||
point at the POST carrying this message too. Reading the stale session anchor
|
||||
would attribute the tool call to whichever request opened the session."""
|
||||
def test_mcp_span_with_propagated_context_nests_under_this_messages_transport():
|
||||
"""With client context propagated, the span must still anchor to the POST
|
||||
carrying this message, not the stale session anchor — otherwise the tool call
|
||||
is attributed to whichever request opened the session."""
|
||||
logger, exporter = _logger()
|
||||
session_opener = logger._emitter.start_span(
|
||||
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
|
|
@ -916,10 +976,10 @@ def test_mcp_span_links_this_messages_transport_when_context_is_propagated():
|
|||
session_opener.end()
|
||||
this_message.end()
|
||||
span = next(s for s in exporter.get_finished_spans() if s.name == "tools/list")
|
||||
assert span.parent is not None and span.parent.span_id == 0x2222222222222222
|
||||
assert [link.context.span_id for link in span.links] == [
|
||||
this_message.get_span_context().span_id
|
||||
]
|
||||
assert span.parent is not None
|
||||
assert span.parent.span_id == this_message.get_span_context().span_id
|
||||
assert span.context.trace_id == this_message.get_span_context().trace_id
|
||||
assert [link.context.span_id for link in span.links] == [0x2222222222222222]
|
||||
|
||||
|
||||
def test_pre_call_idempotent_keeps_first_span():
|
||||
|
|
|
|||
|
|
@ -107,32 +107,29 @@ def test_registry_parent_integrity_no_orphans():
|
|||
|
||||
|
||||
def test_registry_hierarchy_shape():
|
||||
# MCP roles have no in-process parent: per the MCP semconv they root (or adopt
|
||||
# the client's propagated _meta context), so they sit alongside PROXY_REQUEST.
|
||||
assert set(root_roles()) == {
|
||||
SpanRole.PROXY_REQUEST,
|
||||
SpanRole.MCP_TOOL_CALL,
|
||||
SpanRole.MCP_LIST_TOOLS,
|
||||
}
|
||||
assert set(root_roles()) == {SpanRole.PROXY_REQUEST}
|
||||
# Guardrails parent to the request span, not the LLM call: a pre-call
|
||||
# guardrail runs before the LLM call exists, so it's a sibling of it.
|
||||
# guardrail runs before the LLM call exists, so it's a sibling of it. MCP
|
||||
# spans nest under the transport span of the request carrying that message.
|
||||
assert set(child_roles(SpanRole.PROXY_REQUEST)) == {
|
||||
SpanRole.LLM_CALL,
|
||||
SpanRole.GUARDRAIL,
|
||||
SpanRole.DB_CALL,
|
||||
SpanRole.SERVICE,
|
||||
SpanRole.MCP_TOOL_CALL,
|
||||
SpanRole.MCP_LIST_TOOLS,
|
||||
}
|
||||
assert SPAN_REGISTRY[SpanRole.LLM_CALL].kind is LiteLLMSpanKind.CLIENT
|
||||
# The proxy is an MCP client to the upstream tool server: CLIENT span. Listing
|
||||
# tools is the same client relationship, so it's a CLIENT span too.
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].kind is LiteLLMSpanKind.CLIENT
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].kind is LiteLLMSpanKind.CLIENT
|
||||
# MCP spans don't nest under the transport: they link the PROXY_REQUEST span
|
||||
# instead of parenting to it (OTel GenAI MCP semconv).
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].parent is None
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].parent is None
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].links is SpanRole.PROXY_REQUEST
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].links is SpanRole.PROXY_REQUEST
|
||||
# MCP spans nest under the transport span of the request carrying that
|
||||
# message (resolved per message at emit time); a client-propagated context
|
||||
# becomes a span link to that remote context, which is not a registry role
|
||||
# (SpanSpec declares no link field at all).
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].parent is SpanRole.PROXY_REQUEST
|
||||
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].parent is SpanRole.PROXY_REQUEST
|
||||
assert SPAN_REGISTRY[SpanRole.PROXY_REQUEST].kind is LiteLLMSpanKind.SERVER
|
||||
assert SPAN_REGISTRY[SpanRole.GUARDRAIL].parent is SpanRole.PROXY_REQUEST
|
||||
# An outbound datastore call is a CLIENT span; an internal service is INTERNAL.
|
||||
|
|
|
|||
|
|
@ -0,0 +1,134 @@
|
|||
from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
||||
SubtitleToken,
|
||||
render_subtitle_tokens_as_srt,
|
||||
render_subtitle_tokens_as_vtt,
|
||||
synthesize_subtitle_document,
|
||||
)
|
||||
|
||||
|
||||
class TestRenderSubtitleTokensAsSrt:
|
||||
def test_single_cue_full_document(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,000\nHello world.\n"
|
||||
|
||||
def test_speaker_change_starts_a_new_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker="spk:0"),
|
||||
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker="spk:1"),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:01,000\nHi.\n\n2\n00:00:01,500 --> 00:00:02,500\nHey.\n"
|
||||
)
|
||||
|
||||
def test_token_cap_starts_a_new_cue_after_15_tokens(self):
|
||||
tokens = tuple(
|
||||
SubtitleToken(text=f"{index} ", start_ms=index * 100, end_ms=index * 100 + 100) for index in range(16)
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:01,500\n0 1 2 3 4 5 6 7 8 9 10 11 12 13 14\n"
|
||||
"\n2\n00:00:01,500 --> 00:00:01,600\n15\n"
|
||||
)
|
||||
|
||||
def test_duration_cap_starts_a_new_cue_at_5000ms(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Alpha ", start_ms=0, end_ms=400),
|
||||
SubtitleToken(text="beta ", start_ms=2000, end_ms=2400),
|
||||
SubtitleToken(text="gamma.", start_ms=5000, end_ms=5400),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == (
|
||||
"1\n00:00:00,000 --> 00:00:02,400\nAlpha beta\n\n2\n00:00:05,000 --> 00:00:05,400\ngamma.\n"
|
||||
)
|
||||
|
||||
def test_timestampless_token_joins_the_current_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="there "),
|
||||
SubtitleToken(text="world.", start_ms=900, end_ms=1300),
|
||||
)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:01,300\nHello there world.\n"
|
||||
|
||||
def test_only_timestampless_tokens_renders_empty(self):
|
||||
assert render_subtitle_tokens_as_srt((SubtitleToken(text="no timestamps"),)) == ""
|
||||
|
||||
def test_empty_tokens_render_empty(self):
|
||||
assert render_subtitle_tokens_as_srt(()) == ""
|
||||
|
||||
def test_timestamps_past_one_hour(self):
|
||||
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n01:01:01,001 --> 01:01:02,002\nLate.\n"
|
||||
|
||||
def test_negative_timestamps_clamp_to_zero(self):
|
||||
tokens = (SubtitleToken(text="Early.", start_ms=-100, end_ms=-50),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:00,000 --> 00:00:00,000\nEarly.\n"
|
||||
|
||||
def test_missing_end_falls_back_to_cue_start(self):
|
||||
tokens = (SubtitleToken(text="Open.", start_ms=1200),)
|
||||
assert render_subtitle_tokens_as_srt(tokens) == "1\n00:00:01,200 --> 00:00:01,200\nOpen.\n"
|
||||
|
||||
|
||||
class TestRenderSubtitleTokensAsVtt:
|
||||
def test_single_cue_full_document(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hello ", start_ms=0, end_ms=500),
|
||||
SubtitleToken(text="world.", start_ms=500, end_ms=1000),
|
||||
)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHello world.\n"
|
||||
|
||||
def test_empty_tokens_render_header_only(self):
|
||||
assert render_subtitle_tokens_as_vtt(()) == "WEBVTT\n"
|
||||
|
||||
def test_timestamps_past_one_hour_use_dot_separator(self):
|
||||
tokens = (SubtitleToken(text="Late.", start_ms=3_661_001, end_ms=3_662_002),)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == "WEBVTT\n\n01:01:01.001 --> 01:01:02.002\nLate.\n"
|
||||
|
||||
def test_speaker_change_starts_a_new_cue(self):
|
||||
tokens = (
|
||||
SubtitleToken(text="Hi.", start_ms=0, end_ms=1000, speaker=1),
|
||||
SubtitleToken(text="Hey.", start_ms=1500, end_ms=2500, speaker=2),
|
||||
)
|
||||
assert render_subtitle_tokens_as_vtt(tokens) == (
|
||||
"WEBVTT\n\n00:00:00.000 --> 00:00:01.000\nHi.\n\n00:00:01.500 --> 00:00:02.500\nHey.\n"
|
||||
)
|
||||
|
||||
|
||||
class TestSynthesizeSubtitleDocument:
|
||||
WORDS = [
|
||||
{"word": "Four", "start": 0.4, "end": 0.7, "speaker": "spk:0"},
|
||||
{"word": "score", "start": 0.7, "end": 1.1, "speaker": "spk:0"},
|
||||
]
|
||||
|
||||
def test_srt_from_words_converts_seconds_to_milliseconds(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "srt") == "1\n00:00:00,400 --> 00:00:01,100\nFour score\n"
|
||||
|
||||
def test_vtt_from_words_converts_seconds_to_milliseconds(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "vtt") == (
|
||||
"WEBVTT\n\n00:00:00.400 --> 00:00:01.100\nFour score\n"
|
||||
)
|
||||
|
||||
def test_speaker_change_splits_cues(self):
|
||||
words = [
|
||||
{"word": "Hi", "start": 0.0, "end": 0.5, "speaker": "spk:0"},
|
||||
{"word": "Hey", "start": 0.6, "end": 1.0, "speaker": "spk:1"},
|
||||
]
|
||||
assert synthesize_subtitle_document(words, "srt") == (
|
||||
"1\n00:00:00,000 --> 00:00:00,500\nHi\n\n2\n00:00:00,600 --> 00:00:01,000\nHey\n"
|
||||
)
|
||||
|
||||
def test_non_subtitle_format_returns_none(self):
|
||||
assert synthesize_subtitle_document(self.WORDS, "verbose_json") is None
|
||||
assert synthesize_subtitle_document(self.WORDS, "json") is None
|
||||
|
||||
def test_missing_words_returns_none(self):
|
||||
assert synthesize_subtitle_document(None, "srt") is None
|
||||
assert synthesize_subtitle_document([], "srt") is None
|
||||
|
||||
def test_words_without_timestamps_return_none(self):
|
||||
assert synthesize_subtitle_document([{"word": "Hello"}], "srt") is None
|
||||
assert synthesize_subtitle_document([{"word": "Hello"}], "vtt") is None
|
||||
|
||||
def test_malformed_words_return_none(self):
|
||||
assert synthesize_subtitle_document("not words", "srt") is None
|
||||
assert synthesize_subtitle_document([{"word": "ok", "start": "not-a-number"}], "srt") is None
|
||||
|
|
@ -1027,3 +1027,70 @@ def test_update_messages_xlitellm_decode_does_not_override_mapping():
|
|||
updated = update_messages_with_model_file_ids(messages, "model-A", mapping)
|
||||
|
||||
assert updated[0]["content"][0]["file"]["file_id"] == "provider-explicit-id"
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_keeps_text_parts():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
messages = [
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg(
|
||||
[
|
||||
{"type": "text", "text": "WebFetch tool loaded successfully."},
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
]
|
||||
),
|
||||
]
|
||||
|
||||
result = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
|
||||
assert result[1]["content"] == [{"type": "text", "text": "WebFetch tool loaded successfully."}]
|
||||
assert result[1]["tool_call_id"] == "call_1"
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_reference_only_becomes_empty_text():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
messages = [
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg([{"type": "tool_reference", "tool_name": "WebFetch"}]),
|
||||
]
|
||||
|
||||
result = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
|
||||
assert result[1] == {"role": "tool", "tool_call_id": "call_1", "content": ""}
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_without_references_passes_through():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
messages = [
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg([{"type": "text", "text": "plain result"}]),
|
||||
]
|
||||
|
||||
assert drop_tool_reference_parts_from_tool_messages(messages) is messages
|
||||
|
||||
|
||||
def test_drop_tool_reference_parts_leaves_non_tool_messages_alone():
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
)
|
||||
|
||||
user_message = {"role": "user", "content": [{"type": "tool_reference", "tool_name": "WebFetch"}]}
|
||||
messages = [
|
||||
user_message,
|
||||
_assistant_tool_call_msg("call_1"),
|
||||
_tool_msg([{"type": "tool_reference", "tool_name": "WebFetch"}]),
|
||||
]
|
||||
|
||||
result = drop_tool_reference_parts_from_tool_messages(messages)
|
||||
|
||||
assert result[0] == user_message
|
||||
assert result[2]["content"] == ""
|
||||
|
|
|
|||
|
|
@ -3578,3 +3578,52 @@ async def test_bedrock_converse_pdf_only_user_message_gets_text_block_async():
|
|||
assert len(result) == 1
|
||||
assert any("document" in block for block in result[0]["content"])
|
||||
assert _text_blocks(result[0]) == [BEDROCK_DOCUMENT_PLACEHOLDER_TEXT]
|
||||
|
||||
|
||||
def test_convert_to_anthropic_tool_result_keeps_tool_reference_blocks():
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import convert_to_anthropic_tool_result
|
||||
|
||||
result = convert_to_anthropic_tool_result(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "toolu_01",
|
||||
"content": [
|
||||
{"type": "text", "text": "loaded"},
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_01",
|
||||
"content": [
|
||||
{"type": "text", "text": "loaded"},
|
||||
{"type": "tool_reference", "tool_name": "WebFetch"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def test_convert_gemini_tool_call_result_answers_tool_reference_only_result():
|
||||
"""Every Gemini function call needs a function response, even when the tool result carries no text.
|
||||
Fixes: https://github.com/BerriAI/litellm/issues/37462
|
||||
"""
|
||||
result = convert_to_gemini_tool_call_result(
|
||||
message=ChatCompletionToolMessage(
|
||||
role="tool",
|
||||
tool_call_id="toolu_01",
|
||||
content=[{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
),
|
||||
last_message_with_tool_calls={
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "toolu_01",
|
||||
"type": "function",
|
||||
"function": {"name": "ToolSearch", "arguments": '{"query": "select:WebFetch"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
assert result == {"function_response": {"name": "ToolSearch", "response": {"content": ""}}}
|
||||
|
|
|
|||
|
|
@ -4460,3 +4460,51 @@ def test_handle_stream_fallback_error_restores_context_only_after_exception_mapp
|
|||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_chunk_creator_preserves_hidden_provider_specific_fields_from_parsed_chunk():
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=None,
|
||||
model="gemini-3.5-flash",
|
||||
logging_obj=MagicMock(),
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
parsed_chunk = ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)],
|
||||
)
|
||||
parsed_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"}
|
||||
|
||||
result = wrapper.chunk_creator(chunk=parsed_chunk)
|
||||
|
||||
assert result is not None
|
||||
assert result._hidden_params["provider_specific_fields"] == {"traffic_type": "ON_DEMAND_FLEX"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_stream_assembled_response_keeps_vertex_traffic_type(logging_obj: Logging):
|
||||
content_chunk = ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content="hello", role="assistant"), finish_reason=None)],
|
||||
)
|
||||
final_chunk = ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=""), finish_reason="stop")],
|
||||
)
|
||||
setattr(final_chunk, "usage", Usage(prompt_tokens=7, completion_tokens=5, total_tokens=12))
|
||||
final_chunk._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND_FLEX"}
|
||||
|
||||
async def _stream():
|
||||
yield content_chunk
|
||||
yield final_chunk
|
||||
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=_stream(),
|
||||
model="gemini-3.5-flash",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="vertex_ai",
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
|
||||
received = [chunk async for chunk in wrapper]
|
||||
|
||||
assembled = litellm.stream_chunk_builder(chunks=received, messages=[{"role": "user", "content": "hi"}])
|
||||
assert assembled is not None
|
||||
assert assembled._hidden_params["provider_specific_fields"]["traffic_type"] == "ON_DEMAND_FLEX"
|
||||
|
|
|
|||
|
|
@ -290,6 +290,24 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
assert data.get("litellm_metadata", {}).get("guardrails")
|
||||
assert guardrail.dynamic_params == {"policy_id": "policy-123"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_native_tools_survive_guardrail_round_trip(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = MockPassThroughGuardrail(guardrail_name="test")
|
||||
data = {
|
||||
"model": "gemini-2.5-flash",
|
||||
"messages": [{"role": "user", "content": "coffee shops near Union Square?"}],
|
||||
"tools": [
|
||||
{"googleMaps": {"enable_widget": True}},
|
||||
{"name": "get_weather", "input_schema": {"type": "object", "properties": {}}},
|
||||
],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert {"googleMaps": {"enable_widget": True}} in data["tools"]
|
||||
assert [tool["name"] for tool in data["tools"] if "name" in tool] == ["get_weather"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_midturn_system_correction_is_guardrailed_when_top_level_system_is_skipped(
|
||||
self,
|
||||
|
|
@ -1818,3 +1836,72 @@ class TestAnthropicMessagesScanOnlyToolResults:
|
|||
|
||||
assert guardrail.captured_inputs is not None
|
||||
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
|
||||
|
||||
|
||||
class TestStructuredWriteBackKeepsToolResults:
|
||||
"""A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103)."""
|
||||
|
||||
@staticmethod
|
||||
def _claude_code_tool_search_turns(tool_result_content):
|
||||
return [
|
||||
{"role": "user", "content": "load WebFetch for bob@example.com"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01",
|
||||
"name": "ToolSearch",
|
||||
"input": {"query": "select:WebFetch"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": "toolu_01", "content": tool_result_content},
|
||||
{"type": "text", "text": "Now fetch the page."},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _blocks(message):
|
||||
return message["content"] if isinstance(message["content"], list) else []
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("tool_result_content", "expected_written_back_content"),
|
||||
[
|
||||
(
|
||||
[{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
[{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
),
|
||||
([], ""),
|
||||
],
|
||||
ids=["tool_reference", "empty"],
|
||||
)
|
||||
async def test_tool_result_stays_right_after_its_tool_use(
|
||||
self, tool_result_content, expected_written_back_content
|
||||
):
|
||||
handler = AnthropicMessagesHandler()
|
||||
data = {"model": "claude-fable-5", "messages": self._claude_code_tool_search_turns(tool_result_content)}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=MockStructuredMaskingGuardrail())
|
||||
|
||||
serialized = json.dumps(data["messages"])
|
||||
assert "bob@example.com" not in serialized
|
||||
assert "<EMAIL>" in serialized
|
||||
|
||||
messages = data["messages"]
|
||||
tool_use_index = next(
|
||||
i for i, m in enumerate(messages) if any(b.get("type") == "tool_use" for b in self._blocks(m))
|
||||
)
|
||||
answer = messages[tool_use_index + 1]
|
||||
assert answer["role"] == "user"
|
||||
assert answer["content"][0] == {
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_01",
|
||||
"content": expected_written_back_content,
|
||||
}
|
||||
later_blocks = [b for m in messages[tool_use_index + 1 :] for b in self._blocks(m)]
|
||||
assert {"type": "text", "text": "Now fetch the page."} in later_blocks
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
OPENAI_MAX_TOOL_NAME_LENGTH,
|
||||
AnthropicAdapter,
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
create_tool_name_mapping,
|
||||
truncate_tool_name,
|
||||
|
|
@ -2307,6 +2308,53 @@ def test_translate_anthropic_tools_to_openai_fills_missing_tool_name():
|
|||
assert result[1]["function"]["name"] == "litellm_unnamed_tool_1"
|
||||
|
||||
|
||||
def test_translate_anthropic_tools_to_openai_passes_provider_native_tool_dicts_through():
|
||||
"""Deployment-level provider-native tools (e.g. Gemini googleMaps) must reach the provider transformation verbatim (LIT-6286)."""
|
||||
tools = [
|
||||
{"googleMaps": {}},
|
||||
{"googleSearch": {}},
|
||||
{
|
||||
"name": "get_weather",
|
||||
"input_schema": {"type": "object", "properties": {"location": {"type": "string"}}},
|
||||
},
|
||||
]
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
result, tool_name_mapping = adapter.translate_anthropic_tools_to_openai(tools=tools, model=None)
|
||||
assert result[0] == {"googleMaps": {}}
|
||||
assert result[1] == {"googleSearch": {}}
|
||||
assert result[2]["function"]["name"] == "get_weather"
|
||||
assert tool_name_mapping == {}
|
||||
|
||||
|
||||
def test_translate_anthropic_tools_to_openai_passes_openai_function_tools_through():
|
||||
"""A tool already in OpenAI function format must pass through unchanged instead of becoming litellm_unnamed_tool_N."""
|
||||
openai_tool = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"location": {"type": "string"}}},
|
||||
},
|
||||
}
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
result, _ = adapter.translate_anthropic_tools_to_openai(tools=[openai_tool], model=None)
|
||||
assert result == [openai_tool]
|
||||
|
||||
|
||||
def test_translate_completion_input_params_keeps_provider_native_tools():
|
||||
"""/v1/messages request translation must keep router-merged provider-native tools in kwargs['tools'] (LIT-6286)."""
|
||||
adapter = AnthropicAdapter()
|
||||
translated = adapter.translate_completion_input_params(
|
||||
{
|
||||
"model": "gemini/gemini-2.5-flash",
|
||||
"max_tokens": 1024,
|
||||
"messages": [{"role": "user", "content": "coffee shops near Union Square"}],
|
||||
"tools": [{"googleMaps": {}}],
|
||||
}
|
||||
)
|
||||
assert translated is not None
|
||||
assert translated["tools"] == [{"googleMaps": {}}]
|
||||
|
||||
|
||||
def test_translate_openai_content_to_anthropic_reasoning_content_without_thinking_blocks():
|
||||
"""
|
||||
Test that reasoning_content is converted to thinking block when thinking_blocks is not present.
|
||||
|
|
@ -3999,6 +4047,75 @@ def test_translate_anthropic_messages_to_openai_carries_midturn_system_prompt_ca
|
|||
]
|
||||
|
||||
|
||||
def _tool_reference_block(tool_name="WebFetch"):
|
||||
return {"type": "tool_reference", "tool_name": tool_name}
|
||||
|
||||
|
||||
def test_tool_result_tool_reference_is_carried_through_untouched():
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
result = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=[
|
||||
_anthropic_tool_use_turn("toolu_01"),
|
||||
_anthropic_tool_result_turn({"toolu_01": [_tool_reference_block()]}),
|
||||
]
|
||||
)
|
||||
|
||||
assert [m["role"] for m in result] == ["assistant", "tool"]
|
||||
assert result[1]["tool_call_id"] == "toolu_01"
|
||||
assert result[1]["content"] == [{"type": "tool_reference", "tool_name": "WebFetch"}]
|
||||
|
||||
|
||||
def test_tool_result_text_beside_tool_reference_keeps_both_parts_in_order():
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
result = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=[
|
||||
_anthropic_tool_use_turn("toolu_01"),
|
||||
_anthropic_tool_result_turn(
|
||||
{"toolu_01": [{"type": "text", "text": "loaded"}, _tool_reference_block("Grep")]}
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
assert result[1]["content"] == [
|
||||
{"type": "text", "text": "loaded"},
|
||||
{"type": "tool_reference", "tool_name": "Grep"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tool_result_content",
|
||||
[
|
||||
[],
|
||||
None,
|
||||
"",
|
||||
{"not": "a list"},
|
||||
[{"type": "future_block", "payload": 1}],
|
||||
[{"type": "search_result", "source": "https://example.com", "title": "t", "content": []}],
|
||||
],
|
||||
ids=["empty_list", "null", "empty_string", "non_list", "unknown_block", "search_result_only"],
|
||||
)
|
||||
def test_tool_result_without_translatable_content_still_answers_its_tool_use(tool_result_content):
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
result = adapter.translate_anthropic_messages_to_openai(
|
||||
messages=[
|
||||
_anthropic_tool_use_turn("toolu_01"),
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": tool_result_content}],
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
assert result == [
|
||||
result[0],
|
||||
{"role": "tool", "tool_call_id": "toolu_01", "content": ""},
|
||||
]
|
||||
assert result[0]["role"] == "assistant"
|
||||
|
||||
|
||||
def _openai_response_with_usage(usage: Usage) -> ModelResponse:
|
||||
return ModelResponse(
|
||||
id="resp_web_search",
|
||||
|
|
|
|||
|
|
@ -102,6 +102,35 @@ def test_transform_request_hoists_tool_message_image():
|
|||
]
|
||||
|
||||
|
||||
def test_transform_request_drops_tool_reference_parts():
|
||||
"""Azure's transform_request shares the tool-message sanitizing with OpenAI:
|
||||
tool_reference parts are dropped, a reference-only result keeps its tool
|
||||
message with empty text (#37462 round trip)."""
|
||||
messages = [
|
||||
{"role": "user", "content": "load the WebFetch tool"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "ToolSearch", "arguments": "{}"}}],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": [{"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
},
|
||||
]
|
||||
|
||||
request = AzureOpenAIConfig().transform_request(
|
||||
model="gpt-4o",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request["messages"][2]["content"] == ""
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, emitted_key, absent_key",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -59,17 +59,17 @@ class GptProfile(NamedTuple):
|
|||
GPT_5_6_PROFILES = [
|
||||
GptProfile(
|
||||
model_id="us.openai.gpt-5.6-sol",
|
||||
input_cost=5.5e-06, input_cost_above_272k=1.1e-05,
|
||||
cache_write=6.875e-06, cache_write_above_272k=1.375e-05,
|
||||
cache_read=5.5e-07, cache_read_above_272k=1.1e-06,
|
||||
output_cost=3.3e-05, output_cost_above_272k=4.95e-05,
|
||||
input_cost=4.4e-06, input_cost_above_272k=8.8e-06,
|
||||
cache_write=5.5e-06, cache_write_above_272k=1.1e-05,
|
||||
cache_read=4.4e-07, cache_read_above_272k=8.8e-07,
|
||||
output_cost=2.2e-05, output_cost_above_272k=3.3e-05,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="global.openai.gpt-5.6-sol",
|
||||
input_cost=5e-06, input_cost_above_272k=1e-05,
|
||||
cache_write=6.25e-06, cache_write_above_272k=1.25e-05,
|
||||
cache_read=5e-07, cache_read_above_272k=1e-06,
|
||||
output_cost=3e-05, output_cost_above_272k=4.5e-05,
|
||||
input_cost=4e-06, input_cost_above_272k=8e-06,
|
||||
cache_write=5e-06, cache_write_above_272k=1e-05,
|
||||
cache_read=4e-07, cache_read_above_272k=8e-07,
|
||||
output_cost=2e-05, output_cost_above_272k=3e-05,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="us.openai.gpt-5.6-terra",
|
||||
|
|
@ -221,7 +221,7 @@ def test_bedrock_gpt_5_6_above_272k_tier_applies_to_cost(local_model_cost_map):
|
|||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx((300000 * 1.1e-05) + (1000 * 4.95e-05), rel=1e-9)
|
||||
assert cost == pytest.approx((300000 * 8.8e-06) + (1000 * 3.3e-05), rel=1e-9)
|
||||
|
||||
|
||||
def test_bedrock_gpt_5_6_bills_cache_read_tokens(local_model_cost_map):
|
||||
|
|
@ -241,10 +241,10 @@ def test_bedrock_gpt_5_6_bills_cache_read_tokens(local_model_cost_map):
|
|||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
expected = (2 * 5.5e-06) + (15609 * 5.5e-07) + (5 * 3.3e-05)
|
||||
expected = (2 * 4.4e-06) + (15609 * 4.4e-07) + (5 * 2.2e-05)
|
||||
assert cost == pytest.approx(expected, rel=1e-9)
|
||||
# Without cache_read_input_token_cost the cached prefix bills at zero.
|
||||
assert cost > (15611 * 5.5e-06) * 0.1
|
||||
assert cost > (15611 * 4.4e-06) * 0.1
|
||||
|
||||
|
||||
def test_bedrock_gpt_5_6_bills_cache_write_tokens(local_model_cost_map):
|
||||
|
|
@ -263,7 +263,7 @@ def test_bedrock_gpt_5_6_bills_cache_write_tokens(local_model_cost_map):
|
|||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
expected = (2 * 5.5e-06) + (15609 * 6.875e-06) + (5 * 3.3e-05)
|
||||
expected = (2 * 4.4e-06) + (15609 * 5.5e-06) + (5 * 2.2e-05)
|
||||
assert cost == pytest.approx(expected, rel=1e-9)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1683,10 +1683,19 @@ class TestBedrockMantleResponsesPricing:
|
|||
assert info["cache_read_input_token_cost"] == pytest.approx(2.75e-07)
|
||||
assert info["max_input_tokens"] == 1050000
|
||||
|
||||
def test_gpt_5_6_cyber_pricing_and_mode(self, local_cost_map):
|
||||
info = litellm.get_model_info("bedrock_mantle/openai.gpt-5.6-cyber")
|
||||
assert info["mode"] == "responses"
|
||||
assert info["input_cost_per_token"] == pytest.approx(1.375e-05)
|
||||
assert info["cache_creation_input_token_cost"] == pytest.approx(1.71875e-05)
|
||||
assert info["cache_read_input_token_cost"] == pytest.approx(1.375e-06)
|
||||
assert info["output_cost_per_token"] == pytest.approx(8.25e-05)
|
||||
assert info["max_input_tokens"] == 272000
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, input_cost, cache_creation_cost, cache_read_cost, output_cost",
|
||||
[
|
||||
("openai.gpt-5.6-sol", 5.5e-06, 6.875e-06, 5.5e-07, 3.3e-05),
|
||||
("openai.gpt-5.6-sol", 4.4e-06, 5.5e-06, 4.4e-07, 2.2e-05),
|
||||
("openai.gpt-5.6-terra", 2.2e-06, 2.75e-06, 2.2e-07, 1.32e-05),
|
||||
("openai.gpt-5.6-luna", 2.2e-07, 2.75e-07, 2.2e-08, 1.32e-06),
|
||||
],
|
||||
|
|
@ -1709,7 +1718,7 @@ class TestBedrockMantleResponsesPricing:
|
|||
@pytest.mark.parametrize(
|
||||
"model, input_cost, output_cost",
|
||||
[
|
||||
("openai.gpt-5.6-sol", 5.5e-06, 3.3e-05),
|
||||
("openai.gpt-5.6-sol", 4.4e-06, 2.2e-05),
|
||||
("openai.gpt-5.6-terra", 2.2e-06, 1.32e-05),
|
||||
("openai.gpt-5.6-luna", 2.2e-07, 1.32e-06),
|
||||
],
|
||||
|
|
|
|||
|
|
@ -1901,6 +1901,76 @@ async def test_async_audio_transcriptions_sends_dict_data_as_json_body():
|
|||
assert response.text == "transcribed"
|
||||
|
||||
|
||||
class _WordTimestampAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
|
||||
def transform_audio_transcription_response(self, raw_response):
|
||||
payload = raw_response.json()
|
||||
response = TranscriptionResponse(text=payload["text"])
|
||||
response["words"] = payload["words"]
|
||||
return response
|
||||
|
||||
|
||||
def test_transform_audio_transcription_response_without_subtitle_opt_in_keeps_text_and_words():
|
||||
words = [
|
||||
{"word": "hello", "start": 0.0, "end": 0.5},
|
||||
{"word": "world", "start": 0.5, "end": 1.0},
|
||||
]
|
||||
raw_response = httpx.Response(200, json={"text": "hello world", "words": words})
|
||||
|
||||
response = BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=_WordTimestampAudioTranscriptionConfig(),
|
||||
model="test-model",
|
||||
response=raw_response,
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": "srt"},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert response["words"] == words
|
||||
|
||||
|
||||
class _SubtitleSynthesisAudioTranscriptionConfig(_JSONBodyAudioTranscriptionConfig):
|
||||
@property
|
||||
def supports_subtitle_synthesis(self) -> bool:
|
||||
return True
|
||||
|
||||
def transform_audio_transcription_response(self, raw_response):
|
||||
payload = raw_response.json()
|
||||
response = TranscriptionResponse(text=payload["text"])
|
||||
if "words" in payload:
|
||||
response["words"] = payload["words"]
|
||||
return response
|
||||
|
||||
|
||||
def _transform_subtitle_response(payload):
|
||||
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=_SubtitleSynthesisAudioTranscriptionConfig(),
|
||||
model="test-model",
|
||||
response=httpx.Response(200, json=payload),
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": "srt"},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
|
||||
def test_subtitle_synthesis_fallback_without_timings_drops_words():
|
||||
response = _transform_subtitle_response(
|
||||
{"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]}
|
||||
)
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert "words" not in response
|
||||
|
||||
|
||||
def test_subtitle_synthesis_without_words_keeps_plain_text():
|
||||
response = _transform_subtitle_response({"text": "hello world"})
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert "words" not in response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_retrieve_file_content_raises_on_http_error():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -61,6 +61,8 @@ PUBLISHED_DBU_PER_MILLION: Final = {
|
|||
"databricks/databricks-gemini-3-1-flash-lite": ("4.464", "26.786", "4.464", "0.446"),
|
||||
"databricks/databricks-gemini-2-5-pro": ("22.321", "178.571", "22.321", "2.232"),
|
||||
"databricks/databricks-gemini-2-5-flash": ("5.357", "44.643", "5.357", "0.536"),
|
||||
"databricks/databricks-kimi-k3": ("42.857", "214.286", "42.857", "4.286"),
|
||||
"databricks/databricks-glm-5-2": ("20.000", "62.857", "20.000", "3.714"),
|
||||
}
|
||||
PROMOTIONAL_DISCOUNT: Final = 0.80
|
||||
PROMOTION_EXPIRES: Final = "2027-01-31"
|
||||
|
|
|
|||
|
|
@ -169,6 +169,42 @@ class TestTransformRequest:
|
|||
}
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["srt", "vtt"])
|
||||
def test_subtitle_response_format_requests_word_timestamps(self, config, response_format):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": response_format},
|
||||
litellm_params={},
|
||||
)
|
||||
transcription_config = request_data.data["generation_config"]["transcription_config"]
|
||||
assert json.loads(json.dumps(transcription_config)) == {
|
||||
"mode": {
|
||||
"type": "verbatim",
|
||||
"timestamp_granularities": ["word"],
|
||||
"diarization_mode": "speaker",
|
||||
}
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["json", "text", "verbose_json"])
|
||||
def test_non_subtitle_response_format_sends_no_mode(self, config, response_format):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": response_format},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "generation_config" not in request_data.data
|
||||
|
||||
def test_non_string_response_format_sends_no_mode(self, config):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
audio_file=("sample.wav", AUDIO_BYTES, "audio/wav"),
|
||||
optional_params={"response_format": {"type": "json_object"}},
|
||||
litellm_params={},
|
||||
)
|
||||
assert "generation_config" not in request_data.data
|
||||
|
||||
def test_segment_granularity_sends_no_mode(self, config):
|
||||
request_data = config.transform_audio_transcription_request(
|
||||
model="gemini-3.5-transcribe",
|
||||
|
|
@ -214,6 +250,54 @@ class TestTransformResponse:
|
|||
assert response.get("duration") is None
|
||||
|
||||
|
||||
class TestSubtitleSynthesisThroughHandler:
|
||||
def _transform(self, config, response_format):
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.utils import TranscriptionResponse
|
||||
|
||||
return BaseLLMHTTPHandler()._transform_audio_transcription_response(
|
||||
provider_config=config,
|
||||
model="gemini-3.5-transcribe",
|
||||
response=make_response(COMPLETED_RESPONSE),
|
||||
model_response=TranscriptionResponse(),
|
||||
logging_obj=Mock(),
|
||||
optional_params={"response_format": response_format},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
def test_supports_subtitle_synthesis(self, config):
|
||||
assert config.supports_subtitle_synthesis is True
|
||||
|
||||
def test_srt_synthesizes_subtitle_document_and_drops_words(self, config):
|
||||
response = self._transform(config, "srt")
|
||||
assert response.text == (
|
||||
"1\n00:00:00,100 --> 00:00:00,400\nHello\n\n2\n00:00:00,500 --> 00:00:00,900\nworld.\n"
|
||||
)
|
||||
assert "words" not in response
|
||||
assert response["task"] == "transcribe"
|
||||
assert response["duration"] == 0.9
|
||||
assert response.usage.total_tokens == 200
|
||||
|
||||
def test_vtt_synthesizes_subtitle_document_and_drops_words(self, config):
|
||||
response = self._transform(config, "vtt")
|
||||
assert response.text == (
|
||||
"WEBVTT\n\n00:00:00.100 --> 00:00:00.400\nHello\n\n00:00:00.500 --> 00:00:00.900\nworld.\n"
|
||||
)
|
||||
assert "words" not in response
|
||||
assert response.usage.total_tokens == 200
|
||||
|
||||
@pytest.mark.parametrize("response_format", ["json", "verbose_json"])
|
||||
def test_non_subtitle_formats_keep_plain_text_and_words(self, config, response_format):
|
||||
response = self._transform(config, response_format)
|
||||
assert response.text == "Hello world."
|
||||
assert response["words"] == [
|
||||
{"word": "Hello", "start": 0.1, "end": 0.4, "speaker": "spk:0"},
|
||||
{"word": "world.", "start": 0.5, "end": 0.9, "speaker": "spk:1"},
|
||||
]
|
||||
|
||||
|
||||
class TestCostRegression:
|
||||
@pytest.fixture
|
||||
def local_cost_map(self, monkeypatch):
|
||||
|
|
|
|||
|
|
@ -1866,6 +1866,54 @@ def test_map_openai_params_drops_stock_voice_case_insensitively():
|
|||
assert passthrough["generationConfig"]["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
|
||||
|
||||
|
||||
def test_gemini_response_done_bills_audio_output_tokens_at_audio_rate(monkeypatch):
|
||||
"""Regression for the Gemini Live AUDIO output breakdown: responseTokensDetails
|
||||
must survive into response.done usage and bill at output_cost_per_audio_token,
|
||||
not the text rate."""
|
||||
from litellm.cost_calculator import (
|
||||
RealtimeAPITokenUsageProcessor,
|
||||
handle_realtime_stream_cost_calculation,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
|
||||
config = GeminiRealtimeConfig()
|
||||
done_event = config.transform_response_done_event(
|
||||
message={
|
||||
"serverContent": {"turnComplete": True},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 377,
|
||||
"responseTokenCount": 51,
|
||||
"totalTokenCount": 428,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 377}],
|
||||
"responseTokensDetails": [{"modality": "AUDIO", "tokenCount": 51}],
|
||||
"thoughtsTokenCount": 37,
|
||||
},
|
||||
},
|
||||
current_response_id="resp_lit6277",
|
||||
current_conversation_id="conv_lit6277",
|
||||
output_items=None,
|
||||
)
|
||||
|
||||
usage = done_event["response"]["usage"]
|
||||
assert usage["output_tokens_details"]["audio_tokens"] == 51
|
||||
assert usage["output_token_details"]["audio_tokens"] == 51
|
||||
|
||||
results = [done_event]
|
||||
combined_usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
|
||||
results=results,
|
||||
)
|
||||
assert combined_usage.completion_tokens_details is not None
|
||||
assert combined_usage.completion_tokens_details.audio_tokens == 51
|
||||
|
||||
cost = handle_realtime_stream_cost_calculation(
|
||||
results=results,
|
||||
combined_usage_object=combined_usage,
|
||||
custom_llm_provider="gemini",
|
||||
litellm_model_name="gemini-2.5-flash-native-audio-preview-12-2025",
|
||||
)
|
||||
assert cost == pytest.approx(377 * 5e-07 + 51 * 1.2e-05 + 37 * 2e-06)
|
||||
@pytest.fixture(autouse=False)
|
||||
def patch_gemini_transcribe_live_cost_map_entry(monkeypatch):
|
||||
"""Inject the gemini-3.5-transcribe-live registry entry locally.
|
||||
|
|
|
|||
|
|
@ -869,6 +869,69 @@ class TestToolMessageImageHoisting:
|
|||
assert result[3]["content"] == self.HOISTED_USER_CONTENT
|
||||
|
||||
|
||||
class TestToolReferenceStripping:
|
||||
"""transform_request drops tool_reference parts from tool messages: OpenAI's
|
||||
chat API rejects them, and the reference names an already-declared tool
|
||||
rather than carrying content (#37462 round trip)."""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = OpenAIGPTConfig()
|
||||
|
||||
def _messages_with_tool_reference(self, extra_parts=()):
|
||||
return [
|
||||
{"role": "user", "content": "load the WebFetch tool"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "ToolSearch", "arguments": "{}"}}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": [*extra_parts, {"type": "tool_reference", "tool_name": "WebFetch"}],
|
||||
},
|
||||
]
|
||||
|
||||
def test_transform_request_keeps_text_and_drops_reference(self):
|
||||
request = self.config.transform_request(
|
||||
model="gpt-4.1",
|
||||
messages=self._messages_with_tool_reference(extra_parts=({"type": "text", "text": "loaded"},)),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
tool_message = request["messages"][2]
|
||||
assert tool_message["content"] == [{"type": "text", "text": "loaded"}]
|
||||
assert tool_message["tool_call_id"] == "call_1"
|
||||
|
||||
def test_transform_request_reference_only_keeps_tool_message_with_empty_text(self):
|
||||
request = self.config.transform_request(
|
||||
model="gpt-4.1",
|
||||
messages=self._messages_with_tool_reference(),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert [m.get("role") for m in request["messages"]] == ["user", "assistant", "tool"]
|
||||
assert request["messages"][2]["content"] == ""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_transform_request_drops_reference(self):
|
||||
request = await self.config.async_transform_request(
|
||||
model="gpt-4.1",
|
||||
messages=self._messages_with_tool_reference(),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request["messages"][2]["content"] == ""
|
||||
|
||||
|
||||
class TestOpenAIPromptCacheBreakpointChatPath:
|
||||
"""Chat-path shape for OpenAI explicit prompt caching (#37509)."""
|
||||
|
||||
|
|
|
|||
|
|
@ -477,12 +477,12 @@ class TestBuildResponseWithResponseFormat:
|
|||
}
|
||||
}
|
||||
# SRT requested but tokens have no start_ms/end_ms -> empty SRT
|
||||
# falls back gracefully since _group_tokens_into_cues skips them
|
||||
# falls back gracefully since group_subtitle_tokens_into_cues skips them
|
||||
resp = cfg._build_response_from_payload(payload, response_format="srt")
|
||||
# With no timestamp data, SRT rendering produces empty string,
|
||||
# but we still get output because the code checks `tokens` truthiness
|
||||
# before choosing SRT path. Actually the tokens list is truthy but
|
||||
# _group_tokens_into_cues will produce no cues -> empty SRT string.
|
||||
# group_subtitle_tokens_into_cues will produce no cues -> empty SRT string.
|
||||
# Let's verify it doesn't crash.
|
||||
assert isinstance(resp.text, str)
|
||||
|
||||
|
|
|
|||
|
|
@ -18,8 +18,14 @@ from litellm.types.utils import LlmProviders, ModelResponse
|
|||
|
||||
TOOL_CALLING_MODEL = "openai/gpt-oss-20b"
|
||||
REASONING_MODEL = "deepseek-ai/DeepSeek-V3.1"
|
||||
PLAIN_MODEL = "Qwen/Qwen3-235B-A22B-fp8-tput"
|
||||
UNMAPPED_MODEL = "example-org/brand-new-model"
|
||||
NO_TOOLS_MODEL = "example-org/no-tools-model"
|
||||
ADJUSTABLE_REASONING_MODEL = "openai/gpt-oss-120b"
|
||||
HYBRID_REASONING_MODEL = "Qwen/Qwen3.5-9B"
|
||||
HIGH_MAX_REASONING_MODEL = "deepseek-ai/DeepSeek-V4-Pro"
|
||||
REGISTRY_FLAGGED_REASONING_MODEL = "zai-org/GLM-4.6"
|
||||
NON_REASONING_MODEL = "meta-llama/Llama-3.3-70B-Instruct-Turbo"
|
||||
NO_SCHEMA_MODEL = "example-org/no-schema-model"
|
||||
|
||||
TOOL_PARAMS = ("tools", "tool_choice", "function_call")
|
||||
|
|
@ -39,6 +45,15 @@ JSON_SCHEMA_RESPONSE_FORMAT = {
|
|||
REGEX_RESPONSE_FORMAT = {"type": "regex", "pattern": "(positive|neutral|negative)"}
|
||||
|
||||
|
||||
def _map_reasoning_effort(model: str, effort: str) -> dict:
|
||||
return TogetherAIChatConfig().map_openai_params(
|
||||
non_default_params={"reasoning_effort": effort},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def force_local_model_cost(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
|
|
@ -191,6 +206,116 @@ def test_map_openai_params_schema_model_passes_response_format_through(response_
|
|||
assert mapped["response_format"] == response_format
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[ADJUSTABLE_REASONING_MODEL, HYBRID_REASONING_MODEL, HIGH_MAX_REASONING_MODEL, REGISTRY_FLAGGED_REASONING_MODEL],
|
||||
)
|
||||
def test_supported_params_includes_reasoning_effort_for_reasoning_models(model):
|
||||
supported = TogetherAIChatConfig().get_supported_openai_params(model=model)
|
||||
|
||||
assert "reasoning_effort" in supported
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", [NON_REASONING_MODEL, PLAIN_MODEL])
|
||||
def test_supported_params_excludes_reasoning_effort_for_non_reasoning_models(model):
|
||||
supported = TogetherAIChatConfig().get_supported_openai_params(model=model)
|
||||
|
||||
assert "reasoning_effort" not in supported
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"effort, expected",
|
||||
[("low", "low"), ("medium", "medium"), ("high", "high"), ("minimal", "low"), ("xhigh", "high"), ("max", "high")],
|
||||
)
|
||||
def test_adjustable_model_translates_reasoning_effort(effort, expected):
|
||||
mapped = _map_reasoning_effort(ADJUSTABLE_REASONING_MODEL, effort)
|
||||
|
||||
assert mapped["reasoning_effort"] == expected
|
||||
assert "reasoning" not in mapped
|
||||
|
||||
|
||||
def test_adjustable_model_cannot_disable_reasoning_so_none_becomes_low():
|
||||
mapped = _map_reasoning_effort(ADJUSTABLE_REASONING_MODEL, "none")
|
||||
|
||||
assert mapped["reasoning_effort"] == "low"
|
||||
assert "reasoning" not in mapped
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"effort, expected",
|
||||
[("low", "low"), ("medium", "medium"), ("high", "high"), ("minimal", "low"), ("xhigh", "high"), ("max", "high")],
|
||||
)
|
||||
def test_hybrid_model_translates_reasoning_effort(effort, expected):
|
||||
mapped = _map_reasoning_effort(HYBRID_REASONING_MODEL, effort)
|
||||
|
||||
assert mapped["reasoning_effort"] == expected
|
||||
assert "reasoning" not in mapped
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", [HYBRID_REASONING_MODEL, HIGH_MAX_REASONING_MODEL, REGISTRY_FLAGGED_REASONING_MODEL])
|
||||
def test_reasoning_effort_none_becomes_reasoning_toggle(model):
|
||||
mapped = _map_reasoning_effort(model, "none")
|
||||
|
||||
assert mapped["reasoning"] == {"enabled": False}
|
||||
assert "reasoning_effort" not in mapped
|
||||
|
||||
|
||||
def test_reasoning_effort_none_does_not_clobber_user_reasoning():
|
||||
mapped = TogetherAIChatConfig().map_openai_params(
|
||||
non_default_params={"reasoning_effort": "none"},
|
||||
optional_params={"reasoning": {"enabled": True}},
|
||||
model=HYBRID_REASONING_MODEL,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert mapped["reasoning"] == {"enabled": True}
|
||||
assert "reasoning_effort" not in mapped
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"effort, expected",
|
||||
[("minimal", "high"), ("low", "high"), ("medium", "high"), ("high", "high"), ("xhigh", "max"), ("max", "max")],
|
||||
)
|
||||
def test_deepseek_v4_pro_remaps_to_high_max(effort, expected):
|
||||
mapped = _map_reasoning_effort(HIGH_MAX_REASONING_MODEL, effort)
|
||||
|
||||
assert mapped["reasoning_effort"] == expected
|
||||
|
||||
|
||||
def test_deepseek_v4_pro_dated_variant_remaps_via_prefix():
|
||||
mapped = _map_reasoning_effort(f"{HIGH_MAX_REASONING_MODEL}-0813", "low")
|
||||
|
||||
assert mapped["reasoning_effort"] == "high"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", [ADJUSTABLE_REASONING_MODEL, HYBRID_REASONING_MODEL, HIGH_MAX_REASONING_MODEL])
|
||||
def test_reasoning_effort_default_is_dropped(model):
|
||||
mapped = _map_reasoning_effort(model, "default")
|
||||
|
||||
assert "reasoning_effort" not in mapped
|
||||
assert "reasoning" not in mapped
|
||||
|
||||
|
||||
def test_get_optional_params_translates_reasoning_effort_for_together():
|
||||
optional_params = litellm.get_optional_params(
|
||||
model=ADJUSTABLE_REASONING_MODEL,
|
||||
custom_llm_provider="together_ai",
|
||||
reasoning_effort="max",
|
||||
)
|
||||
|
||||
assert optional_params["reasoning_effort"] == "high"
|
||||
|
||||
|
||||
def test_get_optional_params_rejects_reasoning_effort_for_non_reasoning_together_model():
|
||||
with pytest.raises(litellm.UnsupportedParamsError):
|
||||
litellm.get_optional_params(
|
||||
model=NON_REASONING_MODEL,
|
||||
custom_llm_provider="together_ai",
|
||||
reasoning_effort="low",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("drop_params", [False, True])
|
||||
def test_map_openai_params_unmapped_model_passes_response_format_through(drop_params, together_warning_log):
|
||||
mapped = TogetherAIChatConfig().map_openai_params(
|
||||
|
|
|
|||
75
tests/test_litellm/llms/xai/test_xai_model_registry.py
Normal file
75
tests/test_litellm/llms/xai/test_xai_model_registry.py
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
"""
|
||||
Registry regression tests for xAI entries in the model cost map.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
REPO_ROOT = Path(__file__).parents[4]
|
||||
PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json"
|
||||
BACKUP_PRICES_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json"
|
||||
|
||||
# Retired by xAI and no longer served: requests to these slugs 404 rather than
|
||||
# redirecting, and they are absent from https://docs.x.ai/docs/models
|
||||
RETIRED_MODELS = (
|
||||
"xai/grok-2",
|
||||
"xai/grok-2-1212",
|
||||
"xai/grok-2-latest",
|
||||
"xai/grok-2-vision",
|
||||
"xai/grok-2-vision-1212",
|
||||
"xai/grok-2-vision-latest",
|
||||
"xai/grok-beta",
|
||||
"xai/grok-vision-beta",
|
||||
)
|
||||
|
||||
# https://docs.x.ai/developers/model-capabilities/text/multi-agent
|
||||
# "The multi-agent model does not work with the OpenAI Chat Completions API."
|
||||
RESPONSES_ONLY_MODELS = (
|
||||
"xai/grok-4.20-multi-agent-0309",
|
||||
"xai/grok-4.20-multi-agent-beta-0309",
|
||||
)
|
||||
|
||||
MAP_PATHS = (PRICES_PATH, BACKUP_PRICES_PATH)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", params=[p.name for p in MAP_PATHS])
|
||||
def cost_map(request: pytest.FixtureRequest) -> dict:
|
||||
path = next(p for p in MAP_PATHS if p.name == request.param)
|
||||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", RETIRED_MODELS)
|
||||
def test_retired_xai_models_are_not_advertised(cost_map: dict, model: str):
|
||||
assert model not in cost_map
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", RESPONSES_ONLY_MODELS)
|
||||
def test_multi_agent_models_are_responses_only(cost_map: dict, model: str):
|
||||
entry = cost_map[model]
|
||||
assert entry["supported_endpoints"] == ["/v1/responses"]
|
||||
assert entry["mode"] == "responses"
|
||||
assert "/v1/chat/completions" not in entry["supported_endpoints"]
|
||||
|
||||
|
||||
def test_surviving_xai_chat_models_still_serve_chat_completions(cost_map: dict):
|
||||
"""Guard against the removal above over-reaching into live models."""
|
||||
chat_models = [
|
||||
key
|
||||
for key, value in cost_map.items()
|
||||
if isinstance(value, dict) and value.get("litellm_provider") == "xai" and value.get("mode") == "chat"
|
||||
]
|
||||
assert "xai/grok-4.3" in chat_models
|
||||
assert "xai/grok-4.6" in chat_models
|
||||
assert not any(key.startswith("xai/grok-2") for key in chat_models)
|
||||
|
||||
|
||||
def test_both_cost_maps_agree_on_xai_entries():
|
||||
prices = json.loads(PRICES_PATH.read_text(encoding="utf-8"))
|
||||
backup = json.loads(BACKUP_PRICES_PATH.read_text(encoding="utf-8"))
|
||||
xai_keys = {k for k, v in prices.items() if isinstance(v, dict) and v.get("litellm_provider") == "xai"}
|
||||
assert xai_keys
|
||||
assert {k: prices[k] for k in xai_keys} == {k: backup[k] for k in xai_keys}
|
||||
|
|
@ -10,6 +10,7 @@ from types import SimpleNamespace
|
|||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
oauth_protected_resource_path,
|
||||
|
|
@ -598,3 +599,92 @@ def test_client_credentials_uses_admin_entered_token_url_when_issuer_yield_empti
|
|||
assert spec is not None
|
||||
assert isinstance(spec.config, ClientCredentialsConfig)
|
||||
assert spec.config.token_url == "https://idp.example.com/token"
|
||||
|
||||
|
||||
_M2M_FIELDS = dict(
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="client_credentials",
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_url="https://idp.example.com/token",
|
||||
)
|
||||
_OBO_FIELDS = dict(
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_exchange_endpoint="https://idp.example.com/token",
|
||||
)
|
||||
_ID_JAG_FIELDS = dict(
|
||||
auth_type=MCPAuth.oauth2_id_jag,
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_exchange_endpoint="https://idp.example.com/token",
|
||||
id_jag_resource_token_endpoint="https://mcp-as.example.com/token",
|
||||
audience="api://mcp",
|
||||
)
|
||||
_AUTHZ_CODE_FIELDS = dict(auth_type=MCPAuth.oauth2, url="https://up.example.com/mcp")
|
||||
_STATIC_FIELDS = dict(auth_type=MCPAuth.bearer_token, authentication_token="static-tok")
|
||||
|
||||
_ARM_FIELDS = (
|
||||
("client_credentials", _M2M_FIELDS),
|
||||
("token_exchange", _OBO_FIELDS),
|
||||
("id_jag", _ID_JAG_FIELDS),
|
||||
("authorization_code", _AUTHZ_CODE_FIELDS),
|
||||
("api_key", _STATIC_FIELDS),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,fields", _ARM_FIELDS, ids=[n for n, _ in _ARM_FIELDS])
|
||||
def test_upstream_token_header_reaches_every_arms_config(name, fields):
|
||||
# to_server_spec builds each arm's config from a hand-written kwargs list, so an arm that
|
||||
# forgets to read the field fails silently: the server keeps writing to Authorization.
|
||||
spec = to_server_spec(_server(upstream_token_header="esb-oauth", **fields))
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "esb-oauth"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,fields", _ARM_FIELDS, ids=[n for n, _ in _ARM_FIELDS])
|
||||
def test_omitting_the_field_keeps_each_arms_shipped_default(name, fields):
|
||||
spec = to_server_spec(_server(**fields))
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "Authorization"
|
||||
|
||||
|
||||
def test_api_key_scheme_default_survives_when_the_field_is_unset():
|
||||
spec = to_server_spec(_server(auth_type=MCPAuth.api_key, authentication_token="k"))
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "X-API-Key"
|
||||
assert spec.config.value_prefix == ""
|
||||
|
||||
|
||||
def test_the_field_overrides_the_api_key_scheme_default():
|
||||
spec = to_server_spec(_server(auth_type=MCPAuth.api_key, authentication_token="k", upstream_token_header="X-Esb"))
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "X-Esb"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad", ["with space", "has:colon", "trailing\r\nX-Injected", 'quoted"name'])
|
||||
def test_a_malformed_header_name_is_refused_when_the_server_is_built(bad):
|
||||
"""Validation belongs at ingestion, not at spec building. Raising inside to_server_spec would
|
||||
abort the whole aggregate tools/list, so one mistyped server would silently empty the tool list
|
||||
for every other server too. Refusing at MCPServer construction fails the config load loudly
|
||||
instead, and means no malformed value can ever reach an arm.
|
||||
"""
|
||||
with pytest.raises(ValidationError):
|
||||
_server(upstream_token_header=bad, **_M2M_FIELDS)
|
||||
|
||||
|
||||
def test_a_valid_header_name_is_trimmed_at_ingestion():
|
||||
assert _server(upstream_token_header=" esb-oauth ", **_M2M_FIELDS).upstream_token_header == "esb-oauth"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("blank", ["", " ", "\t"])
|
||||
def test_a_blank_header_name_means_unset_rather_than_an_error(blank):
|
||||
"""The management API treats a blank as "not supplied" and stores it, so raising here made every
|
||||
later rebuild of that server 500 instead of falling back to the default Authorization behavior.
|
||||
"""
|
||||
server = _server(upstream_token_header=blank, **_M2M_FIELDS)
|
||||
assert server.upstream_token_header is None
|
||||
spec = to_server_spec(server)
|
||||
assert spec is not None
|
||||
assert spec.config.header_name == "Authorization"
|
||||
|
|
|
|||
|
|
@ -341,7 +341,7 @@ async def test_bearer_auth_sends_the_token_and_leaves_a_success_alone():
|
|||
async def refetch(failed: str) -> "str | None":
|
||||
raise AssertionError("must not refetch on success")
|
||||
|
||||
auth = ClientCredentialsBearerAuth("m2m-token", refetch)
|
||||
auth = ClientCredentialsBearerAuth("m2m-token", refetch, ClientCredentialsConfig())
|
||||
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
|
||||
response = await client.get("https://upstream.example.com/mcp")
|
||||
assert response.status_code == 200
|
||||
|
|
@ -357,7 +357,7 @@ async def test_bearer_auth_retries_a_401_once_with_a_fresh_token():
|
|||
refetched.append(failed)
|
||||
return "fresh-token"
|
||||
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch)
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
|
||||
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
|
||||
response = await client.get("https://upstream.example.com/mcp")
|
||||
assert response.status_code == 200
|
||||
|
|
@ -377,7 +377,7 @@ async def test_bearer_auth_remembers_the_rotated_token_for_later_requests():
|
|||
refetched.append(failed)
|
||||
return "fresh-token"
|
||||
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch)
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
|
||||
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
|
||||
first = await client.get("https://upstream.example.com/mcp")
|
||||
second = await client.get("https://upstream.example.com/mcp")
|
||||
|
|
@ -393,7 +393,7 @@ async def test_bearer_auth_surfaces_the_401_when_the_refetch_fails():
|
|||
async def refetch(failed: str) -> "str | None":
|
||||
return None
|
||||
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch)
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
|
||||
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
|
||||
response = await client.get("https://upstream.example.com/mcp")
|
||||
assert response.status_code == 401
|
||||
|
|
@ -409,7 +409,7 @@ async def test_bearer_auth_gives_up_after_a_second_401():
|
|||
refetched.append(failed)
|
||||
return "fresh-token"
|
||||
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch)
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig())
|
||||
async with httpx.AsyncClient(transport=transport, auth=auth) as client:
|
||||
response = await client.get("https://upstream.example.com/mcp")
|
||||
assert response.status_code == 401
|
||||
|
|
@ -421,7 +421,60 @@ def test_bearer_auth_rejects_sync_clients():
|
|||
async def refetch(failed: str) -> "str | None":
|
||||
return None
|
||||
|
||||
auth = ClientCredentialsBearerAuth("token", refetch)
|
||||
auth = ClientCredentialsBearerAuth("token", refetch, ClientCredentialsConfig())
|
||||
with httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200)), auth=auth) as client:
|
||||
with pytest.raises(RuntimeError):
|
||||
client.get("https://upstream.example.com/mcp")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bearer_auth_writes_the_minted_token_to_the_configured_header():
|
||||
seen: "list[dict[str, str]]" = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen.append(dict(request.headers))
|
||||
return httpx.Response(200)
|
||||
|
||||
async def refetch(failed: str) -> "str | None":
|
||||
raise AssertionError("must not refetch on success")
|
||||
|
||||
auth = ClientCredentialsBearerAuth("m2m-token", refetch, ClientCredentialsConfig(header_name="esb-oauth"))
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(handler), auth=auth) as client:
|
||||
await client.get("https://upstream.example.com/mcp")
|
||||
assert seen[0]["esb-oauth"] == "Bearer m2m-token"
|
||||
assert "authorization" not in seen[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_401_refetch_retry_also_targets_the_configured_header():
|
||||
# The retry is a SECOND write of the credential. Honoring the carrier only on the first write
|
||||
# would silently send the fresh token to Authorization, so the ESB rejects every recovered
|
||||
# request while the first attempt looked correct.
|
||||
seen: "list[dict[str, str]]" = []
|
||||
responses = [httpx.Response(401), httpx.Response(200)]
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen.append(dict(request.headers))
|
||||
return responses[min(len(seen) - 1, len(responses) - 1)]
|
||||
|
||||
async def refetch(failed: str) -> "str | None":
|
||||
return "fresh-token"
|
||||
|
||||
auth = ClientCredentialsBearerAuth("stale-token", refetch, ClientCredentialsConfig(header_name="esb-oauth"))
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(handler), auth=auth) as client:
|
||||
response = await client.get("https://upstream.example.com/mcp")
|
||||
assert response.status_code == 200
|
||||
assert [h["esb-oauth"] for h in seen] == ["Bearer stale-token", "Bearer fresh-token"]
|
||||
assert all("authorization" not in h for h in seen)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bearer_auth_advertises_the_header_it_will_occupy():
|
||||
# _resolve_v2_auth reads header_name off the auth object to decide which injected header
|
||||
# conflicts; an auth object that lies about its slot would drop the wrong one.
|
||||
async def refetch(failed: str) -> "str | None":
|
||||
return None
|
||||
|
||||
assert ClientCredentialsBearerAuth("t", refetch, ClientCredentialsConfig()).header_name == "Authorization"
|
||||
default_carrier = ClientCredentialsConfig(header_name="esb-oauth")
|
||||
assert ClientCredentialsBearerAuth("t", refetch, default_carrier).header_name == "esb-oauth"
|
||||
|
|
|
|||
|
|
@ -1033,3 +1033,70 @@ async def test_invalidate_credentials_for_id_jag_is_a_noop_without_a_caller_toke
|
|||
assert isinstance(first, Ok) and isinstance(second, Ok)
|
||||
assert _emitted(second.ok)["Authorization"] == "Bearer cached-bearer"
|
||||
assert len(endpoint.calls) == 2
|
||||
|
||||
|
||||
async def _resolve_with_carrier(kind: str, header: str):
|
||||
"""Resolve one minted-token arm whose config targets ``header``."""
|
||||
if kind == "client_credentials":
|
||||
source = _FakeM2MSource(Ok(OAuthToken(access_token="minted")))
|
||||
config = _M2M.model_copy(update={"header_name": header})
|
||||
provider = UpstreamCredentialProvider(client_credentials_source=source)
|
||||
return await provider.resolve_credentials(_SUBJECT, _spec(config))
|
||||
if kind == "token_exchange":
|
||||
exchanger = _FakeExchanger(Ok(OAuthToken(access_token="minted")))
|
||||
config = _OBO.model_copy(update={"header_name": header})
|
||||
subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-jwt"))
|
||||
provider = UpstreamCredentialProvider(token_exchanger=exchanger)
|
||||
return await provider.resolve_credentials(subject, _spec(config))
|
||||
if kind == "authorization_code":
|
||||
store = _FakeTokenStore({("alice", "s"): OAuthToken(access_token="minted")})
|
||||
provider = UpstreamCredentialProvider(oauth_token_store=store)
|
||||
return await provider.resolve_credentials(
|
||||
Subject(tenant_id="", subject_id="alice"),
|
||||
_spec(AuthorizationCodeConfig(header_name=header)),
|
||||
)
|
||||
endpoint = _FakeTokenEndpoint(
|
||||
[
|
||||
Ok(ExchangedToken(access_token="id-jag-assertion", expires_in=300)),
|
||||
Ok(ExchangedToken(access_token="minted", expires_in=300)),
|
||||
]
|
||||
)
|
||||
config = _id_jag_config().model_copy(update={"header_name": header})
|
||||
subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-id-token"))
|
||||
provider = UpstreamCredentialProvider(token_endpoint=endpoint)
|
||||
return await provider.resolve_credentials(subject, _spec(config))
|
||||
|
||||
|
||||
_MINTED_ARMS = ("client_credentials", "token_exchange", "authorization_code", "id_jag")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind", _MINTED_ARMS)
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_minted_arm_emits_its_configured_header(kind):
|
||||
# One arm left on a hardcoded Authorization is a silent no-op for exactly the server that
|
||||
# configured the knob, so this is asserted across all four rather than on the M2M arm alone.
|
||||
result = await _resolve_with_carrier(kind, "esb-oauth")
|
||||
assert isinstance(result, Ok)
|
||||
headers, _ = await _emitted_async(result.ok)
|
||||
assert headers["esb-oauth"] == "Bearer minted"
|
||||
assert "authorization" not in headers
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind", _MINTED_ARMS)
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_minted_arm_still_defaults_to_authorization(kind):
|
||||
result = await _resolve_with_carrier(kind, "Authorization")
|
||||
assert isinstance(result, Ok)
|
||||
headers, _ = await _emitted_async(result.ok)
|
||||
assert headers["Authorization"] == "Bearer minted"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passthrough_ignores_the_carrier_and_keeps_the_callers_slot():
|
||||
# Passthrough mints nothing: it forwards the caller's own credential, so it has no carrier to
|
||||
# configure and must keep using the header the caller aimed it at.
|
||||
subject = Subject(tenant_id="", subject_id="", inbound_token=SecretStr("caller-token"))
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(subject, _spec(PassthroughConfig()))
|
||||
assert isinstance(result, Ok)
|
||||
headers, _ = await _emitted_async(result.ok)
|
||||
assert headers["Authorization"] == "caller-token"
|
||||
|
|
|
|||
|
|
@ -14,9 +14,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
|||
Ambient,
|
||||
ApiKeyConfig,
|
||||
AuthConfig,
|
||||
AuthorizationCodeConfig,
|
||||
AuthSpecKind,
|
||||
AwsSigV4Config,
|
||||
Byok,
|
||||
ClientCredentialsConfig,
|
||||
ClientSecretAuth,
|
||||
CredError,
|
||||
Error,
|
||||
|
|
@ -27,7 +29,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
|||
ServerSpec,
|
||||
SharedKey,
|
||||
StaticKeys,
|
||||
TokenExchangeConfig,
|
||||
parse_auth_spec_kind,
|
||||
validate_header_name,
|
||||
)
|
||||
|
||||
_AUTH_CONFIG = TypeAdapter(AuthConfig)
|
||||
|
|
@ -229,3 +233,61 @@ def test_id_jag_server_spec_derives_auth_spec_kind():
|
|||
config=config,
|
||||
)
|
||||
assert spec.auth_spec_kind is AuthSpecKind.id_jag
|
||||
|
||||
|
||||
_CARRIER_CONFIGS = (
|
||||
("client_credentials", ClientCredentialsConfig),
|
||||
("token_exchange", lambda **kw: TokenExchangeConfig(token_exchange_endpoint="https://idp/te", **kw)),
|
||||
("authorization_code", AuthorizationCodeConfig),
|
||||
(
|
||||
"id_jag",
|
||||
lambda **kw: IdJagConfig(
|
||||
org_token_endpoint="https://idp.example.com/token",
|
||||
resource_token_endpoint="https://mcp-as.example.com/token",
|
||||
client_id="litellm",
|
||||
client_auth=ClientSecretAuth(client_secret=SecretStr("s")),
|
||||
**kw,
|
||||
),
|
||||
),
|
||||
("api_key", lambda **kw: ApiKeyConfig(key_source=SharedKey(value=SecretStr("k")), **kw)),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,build", _CARRIER_CONFIGS, ids=[n for n, _ in _CARRIER_CONFIGS])
|
||||
def test_every_resolved_credential_config_defaults_to_rfc6750_authorization(name, build):
|
||||
# The default is what preserves today's wire behavior for every existing server.
|
||||
assert build().header("tok") == ("Authorization", "Bearer tok")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,build", _CARRIER_CONFIGS, ids=[n for n, _ in _CARRIER_CONFIGS])
|
||||
def test_every_resolved_credential_config_honors_a_custom_header(name, build):
|
||||
assert build(header_name="esb-oauth").header("tok") == ("esb-oauth", "Bearer tok")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name,build", _CARRIER_CONFIGS, ids=[n for n, _ in _CARRIER_CONFIGS])
|
||||
def test_every_resolved_credential_config_can_send_a_raw_value(name, build):
|
||||
assert build(header_name="esb-oauth", value_prefix="").header("tok") == ("esb-oauth", "tok")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad",
|
||||
[
|
||||
"with space",
|
||||
"has:colon",
|
||||
"trailing\r\nX-Injected",
|
||||
"",
|
||||
" ",
|
||||
"quoted\"name",
|
||||
],
|
||||
)
|
||||
def test_header_name_outside_the_rfc7230_token_grammar_is_rejected(bad):
|
||||
# An operator-supplied name reaches egress verbatim, so anything that could split a
|
||||
# header must fail closed at construction rather than be sanitized later.
|
||||
with pytest.raises(ValidationError):
|
||||
ClientCredentialsConfig(header_name=bad)
|
||||
assert isinstance(validate_header_name(bad), Error)
|
||||
|
||||
|
||||
def test_header_name_is_trimmed_by_the_one_validator():
|
||||
assert validate_header_name(" esb-oauth ") == Ok("esb-oauth")
|
||||
assert ClientCredentialsConfig(header_name=" esb-oauth ").header_name == "esb-oauth"
|
||||
|
|
|
|||
|
|
@ -43,7 +43,6 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
|||
_obo_retry_applies,
|
||||
_resolve_openapi_tool_auth,
|
||||
_should_strip_caller_authorization,
|
||||
_without_authorization,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
|
|
@ -2405,6 +2404,104 @@ class TestMCPServerManager:
|
|||
assert client._resolved_auth is not None
|
||||
assert "authorization" not in {k.lower() for k in (client.extra_headers or {})}
|
||||
|
||||
@staticmethod
|
||||
def _esb_server(header: "str | None") -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="esb",
|
||||
name="esb-server",
|
||||
url="https://up.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="client_credentials",
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_url="https://idp.example.com/token",
|
||||
upstream_token_header=header,
|
||||
static_headers={"Authorization": "Bearer static-upstream-mcp-token"},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_authorization_survives_a_minted_token_aimed_elsewhere(self):
|
||||
"""The dual-credential case: an ESB wants the gateway-minted token on its own header while a
|
||||
separate static Authorization passes through to the origin. Dropping Authorization here (the
|
||||
old name-blind behavior) deletes the second credential and the upstream 401s."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject, server):
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED-M2M", header_name="esb-oauth"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
client = await manager._create_mcp_client(
|
||||
self._esb_server("esb-oauth"),
|
||||
extra_headers={"Authorization": "Bearer static-upstream-mcp-token"},
|
||||
)
|
||||
|
||||
assert client._resolved_auth is not None
|
||||
assert (client.extra_headers or {})["Authorization"] == "Bearer static-upstream-mcp-token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_minted_token_aimed_at_the_static_header_still_wins_that_slot(self):
|
||||
"""The negative class of the test above: when the two DO collide the resolver-owned
|
||||
credential is still authoritative, so the knob cannot be used to smuggle a second
|
||||
credential into the same slot."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject, server):
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED-M2M", header_name="esb-oauth"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
client = await manager._create_mcp_client(
|
||||
self._esb_server("esb-oauth"),
|
||||
extra_headers={"esb-oauth": "Bearer signer-jwt", "X-Trace": "keep-me"},
|
||||
)
|
||||
|
||||
assert client._resolved_auth is not None
|
||||
assert "esb-oauth" not in {k.lower() for k in (client.extra_headers or {})}
|
||||
assert (client.extra_headers or {})["X-Trace"] == "keep-me"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_differently_cased_injected_header_is_still_recognised_as_the_collision(self):
|
||||
"""HTTP header names are case-insensitive, so the conflict check must be too.
|
||||
|
||||
A case-sensitive check reports no conflict and hands the injected header back untouched, so
|
||||
the returned extra_headers still carries a second copy of the credential slot for every
|
||||
downstream consumer of that dict. httpx happens to collapse the two on the wire, which is
|
||||
exactly why this needs pinning rather than being left to luck.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject, server):
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="esb-oauth"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
client = await manager._create_mcp_client(
|
||||
self._esb_server("esb-oauth"),
|
||||
extra_headers={"ESB-OAuth": "Bearer injected", "X-Trace": "keep"},
|
||||
)
|
||||
|
||||
assert client._resolved_auth is not None
|
||||
assert not any(k.lower() == "esb-oauth" for k in (client.extra_headers or {}))
|
||||
assert (client.extra_headers or {})["X-Trace"] == "keep"
|
||||
|
||||
def test_without_header_drops_only_the_named_header(self):
|
||||
from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, without_header
|
||||
|
||||
headers = {"Authorization": "Bearer a", "esb-oauth": "Bearer b", "X-Trace": "t"}
|
||||
assert without_header(headers, "ESB-OAuth") == {"Authorization": "Bearer a", "X-Trace": "t"}
|
||||
assert without_header(headers, DEFAULT_CREDENTIAL_HEADER) == {"esb-oauth": "Bearer b", "X-Trace": "t"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_challenges_on_rejected_subject(self):
|
||||
"""A subject the IdP rejects must raise the RFC 9728 401 challenge from the preflight, so a
|
||||
|
|
@ -2624,14 +2721,16 @@ class TestMCPServerManager:
|
|||
if captured_extra_headers:
|
||||
assert "authorization" not in {k.lower() for k in captured_extra_headers}
|
||||
|
||||
def test_without_authorization_drops_only_the_credential(self):
|
||||
def test_without_header_drops_only_the_credential(self):
|
||||
from litellm.types.mcp import without_header
|
||||
|
||||
# None / empty -> None
|
||||
assert _without_authorization(None) is None
|
||||
assert _without_authorization({}) is None
|
||||
assert without_header(None, "Authorization") is None
|
||||
assert without_header({}, "Authorization") is None
|
||||
# Only Authorization present -> nothing left -> None (case-insensitive)
|
||||
assert _without_authorization({"authorization": "Bearer x"}) is None
|
||||
assert without_header({"authorization": "Bearer x"}, "Authorization") is None
|
||||
# Authorization dropped, other headers kept
|
||||
assert _without_authorization({"Authorization": "Bearer x", "X-Trace-Id": "t"}) == {"X-Trace-Id": "t"}
|
||||
assert without_header({"Authorization": "Bearer x", "X-Trace-Id": "t"}, "Authorization") == {"X-Trace-Id": "t"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header(
|
||||
|
|
@ -9641,13 +9740,38 @@ class TestMaterializeAuthHeaders:
|
|||
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
|
||||
ClientCredentialsBearerAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ClientCredentialsConfig,
|
||||
)
|
||||
|
||||
async def _refetch(_stale: str):
|
||||
return None
|
||||
|
||||
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch))
|
||||
default_carrier = ClientCredentialsConfig()
|
||||
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch, default_carrier))
|
||||
assert headers == {"Authorization": "Bearer m2m-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_materialize_follows_the_minted_token_to_a_custom_header(self):
|
||||
# The OpenAPI arm reads header_name off the auth object rather than assuming Authorization,
|
||||
# so it carries the knob with no per-arm change. This pins that it stays that way.
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_materialize_auth_headers,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
|
||||
ClientCredentialsBearerAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
ClientCredentialsConfig,
|
||||
)
|
||||
|
||||
async def _refetch(_stale: str):
|
||||
return None
|
||||
|
||||
esb_carrier = ClientCredentialsConfig(header_name="esb-oauth")
|
||||
headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch, esb_carrier))
|
||||
assert headers == {"esb-oauth": "Bearer m2m-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_noop_and_none_materialize_to_none(self):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import pytest
|
|||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
||||
MCPOAuth2TokenCache,
|
||||
resolve_mcp_auth,
|
||||
resolved_token_header,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
|
@ -411,3 +412,50 @@ async def test_m2m_mint_uses_admin_entered_token_url_when_issuer_yield_empties_r
|
|||
|
||||
assert result == "m2m-token-configured"
|
||||
assert mock_client.post.call_args[0][0] == "https://auth.example.com/token"
|
||||
|
||||
|
||||
def _m2m_server(**overrides):
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
fields = dict(
|
||||
server_id="s",
|
||||
name="n",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="client_credentials",
|
||||
client_id="cid",
|
||||
client_secret="csec",
|
||||
token_url="https://idp.example.com/token",
|
||||
)
|
||||
fields.update(overrides)
|
||||
return MCPServer(**fields)
|
||||
|
||||
|
||||
def test_resolved_token_header_follows_the_configured_header_for_a_gateway_resolved_token():
|
||||
# resolve_mcp_auth mints the M2M token on this branch, so the value is the gateway's own and
|
||||
# follows upstream_token_header.
|
||||
assert resolved_token_header(_m2m_server(upstream_token_header="esb-oauth")) == "esb-oauth"
|
||||
|
||||
|
||||
def test_resolved_token_header_is_none_when_the_server_configures_nothing():
|
||||
assert resolved_token_header(_m2m_server()) is None
|
||||
|
||||
|
||||
def test_a_caller_supplied_credential_never_moves():
|
||||
# The caller aimed their own token at the slot the upstream normally uses. Relocating it would
|
||||
# break every existing x-mcp-auth caller on a server that sets the field for its own token.
|
||||
server = _m2m_server(upstream_token_header="esb-oauth")
|
||||
assert resolved_token_header(server, "Bearer caller-token") is None
|
||||
assert resolved_token_header(server, {"Authorization": "Bearer caller-token"}) is None
|
||||
|
||||
|
||||
def test_the_header_and_the_value_agree_on_which_branch_they_took():
|
||||
# The two helpers are read as a pair at one call site, so they must never disagree about
|
||||
# whether the credential came from the caller or from the server's own config.
|
||||
import asyncio
|
||||
|
||||
server = _m2m_server(upstream_token_header="esb-oauth", authentication_token="static-tok")
|
||||
caller = "Bearer caller-token"
|
||||
assert asyncio.run(resolve_mcp_auth(server, caller)) == caller
|
||||
assert resolved_token_header(server, caller) is None
|
||||
|
|
|
|||
|
|
@ -729,3 +729,121 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st
|
|||
# A non-auth upstream failure stays a 200 with isError, so REST does not report it as a gateway 500
|
||||
assert result.isError is True
|
||||
assert "upstream returned HTTP 429" in result.content[0].text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"resolved,expect_guard",
|
||||
[
|
||||
({"esb-oauth": "Bearer minted"}, True),
|
||||
({"Authorization": "Bearer minted"}, False),
|
||||
({}, False),
|
||||
],
|
||||
)
|
||||
def test_only_a_custom_credential_slot_needs_the_redirect_guard(resolved, expect_guard):
|
||||
"""The OpenAPI arm sends resolved credentials through a redirect-following client, so a custom
|
||||
slot needs the same cross-origin guard the MCP client installs. Authorization does not: the HTTP
|
||||
client already strips that one, and taking the guarded path would give up the shared client.
|
||||
"""
|
||||
from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, same_header
|
||||
|
||||
guarded = next((n for n in resolved if not same_header(n, DEFAULT_CREDENTIAL_HEADER)), None)
|
||||
assert (guarded is not None) is expect_guard
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_openapi_arm_drops_a_custom_slot_across_origins():
|
||||
"""End to end on the hook the OpenAPI arm installs: same origin keeps the credential, a redirect
|
||||
to another host does not carry it.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.types.mcp import credential_redirect_hook
|
||||
|
||||
hook = credential_redirect_hook("https://api.example.com/v1/things", "esb-oauth")
|
||||
|
||||
same = httpx.Request("POST", "https://api.example.com/v1/other", headers={"esb-oauth": "Bearer m"})
|
||||
await hook(same)
|
||||
assert same.headers["esb-oauth"] == "Bearer m"
|
||||
|
||||
foreign = httpx.Request("POST", "https://attacker.example.com/collect", headers={"esb-oauth": "Bearer m"})
|
||||
await hook(foreign)
|
||||
assert "esb-oauth" not in foreign.headers
|
||||
|
||||
|
||||
def test_the_openapi_arm_installs_the_guard_when_a_credential_rides_a_custom_slot():
|
||||
"""Pins the wiring, not just the hook: the arm must actually build a guarded client. Testing the
|
||||
hook alone passes even if this arm never installs it.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_resolved_auth_headers,
|
||||
_upstream_client,
|
||||
)
|
||||
|
||||
token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"})
|
||||
try:
|
||||
client = _upstream_client()
|
||||
assert client.client.event_hooks["request"], "custom slot must install a redirect guard"
|
||||
finally:
|
||||
_request_resolved_auth_headers.reset(token)
|
||||
|
||||
|
||||
def test_the_guarded_client_is_reused_rather_than_built_per_call():
|
||||
"""A fresh handler per guarded call is never closed, so every OpenAPI tool call on a server that
|
||||
sets upstream_token_header would leak an httpx client and its connection pool. Both variants
|
||||
have to come from the shared cache.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_resolved_auth_headers,
|
||||
_upstream_client,
|
||||
)
|
||||
|
||||
token = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"})
|
||||
try:
|
||||
assert _upstream_client() is _upstream_client()
|
||||
finally:
|
||||
_request_resolved_auth_headers.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_shared_guard_reads_the_url_from_the_request_context():
|
||||
"""The hook is one stable object so the client stays cacheable, which means the origin it guards
|
||||
against has to arrive per request rather than being closed over.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_drop_credential_across_origin,
|
||||
_request_resolved_auth_headers,
|
||||
_request_upstream_url,
|
||||
)
|
||||
|
||||
creds = _request_resolved_auth_headers.set({"esb-oauth": "Bearer minted"})
|
||||
url = _request_upstream_url.set("https://api.example.com/v1/things")
|
||||
try:
|
||||
same = httpx.Request("POST", "https://api.example.com/v1/other", headers={"esb-oauth": "Bearer m"})
|
||||
await _drop_credential_across_origin(same)
|
||||
assert same.headers["esb-oauth"] == "Bearer m"
|
||||
|
||||
foreign = httpx.Request("POST", "https://attacker.example.com/x", headers={"esb-oauth": "Bearer m"})
|
||||
await _drop_credential_across_origin(foreign)
|
||||
assert "esb-oauth" not in foreign.headers
|
||||
finally:
|
||||
_request_upstream_url.reset(url)
|
||||
_request_resolved_auth_headers.reset(creds)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("resolved", [{"Authorization": "Bearer minted"}, {}, None])
|
||||
def test_the_openapi_arm_keeps_the_shared_client_when_no_guard_is_needed(resolved):
|
||||
# Authorization is already stripped across origins by the HTTP client, so taking the guarded
|
||||
# path for it would give up the shared connection pool for nothing.
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_resolved_auth_headers,
|
||||
_upstream_client,
|
||||
)
|
||||
|
||||
token = _request_resolved_auth_headers.set(resolved)
|
||||
try:
|
||||
client = _upstream_client()
|
||||
assert not client.client.event_hooks.get("request")
|
||||
finally:
|
||||
_request_resolved_auth_headers.reset(token)
|
||||
|
|
|
|||
|
|
@ -5274,3 +5274,74 @@ async def test_terminal_failure_logs_usage_and_cost_of_prior_passed_chunks(monke
|
|||
assert logged["guardrail_cost"] == pytest.approx(0.0003)
|
||||
assert logged["guardrail_response"]["usage"] == {"contentPolicyUnits": 2, "wordPolicyUnits": 1}
|
||||
assert "error" in logged["guardrail_response"]
|
||||
|
||||
|
||||
def test_load_credentials_assumes_role_with_external_id():
|
||||
"""A trust policy requiring sts:ExternalId must be satisfied by the guardrail's aws_external_id."""
|
||||
import datetime
|
||||
|
||||
import boto3
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
class FakeSTSClient:
|
||||
"""STS that mirrors a cross-account role whose trust policy requires an ExternalId."""
|
||||
|
||||
def get_caller_identity(self):
|
||||
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
|
||||
|
||||
def assume_role(self, **params):
|
||||
if params.get("ExternalId") != "external-id-123":
|
||||
raise ClientError(
|
||||
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:AssumeRole"}},
|
||||
"AssumeRole",
|
||||
)
|
||||
return {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "ASIAASSUMEDROLEKEY",
|
||||
"SecretAccessKey": "assumed-secret",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(minutes=30),
|
||||
}
|
||||
}
|
||||
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-external-id",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
guardrailIdentifier="gr-1",
|
||||
guardrailVersion="DRAFT",
|
||||
aws_region_name="us-east-1",
|
||||
aws_access_key_id="AKIAPODCALLERKEY",
|
||||
aws_secret_access_key="pod-caller-secret",
|
||||
aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role",
|
||||
aws_session_name="litellm-session",
|
||||
aws_external_id="external-id-123",
|
||||
)
|
||||
|
||||
with patch.object(boto3, "client", return_value=FakeSTSClient()):
|
||||
credentials, aws_region_name = guardrail._load_credentials()
|
||||
|
||||
assert credentials.access_key == "ASIAASSUMEDROLEKEY"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert aws_region_name == "us-east-1"
|
||||
|
||||
|
||||
def test_initialize_bedrock_forwards_aws_external_id():
|
||||
"""aws_external_id configured on the guardrail must survive LitellmParams and the initializer."""
|
||||
from litellm.proxy.guardrails.guardrail_initializers import initialize_bedrock
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
|
||||
litellm_params = LitellmParams(
|
||||
guardrail="bedrock",
|
||||
mode="pre_call",
|
||||
guardrailIdentifier="gr-1",
|
||||
guardrailVersion="DRAFT",
|
||||
aws_region_name="us-east-1",
|
||||
aws_role_name="arn:aws:iam::999999999999:role/litellm-guardrail-role",
|
||||
aws_external_id="external-id-123",
|
||||
)
|
||||
|
||||
guardrail = initialize_bedrock(litellm_params, {"guardrail_name": "bedrock-external-id"})
|
||||
try:
|
||||
assert guardrail.optional_params["aws_external_id"] == "external-id-123"
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_list_by_object(litellm.callbacks, guardrail)
|
||||
|
|
|
|||
|
|
@ -1573,6 +1573,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
existing_server.aws_region_name = None
|
||||
existing_server.aws_service_name = None
|
||||
existing_server.upstream_resource = None
|
||||
existing_server.upstream_token_header = None
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = existing_server
|
||||
|
|
@ -1608,6 +1609,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
existing_server.aws_region_name = None
|
||||
existing_server.aws_service_name = None
|
||||
existing_server.upstream_resource = None
|
||||
existing_server.upstream_token_header = None
|
||||
for key, value in server_overrides.items():
|
||||
setattr(existing_server, key, value)
|
||||
|
||||
|
|
@ -1639,6 +1641,23 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
assert updated.credentials["client_id"] == "client-123"
|
||||
assert updated.credentials["client_secret"] == "secret-xyz"
|
||||
|
||||
def test_upstream_token_header_is_inherited_like_other_admin_config(self):
|
||||
"""It is admin config rather than a credential, so a session server derived from an existing
|
||||
one must carry it. Miss it and the derived server silently sends its token to Authorization
|
||||
while the original sends it to the gateway's header."""
|
||||
updated = self._inherit_with({}, upstream_token_header="esb-oauth")
|
||||
|
||||
assert updated.credentials["upstream_token_header"] == "esb-oauth"
|
||||
|
||||
def test_a_supplied_upstream_token_header_does_not_read_as_a_credential(self):
|
||||
"""It is in the admin-config key set, so submitting only it must still inherit the declared
|
||||
app rather than reading as "the caller supplied real credentials"."""
|
||||
updated = self._inherit_with({"upstream_token_header": "esb-oauth"})
|
||||
|
||||
assert updated.credentials["client_id"] == "client-123"
|
||||
assert updated.credentials["client_secret"] == "secret-xyz"
|
||||
assert updated.credentials["upstream_token_header"] == "esb-oauth"
|
||||
|
||||
def test_supplied_credential_still_wins_over_inheritance(self):
|
||||
"""A caller that supplies a real credential keeps it; inheritance must not overwrite it."""
|
||||
updated = self._inherit_with({"auth_value": "caller-token"})
|
||||
|
|
@ -2256,6 +2275,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
aws_region_name=None,
|
||||
aws_service_name=None,
|
||||
upstream_resource=None,
|
||||
upstream_token_header=None,
|
||||
)
|
||||
built_server = generate_mock_mcp_server_config_record(server_id="temp-server")
|
||||
mock_manager = MagicMock()
|
||||
|
|
|
|||
|
|
@ -1767,16 +1767,20 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan(monkeypatch):
|
||||
async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan():
|
||||
"""Staleness alone stops being evidence once two hosts hold different configuration: a
|
||||
row this run never considered belongs to a deployment another host is pricing from its
|
||||
own file, and sweeping it drops that charge."""
|
||||
table = _FakeSentinelTable()
|
||||
table.seed("t", DAY, "dep-elsewhere", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
|
||||
entry = _router_entry(model_id="cfg-here", model_info=dict(_VALID_PTU))
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
|
||||
|
||||
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
|
||||
await run_scheduled_ptu_rollup(
|
||||
_prisma_for([], table),
|
||||
pod_lock_manager=_pod_lock(acquired=True),
|
||||
target_date=DAY,
|
||||
router=_router_holding(entry),
|
||||
)
|
||||
|
||||
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-elsewhere") in table.rows
|
||||
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-here") in table.rows
|
||||
|
|
@ -1784,7 +1788,7 @@ async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan(monkeypat
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged(monkeypatch):
|
||||
async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged():
|
||||
"""The accepted cost of bounding the prune, driven through the sequence that produces
|
||||
it: charge the day while the deployment exists, remove it, run the day again. Nothing
|
||||
scans it now, so nothing may judge its row, and the amount it was billed stands."""
|
||||
|
|
@ -1793,18 +1797,19 @@ async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged(
|
|||
live_row = _model_row(model_id="dep-live", model_info=ptu)
|
||||
doomed_row = _model_row(model_id="dep-doomed", model_info=ptu)
|
||||
charged_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-doomed")
|
||||
monkeypatch.setattr(
|
||||
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="cfg", model_info=dict(ptu)))
|
||||
)
|
||||
router = _router_holding(_router_entry(model_id="cfg", model_info=dict(ptu)))
|
||||
|
||||
await run_scheduled_ptu_rollup(
|
||||
_prisma_for([live_row, doomed_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
|
||||
_prisma_for([live_row, doomed_row], table),
|
||||
pod_lock_manager=_pod_lock(acquired=True),
|
||||
target_date=DAY,
|
||||
router=router,
|
||||
)
|
||||
billed = table.rows[charged_key]["ptu_flat_cost"]
|
||||
table.rows[charged_key]["updated_at"] = datetime(2020, 1, 1, tzinfo=timezone.utc)
|
||||
|
||||
await run_scheduled_ptu_rollup(
|
||||
_prisma_for([live_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
|
||||
_prisma_for([live_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY, router=router
|
||||
)
|
||||
|
||||
assert table.rows[charged_key]["ptu_flat_cost"] == billed
|
||||
|
|
@ -1843,7 +1848,7 @@ async def test_every_deployment_that_prices_is_inside_the_set_that_bounds_the_pr
|
|||
table,
|
||||
)
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(prisma)
|
||||
loaded = await ptu_rollup._load_ptu_models(prisma, router=None)
|
||||
|
||||
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
|
||||
assert loaded.scanned_ids == {"dep-a", "dep-b", "dep-unpriced"}
|
||||
|
|
@ -1859,7 +1864,7 @@ async def test_a_priced_deployment_is_in_the_bound_even_with_an_id_the_scan_skip
|
|||
_FakeSentinelTable(),
|
||||
)
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(prisma)
|
||||
loaded = await ptu_rollup._load_ptu_models(prisma, router=None)
|
||||
|
||||
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
|
||||
|
||||
|
|
@ -1873,13 +1878,13 @@ async def test_the_prune_splits_the_id_set_across_statements(monkeypatch):
|
|||
table = _FakeSentinelTable()
|
||||
ptu = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
|
||||
deployments = [_model_row(model_id=f"dep-{n}", model_info=ptu) for n in range(4)]
|
||||
monkeypatch.setattr(
|
||||
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="dep-4", model_info=dict(ptu)))
|
||||
)
|
||||
table.seed("t", DAY, "dep-3", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
|
||||
|
||||
await run_scheduled_ptu_rollup(
|
||||
_prisma_for(deployments, table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
|
||||
_prisma_for(deployments, table),
|
||||
pod_lock_manager=_pod_lock(acquired=True),
|
||||
target_date=DAY,
|
||||
router=_router_holding(_router_entry(model_id="dep-4", model_info=dict(ptu))),
|
||||
)
|
||||
|
||||
chunks = [call["model"]["in"] for call in table.delete_many_calls]
|
||||
|
|
@ -1912,140 +1917,123 @@ def _router_holding(*entries):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_declared_deployment_is_priced(monkeypatch):
|
||||
async def test_a_config_declared_deployment_is_priced():
|
||||
"""The whole point. A PTU deployment the proxy only knows from config.yaml is not in
|
||||
LiteLLM_ProxyModelTable, so a DB-only scan bills the provider's reservation to nobody."""
|
||||
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=_router_holding(entry))
|
||||
|
||||
assert [(m.model_id, m.model_name, m.team_id) for m in loaded.models] == [("cfg-1", "gpt-4o-ptu", "t")]
|
||||
assert "cfg-1" in loaded.scanned_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_database_backed_router_entry_is_not_counted_twice(monkeypatch):
|
||||
async def test_a_database_backed_router_entry_is_not_counted_twice():
|
||||
"""Every deployment loaded from the table is also in the router, flagged db_model. Pricing
|
||||
both copies would write two charges for one reservation."""
|
||||
row = _model_row(model_id="db-1", model_info=dict(_VALID_PTU))
|
||||
mirrored = _router_entry(model_id="db-1", model_info={**_VALID_PTU, "db_model": True})
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(mirrored))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["db-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_router_entry_sharing_an_id_with_the_table_is_priced_once(monkeypatch):
|
||||
"""db_model is data the router carries rather than something this module controls, so the
|
||||
id anti-join is what actually maps onto the failure: two charges under one id."""
|
||||
row = _model_row(model_id="both-1", model_info=dict(_VALID_PTU))
|
||||
unflagged = _router_entry(model_id="both-1", model_info=dict(_VALID_PTU))
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(unflagged))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["both-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_client_credential_clone_is_not_priced(monkeypatch):
|
||||
"""Supplying an api_key on a request mints a clone of the deployment under a fresh id,
|
||||
carrying the source's PTU config. Pricing it bills one reservation per distinct caller key."""
|
||||
source = _router_entry(model_id="cfg-1", model_info=dict(_VALID_PTU))
|
||||
clone = _router_entry(model_id="cfg-1-clone", model_info={**_VALID_PTU, "original_model_id": "cfg-1"})
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(source, clone))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["cfg-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_deployment_without_ptu_config_is_scanned_but_not_priced(monkeypatch):
|
||||
"""It has to stay in the scanned set or its leftover sentinel rows become unprunable."""
|
||||
entry = _router_entry(model_id="cfg-plain", model_info={"team_id": "t"})
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
|
||||
|
||||
assert loaded.models == ()
|
||||
assert "cfg-plain" in loaded.scanned_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_router_in_the_process_prices_the_database_alone(monkeypatch):
|
||||
"""The rollup is importable and callable outside a running proxy."""
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: None)
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(
|
||||
_prisma_for([_model_row(model_id="db-1", model_info=dict(_VALID_PTU))], _FakeSentinelTable())
|
||||
_prisma_for([row], _FakeSentinelTable()), router=_router_holding(mirrored)
|
||||
)
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["db-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_deployment_is_charged_end_to_end(monkeypatch):
|
||||
async def test_a_router_entry_sharing_an_id_with_the_table_is_priced_once():
|
||||
"""db_model is data the router carries rather than something this module controls, so the
|
||||
id anti-join is what actually maps onto the failure: two charges under one id."""
|
||||
row = _model_row(model_id="both-1", model_info=dict(_VALID_PTU))
|
||||
unflagged = _router_entry(model_id="both-1", model_info=dict(_VALID_PTU))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(
|
||||
_prisma_for([row], _FakeSentinelTable()), router=_router_holding(unflagged)
|
||||
)
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["both-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_client_credential_clone_is_not_priced():
|
||||
"""Supplying an api_key on a request mints a clone of the deployment under a fresh id,
|
||||
carrying the source's PTU config. Pricing it bills one reservation per distinct caller key."""
|
||||
source = _router_entry(model_id="cfg-1", model_info=dict(_VALID_PTU))
|
||||
clone = _router_entry(model_id="cfg-1-clone", model_info={**_VALID_PTU, "original_model_id": "cfg-1"})
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(
|
||||
_prisma_for([], _FakeSentinelTable()), router=_router_holding(source, clone)
|
||||
)
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["cfg-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_deployment_without_ptu_config_is_scanned_but_not_priced():
|
||||
"""It has to stay in the scanned set or its leftover sentinel rows become unprunable."""
|
||||
entry = _router_entry(model_id="cfg-plain", model_info={"team_id": "t"})
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=_router_holding(entry))
|
||||
|
||||
assert loaded.models == ()
|
||||
assert "cfg-plain" in loaded.scanned_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_router_in_the_process_prices_the_database_alone():
|
||||
"""The rollup is importable and callable outside a running proxy."""
|
||||
loaded = await ptu_rollup._load_ptu_models(
|
||||
_prisma_for([_model_row(model_id="db-1", model_info=dict(_VALID_PTU))], _FakeSentinelTable()), router=None
|
||||
)
|
||||
|
||||
assert [m.model_id for m in loaded.models] == ["db-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_deployment_is_charged_end_to_end():
|
||||
"""Through the scheduled entry point, so the charge lands in a sentinel row rather than
|
||||
stopping at the loader."""
|
||||
table = _FakeSentinelTable()
|
||||
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
|
||||
|
||||
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
|
||||
await run_scheduled_ptu_rollup(
|
||||
_prisma_for([], table),
|
||||
pod_lock_manager=_pod_lock(acquired=True),
|
||||
target_date=DAY,
|
||||
router=_router_holding(entry),
|
||||
)
|
||||
|
||||
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-1") in table.rows
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_stale_database_backed_router_entry_is_not_treated_as_config(monkeypatch):
|
||||
async def test_a_stale_database_backed_router_entry_is_not_treated_as_config():
|
||||
"""The reconcile can leave a deployment on the router after its row is gone. The id
|
||||
anti-join cannot see that one, so the flag is what keeps it from being priced as though
|
||||
config.yaml had declared it."""
|
||||
stale = _router_entry(model_id="db-gone", model_info={**_VALID_PTU, "db_model": True})
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(stale))
|
||||
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=_router_holding(stale))
|
||||
|
||||
assert loaded.models == ()
|
||||
|
||||
|
||||
def test_the_router_lookup_reads_the_proxys_own_global():
|
||||
"""Every other config test replaces this helper, so without one test driving the real
|
||||
body a typo in the module path or the attribute name leaves the whole feature dead in
|
||||
production with the suite still green."""
|
||||
import sys
|
||||
import types as _types
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_router_left_on_the_proxy_module_is_not_scanned(monkeypatch):
|
||||
"""A run scans the router its caller hands it and nothing else. Reading the proxy module's
|
||||
global instead made every run depend on whatever else in the process had set one, which
|
||||
is what a caller passing no router is asking not to happen."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
assert ptu_rollup._running_router() is None or "litellm.proxy.proxy_server" in sys.modules
|
||||
ambient = _router_holding(_router_entry(model_id="ambient-1", model_info=dict(_VALID_PTU)))
|
||||
monkeypatch.setattr(proxy_server, "llm_router", ambient, raising=False)
|
||||
|
||||
sentinel = object()
|
||||
stub = _types.SimpleNamespace(llm_router=sentinel)
|
||||
real = sys.modules.get("litellm.proxy.proxy_server")
|
||||
sys.modules["litellm.proxy.proxy_server"] = stub
|
||||
try:
|
||||
assert ptu_rollup._running_router() is sentinel
|
||||
del stub.llm_router
|
||||
assert ptu_rollup._running_router() is None
|
||||
finally:
|
||||
if real is None:
|
||||
del sys.modules["litellm.proxy.proxy_server"]
|
||||
else:
|
||||
sys.modules["litellm.proxy.proxy_server"] = real
|
||||
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=None)
|
||||
|
||||
|
||||
def test_the_router_lookup_returns_none_outside_a_proxy():
|
||||
import sys
|
||||
|
||||
real = sys.modules.pop("litellm.proxy.proxy_server", None)
|
||||
try:
|
||||
assert ptu_rollup._running_router() is None
|
||||
finally:
|
||||
if real is not None:
|
||||
sys.modules["litellm.proxy.proxy_server"] = real
|
||||
assert loaded.models == ()
|
||||
assert loaded.scanned_ids == frozenset()
|
||||
|
||||
|
||||
def test_the_prune_filter_is_a_plain_dict():
|
||||
|
|
@ -2074,7 +2062,7 @@ async def test_a_run_that_scanned_nothing_issues_no_delete_statements():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_catch_up_pass_reaches_a_config_declared_deployment(monkeypatch):
|
||||
async def test_the_catch_up_pass_reaches_a_config_declared_deployment():
|
||||
"""The catch-up shares the loader, so config deployments join it without being wired in.
|
||||
That is what prices the elapsed days of a reservation declared before today."""
|
||||
table = _FakeSentinelTable()
|
||||
|
|
@ -2084,9 +2072,10 @@ async def test_the_catch_up_pass_reaches_a_config_declared_deployment(monkeypatc
|
|||
model_id="cfg-back",
|
||||
model_info={"ptu_count": 100, "cost_per_ptu_per_hour": 0.02, "team_id": "t", "ptu_effective_from": started},
|
||||
)
|
||||
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
|
||||
|
||||
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True))
|
||||
await run_scheduled_ptu_rollup(
|
||||
_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), router=_router_holding(entry)
|
||||
)
|
||||
|
||||
charged = sorted(day for (_, day, _, model) in table.rows if model == "cfg-back")
|
||||
yesterday = (now.date() - timedelta(days=1)).isoformat()
|
||||
|
|
|
|||
|
|
@ -11245,6 +11245,35 @@ async def test_ptu_rollup_job_registered_at_startup(monkeypatch):
|
|||
assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ptu_rollup_job_hands_the_rollup_the_proxys_router(monkeypatch):
|
||||
"""The rollup prices PTU deployments declared in config.yaml, which only the router
|
||||
knows about. It takes the router as an argument, so nothing but this call site puts the
|
||||
proxy's own router in front of it: without it that half of the feature is dead."""
|
||||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||||
from litellm.proxy.spend_tracking import ptu_flat_cost_rollup
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
|
||||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import PTU_ROLLUP_JOB_ID
|
||||
|
||||
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
|
||||
calls = []
|
||||
monkeypatch.setattr(
|
||||
ptu_flat_cost_rollup,
|
||||
"run_scheduled_ptu_rollup",
|
||||
AsyncMock(side_effect=lambda *args, **kwargs: calls.append(kwargs)),
|
||||
)
|
||||
|
||||
scheduler = await _run_scheduled_background_jobs()
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
router = MagicMock()
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
await scheduler.get_job(PTU_ROLLUP_JOB_ID).func()
|
||||
|
||||
assert [call["router"] for call in calls] == [router]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ptu_rollup_job_not_registered_without_opt_in(monkeypatch):
|
||||
"""Without LITELLM_ENABLE_PTU_COST_ATTRIBUTION the rollup never runs, so no sentinel row
|
||||
|
|
|
|||
|
|
@ -93,6 +93,32 @@ class TestExtractRequestToolNames:
|
|||
"run_sql",
|
||||
]
|
||||
|
||||
def test_anthropic_openai_format_tools_forwarded_by_bridge(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "get_weather"}},
|
||||
{"name": "run_sql"},
|
||||
{"googleSearch": {}},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/messages", data) == [
|
||||
"get_weather",
|
||||
"run_sql",
|
||||
]
|
||||
|
||||
def test_anthropic_hybrid_tool_yields_every_name(self):
|
||||
data = {
|
||||
"tools": [
|
||||
{"type": "function", "name": "decoy", "function": {"name": "blocked_fn"}},
|
||||
{"type": "function", "name": "", "function": {"name": "hidden_fn"}},
|
||||
]
|
||||
}
|
||||
assert extract_request_tool_names("/v1/messages", data) == [
|
||||
"decoy",
|
||||
"blocked_fn",
|
||||
"hidden_fn",
|
||||
]
|
||||
|
||||
def test_generate_content_tools(self):
|
||||
data = {
|
||||
"tools": [
|
||||
|
|
@ -159,6 +185,34 @@ class TestCheckToolsAllowlist:
|
|||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "get_weather" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disallowed_openai_format_tool_raises_on_messages_route(self):
|
||||
token = _token(metadata={"allowed_tools": ["other_tool"]})
|
||||
body = {"tools": [{"type": "function", "function": {"name": "get_weather"}}]}
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/messages",
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "get_weather" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hybrid_tool_with_decoy_name_raises_on_messages_route(self):
|
||||
token = _token(metadata={"allowed_tools": ["decoy"]})
|
||||
body = {"tools": [{"type": "function", "name": "decoy", "function": {"name": "run_sql"}}]}
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await check_tools_allowlist(
|
||||
request_body=body,
|
||||
valid_token=token,
|
||||
team_object=None,
|
||||
route="/v1/messages",
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.tool_access_denied
|
||||
assert "run_sql" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disallowed_custom_tool_raises_on_responses_route(self):
|
||||
token = _token(metadata={"allowed_tools": ["other_tool"]})
|
||||
|
|
|
|||
|
|
@ -87,3 +87,56 @@ def test_anthropic_sonnet_1hr_cache_write_pricing(
|
|||
), f"{model_key}: long-context 1hr/5min ratio is {ratio_lc}, expected 1.6"
|
||||
else:
|
||||
assert "cache_creation_input_token_cost_above_1hr_above_200k_tokens" not in info
|
||||
|
||||
|
||||
CLAUDE_3_EXPECTED = [
|
||||
("claude-3-haiku-20240307", 5e-07),
|
||||
("claude-3-opus-20240229", 3e-05),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_key, expected_1hr", CLAUDE_3_EXPECTED)
|
||||
def test_claude_3_1hr_cache_write_pricing(model_data, model_key, expected_1hr):
|
||||
"""Haiku 3 and Opus 3 both carried Sonnet's 6e-06 1hr rate, overbilling Haiku 3
|
||||
1-hour cache writes 12x and underbilling Opus 3 5x."""
|
||||
info = model_data[model_key]
|
||||
|
||||
assert info["cache_creation_input_token_cost_above_1hr"] == expected_1hr
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_key, expected_1hr", CLAUDE_3_EXPECTED)
|
||||
def test_backup_matches_main_for_claude_3_1hr_cache_write(model_key, expected_1hr):
|
||||
json_path = os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"../../litellm/model_prices_and_context_window_backup.json",
|
||||
)
|
||||
with open(json_path) as f:
|
||||
backup = json.load(f)
|
||||
|
||||
assert (
|
||||
backup[model_key]["cache_creation_input_token_cost_above_1hr"] == expected_1hr
|
||||
)
|
||||
|
||||
|
||||
def test_first_party_anthropic_1hr_cache_writes_are_2x_base_input(model_data):
|
||||
"""Anthropic charges 1-hour cache writes at 2x base input for every first-party
|
||||
model, so any entry that drifts off that multiple is a copy-paste error."""
|
||||
offenders = tuple(
|
||||
(
|
||||
model_key,
|
||||
info["input_cost_per_token"],
|
||||
info["cache_creation_input_token_cost_above_1hr"],
|
||||
)
|
||||
for model_key, info in model_data.items()
|
||||
if isinstance(info, dict)
|
||||
and info.get("litellm_provider") == "anthropic"
|
||||
and info.get("input_cost_per_token")
|
||||
and info.get("cache_creation_input_token_cost_above_1hr")
|
||||
and abs(
|
||||
info["cache_creation_input_token_cost_above_1hr"]
|
||||
- 2 * info["input_cost_per_token"]
|
||||
)
|
||||
> 1e-12
|
||||
)
|
||||
|
||||
assert offenders == (), f"1hr cache write is not 2x base input for: {offenders}"
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ import os
|
|||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import requests
|
||||
|
||||
_CODE_COVERAGE_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)), "..", "code_coverage_tests"
|
||||
)
|
||||
|
|
@ -122,6 +124,75 @@ def test_get_license_returns_none_on_request_failure(monkeypatch):
|
|||
assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None
|
||||
|
||||
|
||||
def test_get_license_retries_connection_error_then_resolves_license():
|
||||
responses = iter(
|
||||
(
|
||||
requests.ConnectionError("connection reset"),
|
||||
requests.ConnectionError("connection reset"),
|
||||
_FakeResponse({"info": {"license_expression": "MIT"}}),
|
||||
)
|
||||
)
|
||||
calls = []
|
||||
sleeps = []
|
||||
|
||||
def _fake_get(url, timeout=None):
|
||||
calls.append((url, timeout))
|
||||
response = next(responses)
|
||||
if isinstance(response, Exception):
|
||||
raise response
|
||||
return response
|
||||
|
||||
checker = check_licenses.LicenseChecker(
|
||||
config_file=_LICCHECK_INI,
|
||||
http_get=_fake_get,
|
||||
sleep=sleeps.append,
|
||||
)
|
||||
|
||||
assert checker.get_package_license_from_pypi("pkg", "1.0.0") == "MIT"
|
||||
assert len(calls) == 3
|
||||
assert len(sleeps) == 2
|
||||
|
||||
|
||||
def test_get_license_does_not_retry_not_found_http_error():
|
||||
response = requests.Response()
|
||||
response.status_code = 404
|
||||
calls = []
|
||||
sleeps = []
|
||||
|
||||
def _fake_get(url, timeout=None):
|
||||
calls.append((url, timeout))
|
||||
raise requests.HTTPError("not found", response=response)
|
||||
|
||||
checker = check_licenses.LicenseChecker(
|
||||
config_file=_LICCHECK_INI,
|
||||
http_get=_fake_get,
|
||||
sleep=sleeps.append,
|
||||
)
|
||||
|
||||
assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None
|
||||
assert len(calls) == 1
|
||||
assert sleeps == []
|
||||
|
||||
|
||||
def test_get_license_returns_none_after_connection_retry_limit():
|
||||
calls = []
|
||||
sleeps = []
|
||||
|
||||
def _fake_get(url, timeout=None):
|
||||
calls.append((url, timeout))
|
||||
raise requests.ConnectionError("connection reset")
|
||||
|
||||
checker = check_licenses.LicenseChecker(
|
||||
config_file=_LICCHECK_INI,
|
||||
http_get=_fake_get,
|
||||
sleep=sleeps.append,
|
||||
)
|
||||
|
||||
assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None
|
||||
assert len(calls) == 3
|
||||
assert len(sleeps) == 2
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# is_license_acceptable: SPDX identifiers and compound expressions
|
||||
# --------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -3977,6 +3977,74 @@ def test_completion_cost_prices_anthropic_shaped_cache_read_tokens(_local_model_
|
|||
assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9)
|
||||
|
||||
|
||||
def _together_chat_response(model: str, prompt_tokens: int, completion_tokens: int, cached_tokens: int) -> ModelResponse:
|
||||
return ModelResponse(
|
||||
id="chatcmpl-together-cache",
|
||||
choices=[{"finish_reason": "stop", "index": 0, "message": {"content": "acknowledged", "role": "assistant"}}],
|
||||
created=1756164000,
|
||||
model=model,
|
||||
object="chat.completion",
|
||||
usage=Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_completion_cost_prices_together_cached_tokens_at_cache_read_rate(_local_model_cost_map):
|
||||
"""Regression: Together reports prompt_tokens_details.cached_tokens but no together_ai
|
||||
registry entry carried cache_read_input_token_cost, so cache-hit tokens were priced at
|
||||
0.0 and spend on cache-heavy workloads was understated."""
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="deepseek-ai/DeepSeek-V4-Flash-0731", prompt_tokens=7864, completion_tokens=16, cached_tokens=7863
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(1 * 1.4e-07 + 7863 * 3e-08 + 16 * 2.8e-07, rel=1e-9)
|
||||
|
||||
|
||||
def test_completion_cost_together_mapped_model_skips_size_bucket(_local_model_cost_map):
|
||||
"""Regression: any together model whose name matches (\\d+b) was rewritten to a
|
||||
together-ai-* size bucket before the registry lookup, so mapped models like
|
||||
Muse-Glimmer-30B never used their per-model rates, cache fields included."""
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="meta-models/Muse-Glimmer-30B", prompt_tokens=63, completion_tokens=16, cached_tokens=0
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(63 * 3.5e-07 + 16 * 1.5e-06, rel=1e-9)
|
||||
|
||||
|
||||
def test_completion_cost_together_unmapped_model_still_uses_size_bucket(_local_model_cost_map):
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="qwen/Qwen2-72B-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx((23 + 15) * 9e-07, rel=1e-9)
|
||||
|
||||
|
||||
def test_completion_cost_together_metadata_only_model_still_uses_size_bucket(_local_model_cost_map):
|
||||
assert "input_cost_per_token" not in litellm.model_cost["together_ai/togethercomputer/CodeLlama-34b-Instruct"]
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=_together_chat_response(
|
||||
model="togethercomputer/CodeLlama-34b-Instruct", prompt_tokens=23, completion_tokens=15, cached_tokens=0
|
||||
),
|
||||
custom_llm_provider="together_ai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx((23 + 15) * 8e-07, rel=1e-9)
|
||||
def test_select_model_name_strips_unregistered_alias_prefix(_local_model_cost_map):
|
||||
"""A router-facing model_name alias containing "/" whose leading segment is NOT a
|
||||
registered provider must not be double-prefixed into a non-existent cost key.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""
|
||||
Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro).
|
||||
Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro,
|
||||
qwen-image-3.0, qwen-image-3.0-pro).
|
||||
|
||||
Run in docker: pytest tests/test_litellm/test_dashscope_image_generation.py -v
|
||||
"""
|
||||
|
|
@ -30,6 +31,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
|||
[
|
||||
"dashscope/qwen-image-2.0",
|
||||
"dashscope/qwen-image-2.0-pro",
|
||||
"dashscope/qwen-image-3.0",
|
||||
"dashscope/qwen-image-3.0-pro",
|
||||
],
|
||||
)
|
||||
def test_get_llm_provider_returns_dashscope(model_string: str):
|
||||
|
|
@ -48,6 +51,8 @@ def test_get_llm_provider_returns_dashscope(model_string: str):
|
|||
[
|
||||
("dashscope/qwen-image-2.0", "dashscope"),
|
||||
("dashscope/qwen-image-2.0-pro", "dashscope"),
|
||||
("dashscope/qwen-image-3.0", "dashscope"),
|
||||
("dashscope/qwen-image-3.0-pro", "dashscope"),
|
||||
],
|
||||
)
|
||||
def test_get_model_info_mode_is_image_generation(
|
||||
|
|
@ -93,6 +98,19 @@ class TestDashScopeImageGenerationConfig:
|
|||
url = self.cfg.get_complete_url(custom, None, "qwen-image-2.0", {}, {})
|
||||
assert url == custom
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"chat_api_base",
|
||||
[
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1/",
|
||||
],
|
||||
)
|
||||
def test_get_complete_url_ignores_chat_compatible_mode_base(
|
||||
self, chat_api_base: str
|
||||
):
|
||||
url = self.cfg.get_complete_url(chat_api_base, None, "qwen-image-3.0", {}, {})
|
||||
assert url == DEFAULT_API_BASE
|
||||
|
||||
def test_validate_environment_sets_auth_header(self):
|
||||
headers = self.cfg.validate_environment(
|
||||
headers={},
|
||||
|
|
@ -135,6 +153,27 @@ class TestDashScopeImageGenerationConfig:
|
|||
assert messages[0]["content"][0]["text"] == "a puppy on green grass"
|
||||
assert req["parameters"]["size"] == "1024*1024"
|
||||
|
||||
@pytest.mark.parametrize("model", ["qwen-image-3.0", "qwen-image-3.0-pro"])
|
||||
def test_transform_request_qwen_image_3(self, model: str):
|
||||
req = self.cfg.transform_image_generation_request(
|
||||
model=model,
|
||||
prompt="a poster with small multilingual text",
|
||||
optional_params=self.cfg.map_openai_params(
|
||||
non_default_params={"size": "2048x2048", "n": 6},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
),
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert req["model"] == model
|
||||
assert req["input"]["messages"][0]["content"][0]["text"] == (
|
||||
"a poster with small multilingual text"
|
||||
)
|
||||
assert req["parameters"]["size"] == "2048*2048"
|
||||
assert req["parameters"]["n"] == 6
|
||||
|
||||
def test_transform_request_empty_params(self):
|
||||
req = self.cfg.transform_image_generation_request(
|
||||
model="qwen-image-2.0-pro",
|
||||
|
|
@ -238,6 +277,48 @@ class TestDashScopeImageGenerationConfig:
|
|||
assert result.data[0].url == "https://example.com/img1.png"
|
||||
assert result.data[1].url == "https://example.com/img2.png"
|
||||
|
||||
def test_transform_response_multiple_images_in_one_choice(self):
|
||||
body = {
|
||||
"output": {
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"image": "https://example.com/img1.png", "type": "image"},
|
||||
{"image": "https://example.com/img2.png", "type": "image"},
|
||||
],
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
"usage": {
|
||||
"output_width": 1024,
|
||||
"output_height": 1024,
|
||||
"output_image_count": 2,
|
||||
},
|
||||
}
|
||||
mock_resp = MagicMock(spec=httpx.Response)
|
||||
mock_resp.status_code = 200
|
||||
mock_resp.headers = {}
|
||||
mock_resp.json.return_value = body
|
||||
|
||||
result = self.cfg.transform_image_generation_response(
|
||||
model="qwen-image-3.0",
|
||||
raw_response=mock_resp,
|
||||
model_response=ImageResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert [image.url for image in result.data] == [
|
||||
"https://example.com/img1.png",
|
||||
"https://example.com/img2.png",
|
||||
]
|
||||
|
||||
def test_transform_response_raises_on_non_200_status(self):
|
||||
mock_resp = MagicMock(spec=httpx.Response)
|
||||
mock_resp.status_code = 400
|
||||
|
|
@ -294,14 +375,14 @@ class TestDashScopeImageGenerationConfig:
|
|||
)
|
||||
assert mapped["size"] == "1024*1024"
|
||||
|
||||
def test_map_openai_params_n_to_image_count(self):
|
||||
def test_map_openai_params_n_passthrough(self):
|
||||
mapped = self.cfg.map_openai_params(
|
||||
non_default_params={"n": 2},
|
||||
optional_params={},
|
||||
model="qwen-image-2.0",
|
||||
drop_params=False,
|
||||
)
|
||||
assert mapped["image_count"] == 2
|
||||
assert mapped == {"n": 2}
|
||||
|
||||
def test_map_openai_params_unknown_size_uses_asterisk(self):
|
||||
mapped = self.cfg.map_openai_params(
|
||||
|
|
@ -338,7 +419,15 @@ class TestDashScopeImageGenerationConfig:
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_litellm_image_generation_dashscope_end_to_end():
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"dashscope/qwen-image-2.0",
|
||||
"dashscope/qwen-image-3.0",
|
||||
"dashscope/qwen-image-3.0-pro",
|
||||
],
|
||||
)
|
||||
def test_litellm_image_generation_dashscope_end_to_end(model: str):
|
||||
mock_response_body = {
|
||||
"output": {
|
||||
"choices": [
|
||||
|
|
@ -374,7 +463,7 @@ def test_litellm_image_generation_dashscope_end_to_end():
|
|||
mock_post.return_value = mock_http_response
|
||||
|
||||
response = litellm.image_generation(
|
||||
model="dashscope/qwen-image-2.0",
|
||||
model=model,
|
||||
prompt="a puppy playing on green grass",
|
||||
api_key="sk-test-key",
|
||||
size="1024x1024",
|
||||
|
|
@ -392,7 +481,7 @@ def test_litellm_image_generation_dashscope_end_to_end():
|
|||
called_url = (
|
||||
call_args[0][0] if call_args[0] else call_args.kwargs.get("url", "")
|
||||
)
|
||||
assert "dashscope" in called_url or "aliyuncs" in called_url
|
||||
assert called_url == DEFAULT_API_BASE
|
||||
|
||||
# Verify request body contains DashScope format
|
||||
call_kwargs = call_args[1] if call_args[1] else {}
|
||||
|
|
@ -400,3 +489,4 @@ def test_litellm_image_generation_dashscope_end_to_end():
|
|||
body = call_kwargs["json"]
|
||||
assert "input" in body
|
||||
assert "messages" in body["input"]
|
||||
assert body["parameters"]["size"] == "1024*1024"
|
||||
|
|
|
|||
86
tests/test_litellm/test_fireworks_serverless_model_costs.py
Normal file
86
tests/test_litellm/test_fireworks_serverless_model_costs.py
Normal file
|
|
@ -0,0 +1,86 @@
|
|||
"""
|
||||
Validate the Fireworks AI Serverless entry added for #37274 exists in
|
||||
`model_prices_and_context_window.json` and that the bare Fireworks model ID
|
||||
resolves through `get_model_info`.
|
||||
|
||||
Pricing as published at https://docs.fireworks.ai/serverless/pricing
|
||||
(USD per 1M tokens, uncached input / cached input / output):
|
||||
|
||||
accounts/fireworks/models/deepseek-v4-pro-0813 -> $1.32 / $0.044 / $3.96
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def _local_model_cost_map():
|
||||
"""
|
||||
Point litellm at the bundled cost map for the duration of this module
|
||||
only. ``mp.undo()`` restores both the environment variable and
|
||||
``litellm.model_cost`` so nothing leaks into later tests.
|
||||
"""
|
||||
mp = pytest.MonkeyPatch()
|
||||
mp.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
mp.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
get_model_info.cache_clear()
|
||||
yield
|
||||
mp.undo()
|
||||
get_model_info.cache_clear()
|
||||
|
||||
|
||||
NEW_ENTRIES = {
|
||||
"fireworks_ai/accounts/fireworks/models/deepseek-v4-pro-0813": {
|
||||
"input_cost_per_token": 1.32e-06,
|
||||
"cache_read_input_token_cost": 4.4e-08,
|
||||
"output_cost_per_token": 3.96e-06,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 131072,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_data():
|
||||
json_path = os.path.join(
|
||||
os.path.dirname(__file__), "../../model_prices_and_context_window.json"
|
||||
)
|
||||
with open(json_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def test_fireworks_serverless_entries_exist(model_data):
|
||||
"""The new prefixed entry carries the pricing and metadata from #37274."""
|
||||
for key, expected in NEW_ENTRIES.items():
|
||||
assert key in model_data, f"{key} is missing from model_prices_and_context_window.json"
|
||||
entry = model_data[key]
|
||||
for field, value in expected.items():
|
||||
assert entry[field] == pytest.approx(value), f"{key}.{field}"
|
||||
assert entry["litellm_provider"] == "fireworks_ai"
|
||||
assert entry["mode"] == "chat"
|
||||
assert entry["supports_function_calling"] is True
|
||||
assert entry["supports_vision"] is False
|
||||
|
||||
|
||||
def test_bare_fireworks_ids_resolve_through_prefixed_entries():
|
||||
"""Bare IDs from #37274 resolve via the provider-prefix lookup path."""
|
||||
for bare_id, prefixed_key in [
|
||||
(
|
||||
"accounts/fireworks/models/deepseek-v4-pro-0813",
|
||||
"fireworks_ai/accounts/fireworks/models/deepseek-v4-pro-0813",
|
||||
),
|
||||
]:
|
||||
info = get_model_info(model=bare_id, custom_llm_provider="fireworks_ai")
|
||||
expected = NEW_ENTRIES[prefixed_key]
|
||||
assert info.get("key") == prefixed_key
|
||||
assert info["litellm_provider"] == "fireworks_ai"
|
||||
assert info["input_cost_per_token"] == pytest.approx(expected["input_cost_per_token"])
|
||||
assert info["cache_read_input_token_cost"] == pytest.approx(expected["cache_read_input_token_cost"])
|
||||
assert info["output_cost_per_token"] == pytest.approx(expected["output_cost_per_token"])
|
||||
assert info["max_input_tokens"] == expected["max_input_tokens"]
|
||||
assert info["max_output_tokens"] == expected["max_output_tokens"]
|
||||
|
|
@ -15,10 +15,9 @@ COST_MAP_ADAPTER: Final = TypeAdapter(CostMap)
|
|||
SERVERLESS_CHAT_MODELS: Final = (
|
||||
"together_ai/moonshotai/Kimi-K3",
|
||||
"together_ai/zai-org/GLM-5.2",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro",
|
||||
"together_ai/zai-org/GLM-5.3-Flash",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro-0813",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Flash-0731",
|
||||
"together_ai/moonshotai/Kimi-K2.7-Code",
|
||||
"together_ai/MiniMaxAI/MiniMax-M3",
|
||||
"together_ai/thinkingmachines/Inkling",
|
||||
"together_ai/thinkingmachines/Inkling-Small",
|
||||
|
|
@ -27,20 +26,22 @@ SERVERLESS_CHAT_MODELS: Final = (
|
|||
"together_ai/Qwen/Qwen3.7-Plus",
|
||||
"together_ai/Qwen/Qwen3.6-Plus",
|
||||
"together_ai/Qwen/Qwen3.5-9B",
|
||||
"together_ai/nvidia/nemotron-3-ultra-550b-a55b",
|
||||
"together_ai/meta-models/Muse-Glimmer-30B",
|
||||
"together_ai/google/gemma-4-31B-it",
|
||||
"together_ai/pearl-ai/gemma-4-31b-it",
|
||||
"together_ai/google/gemma-3n-E4B-it",
|
||||
"together_ai/arize-ai/qwen-2-1.5b-instruct",
|
||||
"together_ai/Prism-ML/Ternary-Bonsai-27B",
|
||||
"together_ai/meta-llama/Llama-Guard-4-12B",
|
||||
"together_ai/openai/gpt-oss-120b",
|
||||
"together_ai/openai/gpt-oss-20b",
|
||||
"together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo",
|
||||
)
|
||||
|
||||
DEPRECATED_MODELS: Final = {
|
||||
"together_ai/nvidia/nemotron-3-ultra-550b-a55b": "2026-08-27",
|
||||
"together_ai/pearl-ai/gemma-4-31b-it": "2026-08-27",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro": "2026-08-27",
|
||||
"together_ai/moonshotai/Kimi-K2.7-Code": "2026-08-27",
|
||||
"together_ai/google/gemma-3n-E4B-it": "2026-08-25",
|
||||
"together_ai/meta-llama/Llama-Guard-4-12B": "2026-08-25",
|
||||
"together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": "2026-07-10",
|
||||
"together_ai/Qwen/Qwen3.5-397B-A17B": "2026-06-29",
|
||||
"together_ai/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": "2026-06-04",
|
||||
|
|
@ -110,6 +111,22 @@ def test_together_glm_52_pricing(cost_map: CostMap):
|
|||
assert info["supports_reasoning"] is True
|
||||
|
||||
|
||||
def test_together_glm_53_flash_pricing_and_capabilities(cost_map: CostMap):
|
||||
info = cost_map["together_ai/zai-org/GLM-5.3-Flash"]
|
||||
assert info["input_cost_per_token"] == 1.5e-07
|
||||
assert info["output_cost_per_token"] == 5e-07
|
||||
assert info["cache_read_input_token_cost"] == 3e-08
|
||||
assert info["max_input_tokens"] == 1048575
|
||||
assert info["max_output_tokens"] == 1048575
|
||||
assert info["supports_function_calling"] is True
|
||||
assert info["supports_parallel_function_calling"] is True
|
||||
assert info["supports_prompt_caching"] is True
|
||||
assert info["supports_tool_choice"] is True
|
||||
assert info["supports_response_schema"] is True
|
||||
assert info["supports_vision"] is True
|
||||
assert info["supports_reasoning"] is True
|
||||
|
||||
|
||||
def test_together_multilingual_e5_embedding_entry(cost_map: CostMap):
|
||||
info = cost_map["together_ai/intfloat/multilingual-e5-large-instruct"]
|
||||
assert info["mode"] == "embedding"
|
||||
|
|
@ -159,3 +176,51 @@ def test_together_backup_cost_map_in_sync(cost_map: CostMap):
|
|||
together_main = {k: v for k, v in cost_map.items() if k.startswith("together_ai/")}
|
||||
together_backup = {k: v for k, v in backup.items() if k.startswith("together_ai/")}
|
||||
assert together_backup == together_main
|
||||
|
||||
|
||||
CACHED_INPUT_MODELS: Final = (
|
||||
"together_ai/moonshotai/Kimi-K3",
|
||||
"together_ai/zai-org/GLM-5.2",
|
||||
"together_ai/meta-models/Muse-Glimmer-30B",
|
||||
"together_ai/Qwen/Qwen3.8-2.4T-A95B",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro-0813",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Flash-0731",
|
||||
"together_ai/thinkingmachines/Inkling",
|
||||
"together_ai/MiniMaxAI/MiniMax-M3",
|
||||
"together_ai/thinkingmachines/Inkling-Small",
|
||||
"together_ai/moonshotai/Kimi-K2.7-Code",
|
||||
"together_ai/deepseek-ai/DeepSeek-V4-Pro",
|
||||
"together_ai/nvidia/nemotron-3-ultra-550b-a55b",
|
||||
"together_ai/Qwen/Qwen3.7-Max",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", CACHED_INPUT_MODELS)
|
||||
def test_together_cached_input_model_carries_cache_read_pricing(cost_map: CostMap, model: str):
|
||||
info = cost_map.get(model)
|
||||
assert info is not None, f"{model} missing from model_prices_and_context_window.json"
|
||||
assert info.get("supports_prompt_caching") is True
|
||||
cache_read = info.get("cache_read_input_token_cost")
|
||||
assert isinstance(cache_read, float)
|
||||
assert 0 < cache_read < info["input_cost_per_token"]
|
||||
assert "cache_creation_input_token_cost" not in info
|
||||
|
||||
|
||||
def test_together_prompt_caching_flag_implies_cache_read_rate(cost_map: CostMap):
|
||||
for model, info in cost_map.items():
|
||||
if model.startswith("together_ai/") and info.get("supports_prompt_caching"):
|
||||
assert "cache_read_input_token_cost" in info, f"{model} flags caching without a cache read rate"
|
||||
|
||||
|
||||
def test_together_deepseek_v4_flash_cache_read_rate(cost_map: CostMap):
|
||||
info = cost_map["together_ai/deepseek-ai/DeepSeek-V4-Flash-0731"]
|
||||
assert info["input_cost_per_token"] == 1.4e-07
|
||||
assert info["cache_read_input_token_cost"] == 3e-08
|
||||
assert info["output_cost_per_token"] == 2.8e-07
|
||||
|
||||
|
||||
def test_together_qwen_37_max_repriced_to_current_together_rate(cost_map: CostMap):
|
||||
info = cost_map["together_ai/Qwen/Qwen3.7-Max"]
|
||||
assert info["input_cost_per_token"] == 2.5e-06
|
||||
assert info["output_cost_per_token"] == 7.5e-06
|
||||
assert info["cache_read_input_token_cost"] == 5e-07
|
||||
|
|
|
|||
87
tests/test_litellm/types/test_mcp.py
Normal file
87
tests/test_litellm/types/test_mcp.py
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
"""Tests for the shared MCP header primitives.
|
||||
|
||||
``same_header`` / ``has_header`` / ``without_header`` are the one owner of "is this the credential's
|
||||
header", used by both MCP stacks and the upstream-credential resolver. They live here rather than in
|
||||
either stack because a second implementation is exactly how an injected header came to shadow a
|
||||
resolved credential on one path and not the other.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.types.mcp import (
|
||||
credential_redirect_hook,
|
||||
crosses_origin,
|
||||
has_header,
|
||||
same_header,
|
||||
without_header,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"a,b,expected",
|
||||
[
|
||||
("Authorization", "authorization", True),
|
||||
("ESB-OAuth", "esb-oauth", True),
|
||||
("esb-oauth", "esb-oauth", True),
|
||||
("esb-oauth", "esb_oauth", False),
|
||||
("esb-oauth", "Authorization", False),
|
||||
],
|
||||
)
|
||||
def test_header_names_compare_case_insensitively(a: str, b: str, expected: bool) -> None:
|
||||
# RFC 7230 3.2. Every consumer of a credential slot routes through this, so a case-sensitive
|
||||
# comparison anywhere would let an injected header shadow a resolved credential.
|
||||
assert same_header(a, b) is expected
|
||||
|
||||
|
||||
def test_without_header_drops_every_casing_and_keeps_the_rest() -> None:
|
||||
headers = {"ESB-OAuth": "injected", "esb-oauth": "also injected", "X-Trace": "keep"}
|
||||
assert without_header(headers, "esb-oauth") == {"X-Trace": "keep"}
|
||||
|
||||
|
||||
def test_without_header_collapses_to_none_when_nothing_remains() -> None:
|
||||
assert without_header({"Authorization": "Bearer x"}, "AUTHORIZATION") is None
|
||||
assert without_header(None, "esb-oauth") is None
|
||||
assert without_header({}, "esb-oauth") is None
|
||||
|
||||
|
||||
def test_has_header_matches_any_casing() -> None:
|
||||
assert has_header({"ESB-OAuth": "v"}, "esb-oauth") is True
|
||||
assert has_header({"X-Other": "v"}, "esb-oauth") is False
|
||||
assert has_header(None, "esb-oauth") is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"target,expected",
|
||||
[
|
||||
("https://upstream.example.com/other", False), # same origin
|
||||
("https://upstream.example.com:443/other", False), # explicit default port
|
||||
("https://attacker.example.com/collect", True), # different host
|
||||
("http://upstream.example.com/collect", True), # scheme downgrade, same host
|
||||
("https://upstream.example.com:8443/other", True), # different port, same host
|
||||
("https://sub.upstream.example.com/x", True), # different host
|
||||
],
|
||||
)
|
||||
def test_origin_is_scheme_host_and_port_not_host_alone(target: str, expected: bool) -> None:
|
||||
assert crosses_origin("https://upstream.example.com/mcp", target) is expected
|
||||
|
||||
|
||||
def test_an_https_upgrade_of_the_same_host_is_not_crossing() -> None:
|
||||
# HTTP clients exempt this when deciding to keep Authorization, so a credential slot that did
|
||||
# not would lose the credential on every such redirect.
|
||||
assert crosses_origin("http://upstream.example.com/mcp", "https://upstream.example.com/x") is False
|
||||
assert crosses_origin("http://upstream.example.com/mcp", "http://upstream.example.com/x") is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_hook_drops_the_slot_only_once_the_origin_changes() -> None:
|
||||
import httpx
|
||||
|
||||
hook = credential_redirect_hook("https://upstream.example.com/mcp", "esb-oauth")
|
||||
|
||||
same = httpx.Request("GET", "https://upstream.example.com/other", headers={"esb-oauth": "Bearer x"})
|
||||
await hook(same)
|
||||
assert same.headers["esb-oauth"] == "Bearer x"
|
||||
|
||||
foreign = httpx.Request("GET", "https://attacker.example.com/x", headers={"esb-oauth": "Bearer x"})
|
||||
await hook(foreign)
|
||||
assert "esb-oauth" not in foreign.headers
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue