Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_cli_login_team_option

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

# Conflicts:
#	tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py
This commit is contained in:
jesus 2026-09-10 05:00:41 +00:00
commit 422246978d
230 changed files with 16659 additions and 1961 deletions

View file

@ -0,0 +1,15 @@
-- DropForeignKey
DO $$
BEGIN
IF EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_JWTKeyMapping_token_fkey') THEN
ALTER TABLE "LiteLLM_JWTKeyMapping" DROP CONSTRAINT "LiteLLM_JWTKeyMapping_token_fkey";
END IF;
END $$;
-- AddForeignKey
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_JWTKeyMapping_token_fkey') THEN
ALTER TABLE "LiteLLM_JWTKeyMapping" ADD CONSTRAINT "LiteLLM_JWTKeyMapping_token_fkey" FOREIGN KEY ("token") REFERENCES "LiteLLM_VerificationToken"("token") ON DELETE CASCADE ON UPDATE CASCADE;
END IF;
END $$;

View file

@ -492,7 +492,7 @@ model LiteLLM_JWTKeyMapping {
updated_at DateTime @default(now()) @updatedAt
updated_by String?
litellm_verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token])
litellm_verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token], onDelete: Cascade)
@@unique([jwt_claim_name, jwt_claim_value])
@@index([jwt_claim_name, jwt_claim_value, is_active])

View file

@ -13,6 +13,7 @@ from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
from litellm.llms.bedrock.base_aws_llm import run_aws_signing
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -45,7 +46,8 @@ class BedrockAgentCoreA2AHandler:
Returns:
A2A JSON-RPC response dict from the AgentCore agent
"""
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
url, headers, body = await run_aws_signing(
BedrockAgentCoreA2ATransformation.get_url_and_signed_request,
request_id=request_id,
params=params,
litellm_params=litellm_params,
@ -91,7 +93,8 @@ class BedrockAgentCoreA2AHandler:
Yields:
A2A streaming response events from the AgentCore agent
"""
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
url, headers, body = await run_aws_signing(
BedrockAgentCoreA2ATransformation.get_url_and_signed_request,
request_id=request_id,
params=params,
litellm_params=litellm_params,

View file

@ -25,7 +25,7 @@ import litellm
from litellm import ModelResponse
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
responses_reasoning_item_from_thinking_blocks,
responses_reasoning_items_from_thinking_blocks,
)
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.bridges.completion_transformation import (
@ -129,8 +129,8 @@ def _reasoning_input_items(msg: "AllMessageValues") -> list[dict[str, object]]:
return stored
raw_blocks: Final = msg.get("thinking_blocks") or ()
blocks: Final = cast("Iterable[ChatCompletionThinkingBlock]", raw_blocks) # cast-ok: untyped client json
from_thinking: Final = responses_reasoning_item_from_thinking_blocks(blocks)
return [] if from_thinking is None else [dict(from_thinking)] # mutable-ok: API message payload
replayed: Final = responses_reasoning_items_from_thinking_blocks(blocks)
return [dict(item) for item in replayed] # mutable-ok: API message payload
def _build_reasoning_item(
@ -227,7 +227,7 @@ class _ChatToolCallDict(ChatCompletionToolCallChunk, total=False):
provider_specific_fields: Mapping[str, object]
def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _ChatToolCallDict:
def tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _ChatToolCallDict:
"""Convert a ``function_call`` or ``custom_tool_call`` output item dict to a chat
completions tool_call dict. Custom (grammar/freeform) tool calls carry their raw
string payload in ``input`` rather than ``arguments``; both map to
@ -755,7 +755,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
# Tool calls accumulate into the single trailing tool_calls choice
# like the typed branches above; a choice per call would hide every
# call after choices[0] from chat clients
accumulated_tool_calls.append(_tool_call_dict_from_output_item(raw_item, tool_call_index))
accumulated_tool_calls.append(tool_call_dict_from_output_item(raw_item, tool_call_index))
tool_call_index += 1
elif handle_raw_dict_callback is not None:
choice, index = handle_raw_dict_callback(item=raw_item, index=index)
@ -1409,7 +1409,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
# New output item added
output_item = parsed_chunk.get("item", {})
if output_item.get("type") in ("function_call", "custom_tool_call"):
converted: Final = _tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0))
converted: Final = tool_call_dict_from_output_item(output_item, parsed_chunk.get("output_index", 0))
provider_specific_fields: Final = converted.get("provider_specific_fields")
function_chunk: Final = ChatCompletionToolCallFunctionChunk(
@ -1484,7 +1484,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
index=0,
delta=Delta(
tool_calls=(
_tool_call_dict_from_output_item(
tool_call_dict_from_output_item(
output_item, parsed_chunk.get("output_index", 0)
),
)

View file

@ -398,6 +398,18 @@ TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS: Final = get_env_int_in_range(
minimum=1,
maximum=TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS,
)
TOKEN_COUNTER_MAX_EXACT_CHARS: Final = get_env_int_in_range(
"TOKEN_COUNTER_MAX_EXACT_CHARS",
default=4_000_000,
minimum=1,
maximum=1_000_000_000,
)
TOKEN_COUNTER_MAX_CONCURRENT_COUNTS: Final = get_env_int_in_range(
"TOKEN_COUNTER_MAX_CONCURRENT_COUNTS",
default=4,
minimum=1,
maximum=256,
)
MAX_TILE_WIDTH: Final = int(os.getenv("MAX_TILE_WIDTH", 512))
MAX_TILE_HEIGHT: Final = int(os.getenv("MAX_TILE_HEIGHT", 512))
OPENAI_FILE_SEARCH_COST_PER_1K_CALLS: Final = float(os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000))
@ -570,6 +582,7 @@ LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS: Final = float(
LOGGING_EXECUTOR_MAX_THREADS: Final = get_env_int("LOGGING_EXECUTOR_MAX_THREADS", 100)
LOGGING_EXECUTOR_MAX_PENDING_TASKS: Final = get_env_int("LOGGING_EXECUTOR_MAX_PENDING_TASKS", 10_000)
LOGGING_EXECUTOR_DROPPED_TASK_LOG_INTERVAL_SECONDS: Final = 30.0
AWS_SIGNING_MAX_THREADS: Final = 16
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE: Final = os.getenv(
"DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE", "streaming.chunk.yield"
)

View file

@ -54,18 +54,22 @@ def missing_streamable_http_client_error() -> ImportError:
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
METHOD_NOT_FOUND,
ClientResult,
GetPromptRequestParams,
GetPromptResult,
ListPromptsResult,
ListResourcesResult,
ListResourceTemplatesResult,
Prompt,
ResourceTemplate,
ServerNotification,
ServerRequest,
TextContent,
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
@ -777,8 +781,19 @@ class MCPClient:
"""List available prompts from the server."""
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
async def _list_prompts_operation(session: ClientSession):
return await session.list_prompts()
async def _list_prompts_operation(session: ClientSession) -> ListPromptsResult:
capabilities: Final = session.get_server_capabilities()
if capabilities is not None and capabilities.prompts is None:
return ListPromptsResult(prompts=[])
try:
return await session.list_prompts()
except McpError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
verbose_logger.debug(
"MCP client list_prompts is unsupported by %s: %s", self.server_url or "stdio", error
)
return ListPromptsResult(prompts=[])
try:
result: Final = await self.run_with_session(_list_prompts_operation)
@ -854,8 +869,19 @@ class MCPClient:
"""List available resources from the server."""
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
async def _list_resources_operation(session: ClientSession):
return await session.list_resources()
async def _list_resources_operation(session: ClientSession) -> ListResourcesResult:
capabilities: Final = session.get_server_capabilities()
if capabilities is not None and capabilities.resources is None:
return ListResourcesResult(resources=[])
try:
return await session.list_resources()
except McpError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
verbose_logger.debug(
"MCP client list_resources is unsupported by %s: %s", self.server_url or "stdio", error
)
return ListResourcesResult(resources=[])
try:
result: Final = await self.run_with_session(_list_resources_operation)
@ -890,8 +916,19 @@ class MCPClient:
"""List available resource templates from the server."""
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
async def _list_resource_templates_operation(session: ClientSession):
return await session.list_resource_templates()
async def _list_resource_templates_operation(session: ClientSession) -> ListResourceTemplatesResult:
capabilities: Final = session.get_server_capabilities()
if capabilities is not None and capabilities.resources is None:
return ListResourceTemplatesResult(resourceTemplates=[])
try:
return await session.list_resource_templates()
except McpError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
verbose_logger.debug(
"MCP client list_resource_templates is unsupported by %s: %s", self.server_url or "stdio", error
)
return ListResourceTemplatesResult(resourceTemplates=[])
try:
result: Final = await self.run_with_session(_list_resource_templates_operation)

View file

@ -2,6 +2,7 @@
# On success, logs events to Langfuse
import inspect
import os
import re
import traceback
from collections.abc import Callable, Iterable, Mapping
from datetime import datetime
@ -63,6 +64,44 @@ def _object_mapping(value: object) -> Mapping[str, object] | None:
return value if isinstance(value, dict) else None
def _widened_items(mapping: Mapping[str, object]) -> Iterable[tuple[object, object]]:
"""Header pairs with the key type widened back to what a caller-supplied dict can actually hold."""
return mapping.items()
def _is_session_header_trace(trace_id: object, session_id: object, proxy_server_request: object) -> bool:
if not isinstance(trace_id, str) or not isinstance(session_id, str):
return False
request: Final = _object_mapping(proxy_server_request)
raw_headers: Final = _object_mapping(request.get("headers")) if request is not None else None
if raw_headers is None:
return False
headers: Final = MappingProxyType(
{key.lower(): value for key, value in _widened_items(raw_headers) if isinstance(key, str)}
)
if headers.get("x-litellm-trace-id"):
return False
if headers.get("langfuse_trace_id") is not None:
return False
if trace_id != session_id and headers.get("langfuse_session_id") != session_id:
return False
if headers.get("x-litellm-session-id") == trace_id:
return True
if re.fullmatch(r"[a-zA-Z0-9_\-]{8,}", trace_id) is None:
return False
user_agent: Final = headers.get("user-agent")
codex: Final = isinstance(user_agent, str) and re.match(r"^codex[-_ /]", user_agent, re.IGNORECASE) is not None
return any(
value == trace_id
and (
key == "x-session-id"
or re.fullmatch(r"x-.+-session-id", key) is not None
or (codex and key in ("session-id", "session_id", "thread-id", "conversation_id"))
)
for key, value in headers.items()
)
class _UsageObject(Protocol):
"""Token-count surface the Langfuse logger reads off a response usage payload."""
@ -609,6 +648,18 @@ class LangFuseLogger:
# This allows continuing an existing trace while still returning the correct trace_id
if existing_trace_id is not None:
trace_id = existing_trace_id
resolved_trace_id: Final = (
litellm_call_id or trace_id
if existing_trace_id is None
and _is_session_header_trace(trace_id, session_id, litellm_params.get("proxy_server_request"))
else trace_id
)
if resolved_trace_id != trace_id:
verbose_logger.debug(
"Langfuse: trace_id %s came from a session header; using call id %s so each call gets its own trace",
trace_id,
resolved_trace_id,
)
requested_trace_keys: Final = _as_steering_key_sequence(clean_metadata.pop("update_trace_keys", ()))
update_trace_keys: Final = (
requested_trace_keys if _as_steering_flag(litellm.langfuse_enable_update_trace_keys) else ()
@ -663,7 +714,7 @@ class LangFuseLogger:
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
else: # don't overwrite an existing trace
trace_params = {
"id": trace_id,
"id": resolved_trace_id,
"name": trace_name,
"session_id": session_id,
"input": masked_input if not mask_input else "redacted-by-litellm",
@ -845,13 +896,13 @@ class LangFuseLogger:
# Verify langfuse accepted our trace_id; if it differs, log a warning but still return our intended value
# to match expected test behavior
if hasattr(generation_client, "trace_id") and generation_client.trace_id:
if generation_client.trace_id != trace_id:
if generation_client.trace_id != resolved_trace_id:
verbose_logger.warning(
"Langfuse trace_id mismatch: set %s, but langfuse returned %s. Using our intended trace_id for consistency.",
trace_id,
resolved_trace_id,
generation_client.trace_id,
)
return trace_id, generation_id
return resolved_trace_id, generation_id
except Exception:
verbose_logger.error("Langfuse Layer Error - %s", traceback.format_exc())
return None, None

View file

@ -24,7 +24,7 @@ from litellm.integrations.s3 import (
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
@ -366,7 +366,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Sign the request
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
await run_aws_signing(S3SigV4Auth(credentials, "s3", aws_region_name).add_auth, aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
@ -597,7 +597,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Sign the request
aws_request: Final = AWSRequest(method="GET", url=url, headers=headers)
S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
await run_aws_signing(S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth, aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())

View file

@ -22,7 +22,7 @@ from litellm.constants import (
SQS_SEND_MESSAGE_ACTION,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -295,7 +295,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
data=prepped.body,
headers=prepped.headers,
)
SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth(aws_request)
await run_aws_signing(SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth, aws_request)
signed_headers: Final = dict(aws_request.headers.items())

View file

@ -120,7 +120,6 @@ from litellm.types.utils import (
CachingDetails,
CallTypes,
CostBreakdown,
CostResponseTypes,
CustomPricingLiteLLMParams,
DynamicPromptManagementParamLiteral,
EmbeddingResponse,
@ -204,7 +203,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, LoggedRelayResponse
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
@ -2381,7 +2380,7 @@ class Logging(LiteLLMLoggingBaseClass):
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
) -> Optional["CostResponseTypes"]:
) -> Optional["LoggedRelayResponse"]:
all_chunks: Final = provider_config._convert_raw_bytes_to_str_lines(raw_bytes)
complete_streaming_response: Final = provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,

View file

@ -3,7 +3,7 @@ import json
import re
import time
import traceback
from collections.abc import Iterable, Sequence
from collections.abc import Mapping, Sequence
from typing import Final, Literal, cast
import litellm
@ -151,6 +151,16 @@ def _clear_later_replay_slice_metadata(choice: StreamingChoices) -> None:
del choice.enhancements
def _invalid_choices_message(response_object: Mapping[str, object]) -> str:
raw_keys: Final = list(response_object.keys())
if "choices" not in response_object:
return f"LiteLLM: provider returned a response with no 'choices'. Raw keys: {raw_keys}"
return (
f"LiteLLM: provider returned 'choices' that is not a list ({type(response_object['choices']).__name__}). "
f"Raw keys: {raw_keys}"
)
async def convert_to_streaming_response_async(
response_object: dict | None = None,
):
@ -179,14 +189,12 @@ async def convert_to_streaming_response_async(
choice_list: Final[list[StreamingChoices]] = []
if not response_object.get("choices"):
if not isinstance(response_object.get("choices"), list):
from litellm.exceptions import APIError
raise APIError(
status_code=500,
message=(
f"LiteLLM: provider returned a response with no 'choices'. Raw keys: {list(response_object.keys())}"
),
message=_invalid_choices_message(response_object),
llm_provider="",
model="",
)
@ -287,14 +295,12 @@ def convert_to_streaming_response(
model_response_object: Final = ModelResponseStream()
choice_list: Final[list[StreamingChoices]] = []
if not response_object.get("choices"):
if not isinstance(response_object.get("choices"), list):
from litellm.exceptions import APIError
raise APIError(
status_code=500,
message=(
f"LiteLLM: provider returned a response with no 'choices'. Raw keys: {list(response_object.keys())}"
),
message=_invalid_choices_message(response_object),
llm_provider="",
model="",
)
@ -623,15 +629,12 @@ def convert_to_model_response_object(
return convert_to_streaming_response(response_object=response_object)
choice_list: Final[list[Choices]] = []
if not response_object.get("choices") or not isinstance(response_object["choices"], Iterable):
if not isinstance(response_object.get("choices"), list):
from litellm.exceptions import APIError
raise APIError(
status_code=500,
message=(
"LiteLLM: provider returned a response with no 'choices'. "
f"Raw keys: {list(response_object.keys())}"
),
message=_invalid_choices_message(response_object),
llm_provider="",
model="",
)

View file

@ -6,7 +6,7 @@ import io
import json
import mimetypes
import re
from collections.abc import Iterable, Mapping, Sequence
from collections.abc import Iterable, Iterator, Mapping, Sequence
from itertools import groupby
from os import PathLike
from pathlib import Path
@ -1823,14 +1823,11 @@ def _extract_reasoning_content(message: dict) -> tuple[str | None, str | None]:
return None, message_content
def _readable_thinking_text(
block: ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock,
) -> str:
def _readable_thinking_text(block: Mapping[str, object]) -> str:
"""The text a chat model can read back, empty for redacted blocks and malformed ones."""
if block.get("type") != "thinking":
return ""
thinking: Final = cast(ChatCompletionThinkingBlock, block).get("thinking") # cast-ok: narrowed by the type tag
return str(thinking or "")
return str(block.get("thinking") or "")
def reasoning_content_from_thinking_blocks(
@ -1843,24 +1840,125 @@ def reasoning_content_from_thinking_blocks(
return "\n".join(text for block in thinking_blocks if (text := _readable_thinking_text(block)))
def responses_reasoning_item_from_thinking_blocks(
thinking_blocks: Iterable[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock],
) -> ChatCompletionReasoningItem | None:
"""Build a Responses API `reasoning` input item from Anthropic thinking blocks.
ENCRYPTED_REASONING_SIGNATURE_PREFIX: Final = "litellm_encrypted_reasoning:"
The item carries no `id`: the Responses API rejects an empty one and 404s on any id it
did not mint itself, while an item without an id is always accepted.
def encrypted_reasoning_signature(encrypted_content: str) -> str:
"""The opaque value a Responses API reasoning item's `encrypted_content` travels in.
Anthropic clients echo a thinking block's `signature` and a redacted block's `data`
back verbatim, so either field can carry the encrypted reasoning across turns; the
prefix tells the two apart from a signature Anthropic minted.
"""
return f"{ENCRYPTED_REASONING_SIGNATURE_PREFIX}{encrypted_content}"
def _carries_encrypted_reasoning(signature: object) -> bool:
return isinstance(signature, str) and signature.startswith(ENCRYPTED_REASONING_SIGNATURE_PREFIX)
def encrypted_content_from_signature(signature: object) -> str | None:
if not isinstance(signature, str) or not _carries_encrypted_reasoning(signature):
return None
return signature.removeprefix(ENCRYPTED_REASONING_SIGNATURE_PREFIX) or None
def _encrypted_reasoning_field(block: Mapping[str, object]) -> object:
match block.get("type"):
case "thinking":
return block.get("signature")
case "redacted_thinking":
return block.get("data")
case _:
return None
def encrypted_content_of_block(block: Mapping[str, object]) -> str | None:
return encrypted_content_from_signature(_encrypted_reasoning_field(block))
def is_encrypted_reasoning_block(block: object) -> bool:
"""A thinking or redacted_thinking block carrying Responses API encrypted reasoning.
Only the Responses API that minted the content can read it back, so an Anthropic
backend has to drop such a block rather than fail signature verification on it.
"""
if not isinstance(block, Mapping):
return False
mapping: Final = cast(Mapping[str, object], block) # cast-ok: narrowed by isinstance
return _carries_encrypted_reasoning(_encrypted_reasoning_field(mapping))
def strip_encrypted_reasoning_from_messages(messages: object) -> None:
"""Drop the bridge-tagged reasoning blocks a routed deployment cannot decrypt from
Anthropic-shaped history.
The whole block goes, the way #40280 drops undecryptable Responses ``input`` items: a
provider that did not mint the block rejects it signed (a foreign signature) and unsigned
(a missing signature) alike, so keeping its text as an unsigned thinking block only moves
the 400 from the router to the provider.
Mutates the content lists in place: the router's fallback snapshot shares these
message objects, so a rebound list would replay the stripped blocks on the fallback hop.
"""
if not isinstance(messages, list):
return
for content in _anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
_strip_encrypted_reasoning_from_blocks(content)
def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
return (
cast(list[object], content) # cast-ok: narrowed by isinstance
for message in messages
if isinstance(message, Mapping)
for content in (cast(Mapping[str, object], message).get("content"),) # cast-ok: narrowed by isinstance
if isinstance(content, list)
)
def _strip_encrypted_reasoning_from_blocks(content: object) -> None:
blocks: Final = cast(list[object], content) # cast-ok: narrowed by the caller's isinstance
kept: Final = tuple(block for block in blocks if not is_encrypted_reasoning_block(block))
blocks[:] = kept # rebind-ok: shared with fallback snapshot
def _reasoning_replay_group_key(indexed_block: tuple[int, Mapping[str, object]]) -> str:
index, block = indexed_block
return f"encrypted:{index}" if is_encrypted_reasoning_block(block) else "summary"
def _reasoning_item_from_block_group(group: tuple[Mapping[str, object], ...]) -> ChatCompletionReasoningItem | None:
summary: Final[list[ChatCompletionReasoningSummaryTextBlock]] = [ # mutable-ok: API message payload
ChatCompletionReasoningSummaryTextBlock(type="summary_text", text=text)
for block in thinking_blocks
for block in group
if (text := _readable_thinking_text(block))
]
encrypted_content: Final = encrypted_content_of_block(group[0])
if encrypted_content is not None:
return ChatCompletionReasoningItem(type="reasoning", summary=summary, encrypted_content=encrypted_content)
if not summary:
return None
return ChatCompletionReasoningItem(type="reasoning", summary=summary)
def responses_reasoning_items_from_thinking_blocks(
thinking_blocks: Iterable[Mapping[str, object]],
) -> tuple[ChatCompletionReasoningItem, ...]:
"""Build Responses API `reasoning` input items from Anthropic thinking blocks.
A block carrying encrypted reasoning replays the item it came from byte for byte;
a run of plain thinking blocks collapses into one summary-only item. No item carries
an `id`: the Responses API 404s on any id it did not mint itself and rejects an empty
one, while an item without an id is always accepted.
"""
return tuple(
item
for _, group in groupby(enumerate(thinking_blocks), key=_reasoning_replay_group_key)
if (item := _reasoning_item_from_block_group(tuple(block for _, block in group))) is not None
)
def _parse_content_for_reasoning(
message_text: str | None,
) -> tuple[str | None, str | None]:

View file

@ -46,6 +46,7 @@ from litellm.types.utils import GenericImageParsingChunk
from .common_utils import (
convert_content_list_to_str,
infer_content_type_from_url_and_content,
is_encrypted_reasoning_block,
is_non_content_values_set,
parse_tool_call_arguments,
)
@ -2299,13 +2300,16 @@ def sanitize_messages_for_tool_calling(
def _is_unsignable_thinking_block(block: object) -> bool:
"""A `thinking` block that Anthropic cannot accept on input.
"""A thinking block that Anthropic cannot accept on input.
Anthropic verifies the thinking signature cryptographically, so a block whose
signature is null, empty, or missing (e.g. from an open-source reasoning model)
is rejected with a 400 and must be dropped rather than blanked or repaired.
`redacted_thinking` blocks carry no signature and are always kept.
is rejected with a 400 and must be dropped rather than blanked or repaired, and
so is a block whose signature or data carries another provider's encrypted
reasoning. A `redacted_thinking` block Anthropic minted is always kept.
"""
if is_encrypted_reasoning_block(block):
return True
if not isinstance(block, dict) or block.get("type") != "thinking":
return False
signature: Final = block.get("signature")

View file

@ -1,6 +1,6 @@
import base64
import time
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Callable, Iterator, Mapping, Sequence
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypeAlias, TypedDict, Union, cast
@ -210,7 +210,7 @@ def apply_grounding_request_counts(
class ChunkProcessor:
def __init__(self, chunks: list, messages: list | None = None):
def __init__(self, chunks: list, messages: Sequence | None = None):
self.chunks = self._sort_chunks(chunks)
self.messages = messages
self.first_chunk = chunks[0]
@ -1004,8 +1004,9 @@ class ChunkProcessor:
chunks: Sequence["_UsageBearingChunk | ModelResponse"],
model: str,
completion_output: str,
messages: list | None = None,
messages: Sequence | None = None,
reasoning_tokens: int | None = None,
count_prompt_tokens: Callable[[], int] | None = None,
) -> Usage:
"""
Calculate usage for the given chunks.
@ -1030,7 +1031,9 @@ class ChunkProcessor:
cost: Final[float | None] = calculated_usage_per_chunk["cost"]
try:
returned_usage.prompt_tokens = prompt_tokens or token_counter(model=model, messages=messages)
returned_usage.prompt_tokens = prompt_tokens or (
count_prompt_tokens() if count_prompt_tokens else token_counter(model=model, messages=messages)
)
except Exception: # don't allow this failing to block a complete streaming response from being returned
print_verbose("token_counter failed, assuming prompt tokens is 0")
returned_usage.prompt_tokens = 0

View file

@ -1473,17 +1473,14 @@ class CustomStreamWrapper:
self.received_finish_reason = response_obj["finish_reason"]
elif self.custom_llm_provider == "cached_response":
cached_chunk: Final = cast(ModelResponseStream, chunk)
chunk_finish_reason: Final = cached_chunk.choices[0].finish_reason
cached_choice: Final = cached_chunk.choices[0] if cached_chunk.choices else None
chunk_finish_reason: Final = cached_choice.finish_reason if cached_choice is not None else None
response_obj = {
"text": cached_chunk.choices[0].delta.content,
"text": cached_choice.delta.content if cached_choice is not None else None,
"is_finished": chunk_finish_reason is not None,
"finish_reason": chunk_finish_reason,
"original_chunk": cached_chunk,
"tool_calls": (
cached_chunk.choices[0].delta.tool_calls
if hasattr(cached_chunk.choices[0].delta, "tool_calls")
else None
),
"tool_calls": (getattr(cached_choice.delta, "tool_calls", None) if cached_choice is not None else None),
}
completion_obj["content"] = response_obj["text"]

View file

@ -3,11 +3,15 @@
import base64
import io
import struct
from collections.abc import Callable, Iterable, Mapping, Sequence
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from typing import Final, Literal, cast
import anyio
import anyio.lowlevel
import httpx
import tiktoken
from tokenizers import Tokenizer
from typing_extensions import ParamSpec, TypeVar
import litellm
from litellm import verbose_logger
@ -21,7 +25,10 @@ from litellm.constants import (
MAX_TILE_HEIGHT,
MAX_TILE_WIDTH,
TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS,
TOKEN_COUNTER_MAX_CONCURRENT_COUNTS,
TOKEN_COUNTER_MAX_EXACT_CHARS,
)
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.default_encoding import encoding as default_encoding
from litellm.litellm_core_utils.url_utils import safe_get
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
@ -172,6 +179,13 @@ def calculate_tiles_needed(
return total_tiles
def high_detail_image_token_upper_bound(base_tokens: int = 85) -> int:
largest_tile_count: Final = calculate_tiles_needed(
MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES, MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES
)
return base_tokens + (base_tokens * 2) * largest_tile_count
def _unpack_ints(fmt: str, buffer: bytes) -> tuple[int, ...]:
return struct.unpack(fmt, buffer)
@ -317,6 +331,32 @@ TokenCounterFunction = Callable[[str], int]
Type for a function that counts tokens in a string.
"""
EXTRAPOLATION_SAMPLES: Final = 16
T_ParamSpec: Final = ParamSpec("T_ParamSpec")
T_Retval = TypeVar("T_Retval")
_COUNT_OFFLOAD_LIMITER: Final = anyio.lowlevel.RunVar[anyio.CapacityLimiter]("litellm_count_offload_limiter")
def _count_offload_limiter_for_this_loop() -> anyio.CapacityLimiter:
existing: Final = _COUNT_OFFLOAD_LIMITER.get(None)
if existing is not None:
return existing
created: Final = anyio.CapacityLimiter(TOKEN_COUNTER_MAX_CONCURRENT_COUNTS)
_COUNT_OFFLOAD_LIMITER.set(created)
return created
def offload_token_count(
function: Callable[T_ParamSpec, T_Retval],
) -> Callable[T_ParamSpec, Awaitable[T_Retval]]:
async def offloaded(
*args: T_ParamSpec.args,
**kwargs: T_ParamSpec.kwargs, # kwargs-ok: ParamSpec keeps the wrapped function's own keyword contract
) -> T_Retval:
return await asyncify(function, limiter=_count_offload_limiter_for_this_loop())(*args, **kwargs)
return offloaded
def _get_tiktoken_count_function(
encode_length: Callable[[str], int],
@ -538,9 +578,40 @@ def _count_extra(
return num_tokens
def _get_extrapolating_count_function(
count_exactly: TokenCounterFunction,
max_exact_chars: int = TOKEN_COUNTER_MAX_EXACT_CHARS,
) -> TokenCounterFunction:
def count_tokens(text: str) -> int:
if len(text) <= max_exact_chars:
return count_exactly(text)
samples: Final = _evenly_spaced_samples(text, max_exact_chars)
sampled_chars: Final = sum(len(sample) for sample in samples)
return round(sum(count_exactly(sample) for sample in samples) * len(text) / sampled_chars)
return count_tokens
def _evenly_spaced_samples(text: str, total_chars: int) -> tuple[str, ...]:
sample_count: Final = min(EXTRAPOLATION_SAMPLES, total_chars)
sample_chars: Final = total_chars // sample_count
last_start: Final = len(text) - sample_chars
return tuple(
text[start : start + sample_chars]
for start in (last_start * index // max(sample_count - 1, 1) for index in range(sample_count))
)
def _get_count_function(
model: str | None,
custom_tokenizer: dict | SelectTokenizerResponse | None = None,
) -> TokenCounterFunction:
return _get_extrapolating_count_function(_get_exact_count_function(model, custom_tokenizer))
def _get_exact_count_function(
model: str | None,
custom_tokenizer: dict | SelectTokenizerResponse | None = None,
) -> TokenCounterFunction:
"""
Get the function to count tokens based on the model and custom tokenizer."""
@ -549,10 +620,10 @@ def _get_count_function(
if model is not None or custom_tokenizer is not None:
tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model)
if tokenizer_json["type"] == "huggingface_tokenizer":
tokenizer: Final[Tokenizer] = tokenizer_json["tokenizer"]
def count_tokens(text: str) -> int:
enc: Final = tokenizer_json["tokenizer"].encode(text)
return len(enc.ids)
return len(tokenizer.encode_batch_fast([text])[0])
return count_tokens
elif tokenizer_json["type"] == "openai_tokenizer":

View file

@ -13,10 +13,11 @@ Pattern Overview:
"""
import json
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Mapping, MutableSequence, Sequence
from copy import deepcopy
from dataclasses import dataclass
from itertools import chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
from typing_extensions import ReadOnly, TypedDict, assert_never
@ -41,6 +42,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
merge_guardrailed_scoped_messages,
merge_returned_tools_into_request_tools,
scoped_structured_message_indices,
stream_item_field,
stream_item_fingerprint,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
@ -153,6 +155,46 @@ class ExtractedInput:
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
@dataclass(frozen=True, slots=True)
class _ToolCallShape:
name: str | None
arguments: str
@dataclass(frozen=True, slots=True)
class _SSEFieldRewrite:
"""One field of one nested section of a buffered SSE event, rewritten."""
section: str
field: str
value: object
class _SSEEventRewriter(Protocol):
def __call__(self, event: Mapping[str, object]) -> _SSEFieldRewrite | None: ...
def _rewritten_event(event: Mapping[str, object], rewrite_event: _SSEEventRewriter) -> Mapping[str, object]:
rewrite: Final = rewrite_event(event)
section: Final = None if rewrite is None else event.get(rewrite.section)
if rewrite is None or not isinstance(section, Mapping):
return event
return {**event, rewrite.section: {**section, rewrite.field: rewrite.value}} # mutable-ok: json.dumps needs a dict
def _tool_call_shapes(tool_calls: Sequence[object]) -> tuple[_ToolCallShape, ...]:
"""The guardrail-visible shape of each tool call, whether the guardrail handed
back the ``ChatCompletionMessageToolCall`` objects it was given or plain dicts."""
functions: Final = tuple(stream_item_field(tool_call, "function") for tool_call in tool_calls)
return tuple(
_ToolCallShape(
name=name if isinstance(name := stream_item_field(function, "name"), str) else None,
arguments=arguments if isinstance(arguments := stream_item_field(function, "arguments"), str) else "",
)
for function in functions
)
class _AnthropicSSEDelta(TypedDict, total=False):
type: ReadOnly[str]
text: ReadOnly[str]
@ -170,12 +212,18 @@ class AnthropicMessagesHandler(BaseTranslation):
them through guardrail rewrites; downstream provider handling is out of scope.
"""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
assembles_streamed_response = True
def __init__(self):
super().__init__()
self.adapter = LiteLLMAnthropicMessagesAdapter()
def post_call_hook_response(self, response: object) -> object:
if not isinstance(response, ModelResponse):
return response
return self.adapter.translate_openai_response_to_anthropic(response)
@staticmethod
def _build_streaming_usage_response(
responses_so_far: Sequence[object],
@ -1050,6 +1098,7 @@ class AnthropicMessagesHandler(BaseTranslation):
first_choice.message.tool_calls,
)
string_so_far = first_choice.message.content
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_list or ())
guardrail_inputs: Final = GenericGuardrailAPIInputs()
if string_so_far:
guardrail_inputs["texts"] = [string_so_far]
@ -1084,6 +1133,19 @@ class AnthropicMessagesHandler(BaseTranslation):
and guardrailed_texts[0] != string_so_far
):
self._write_ended_stream_text_rewrite(responses_so_far, guardrailed_texts[0])
if deliver_ended_stream_rewrites:
returned_tool_calls: Final = _guardrailed_inputs.get("tool_calls")
self._write_ended_stream_tool_call_rewrites(
responses_so_far,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=_tool_call_shapes(
returned_tool_calls
if isinstance(returned_tool_calls, list)
and len(returned_tool_calls) == len(pre_guardrail_tool_calls)
else tool_calls_list or ()
),
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
else:
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
return responses_so_far
@ -1206,44 +1268,124 @@ class AnthropicMessagesHandler(BaseTranslation):
@staticmethod
def _write_ended_stream_text_rewrite(
responses_so_far: list[Any], # mutable-ok: rewrites the caller's buffered chunks in place
responses_so_far: MutableSequence[object], # mutable-ok: rewrites the caller's buffered chunks in place
rewritten_text: str,
) -> None:
"""Deliver an ended-stream guardrail text rewrite by rewriting the
buffered chunks in place: the first ``text_delta`` carries the full
rewritten text and every later one is blanked, leaving the surrounding
message and content-block framing untouched. Handles both chunk formats
this stream carries (parsed event dicts and raw SSE bytes)."""
message and content-block framing untouched."""
replacements: Final = chain((rewritten_text,), repeat(""))
for idx, item in enumerate(responses_so_far):
if isinstance(item, dict):
delta = item.get("delta")
if item.get("type") == "content_block_delta" and isinstance(delta, dict):
if delta.get("type") == "text_delta":
delta["text"] = next(replacements)
elif isinstance(item, (bytes, bytearray)):
responses_so_far[idx] = ( # rebind-ok: delivers the rewrite into the caller's buffer
AnthropicMessagesHandler._rewrite_sse_text_deltas(bytes(item), replacements)
)
def rewrite_text_delta(event: Mapping[str, object]) -> _SSEFieldRewrite | None:
delta: Final = event.get("delta")
if event.get("type") != "content_block_delta" or not isinstance(delta, Mapping):
return None
if delta.get("type") != "text_delta":
return None
return _SSEFieldRewrite("delta", "text", next(replacements))
AnthropicMessagesHandler._rewrite_ended_stream_events(responses_so_far, rewrite_text_delta)
@classmethod
def _write_ended_stream_tool_call_rewrites(
cls,
responses_so_far: MutableSequence[object], # mutable-ok: rewrites the caller's buffered chunks in place
*,
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str,
) -> None:
"""Deliver ended-stream guardrail tool-call rewrites by rewriting the
buffered chunks in place: the rebuilt response lists tool calls in the
order of the stream's ``tool_use`` blocks, so the nth rewritten call lands
on the nth block, its first ``input_json_delta`` carrying the full rewritten
arguments, every later one blanked, and ``content_block_start`` carrying the
rewritten name. Blocks that do not line up with the rebuilt tool calls make
the rewrite undeliverable, so the pipeline executor discards it and releases
the original chunks."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
block_indices: Final = tuple(
index
for item in responses_so_far
for event in cls._iter_sse_events(item)
if event.get("type") == "content_block_start"
and isinstance(block := event.get("content_block"), Mapping)
and block.get("type") == "tool_use"
and isinstance(index := event.get("index"), int)
)
if len(block_indices) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
rewrites_by_block: Final = MappingProxyType(
{
index: after
for index, before, after in zip(block_indices, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
}
)
argument_replacements: Final = MappingProxyType(
{index: chain((rewrite.arguments,), repeat("")) for index, rewrite in rewrites_by_block.items()}
)
def rewrite_tool_use(event: Mapping[str, object]) -> _SSEFieldRewrite | None:
index: Final = event.get("index")
if not isinstance(index, int) or index not in rewrites_by_block:
return None
match event.get("type"):
case "content_block_start":
name: Final = rewrites_by_block[index].name
if name is None:
return None
return _SSEFieldRewrite("content_block", "name", name)
case "content_block_delta":
delta: Final = event.get("delta")
if not isinstance(delta, Mapping) or delta.get("type") != "input_json_delta":
return None
return _SSEFieldRewrite("delta", "partial_json", next(argument_replacements[index]))
case _:
return None
cls._rewrite_ended_stream_events(responses_so_far, rewrite_tool_use)
@staticmethod
def _rewrite_sse_text_deltas(sse_bytes: bytes, replacements: "Iterator[str]") -> bytes:
"""Rewrite every ``text_delta`` data line in one SSE chunk with the next
replacement text, leaving all other events and framing byte-identical."""
def _rewrite_ended_stream_events(
responses_so_far: MutableSequence[object], # mutable-ok: rewrites the caller's buffered chunks in place
rewrite_event: _SSEEventRewriter,
) -> None:
"""Replace every buffered event ``rewrite_event`` returns a rewrite for, in
both chunk formats this stream carries (parsed event dicts and raw SSE
bytes), leaving every other event and the framing untouched."""
rewritten_items: Final = tuple(
AnthropicMessagesHandler._rewrite_buffered_item(item, rewrite_event) for item in responses_so_far
)
responses_so_far[:] = rewritten_items # rebind-ok: delivers the rewrites into the caller's buffer
@staticmethod
def _rewrite_buffered_item(item: object, rewrite_event: _SSEEventRewriter) -> object:
if isinstance(item, dict):
return _rewritten_event(_as_str_mapping(item), rewrite_event)
if isinstance(item, (bytes, bytearray)):
return AnthropicMessagesHandler._rewrite_sse_events(bytes(item), rewrite_event)
return item
@staticmethod
def _rewrite_sse_events(sse_bytes: bytes, rewrite_event: _SSEEventRewriter) -> bytes:
"""Rewrite the data lines of one SSE chunk that ``rewrite_event`` rewrites,
leaving all other events and framing byte-identical."""
try:
decoded: Final = sse_bytes.decode("utf-8")
except UnicodeDecodeError:
return sse_bytes
return "\n\n".join(
AnthropicMessagesHandler._rewrite_sse_block(block, replacements) for block in decoded.split("\n\n")
"\n".join(AnthropicMessagesHandler._rewrite_sse_line(line, rewrite_event) for line in block.split("\n"))
for block in decoded.split("\n\n")
).encode("utf-8")
@staticmethod
def _rewrite_sse_block(block: str, replacements: "Iterator[str]") -> str:
return "\n".join(AnthropicMessagesHandler._rewrite_sse_line(line, replacements) for line in block.split("\n"))
@staticmethod
def _rewrite_sse_line(line: str, replacements: "Iterator[str]") -> str:
def _rewrite_sse_line(line: str, rewrite_event: _SSEEventRewriter) -> str:
if not line.startswith("data:"):
return line
try:
@ -1252,14 +1394,10 @@ class AnthropicMessagesHandler(BaseTranslation):
)
except json.JSONDecodeError:
return line
if not isinstance(data, dict) or data.get("type") != "content_block_delta":
if not isinstance(data, dict):
return line
delta: Final = data.get("delta")
if not isinstance(delta, dict) or delta.get("type") != "text_delta":
return line
return "data: " + json.dumps(
{**data, "delta": {**delta, "text": next(replacements)}} # mutable-ok: json.dumps needs plain dicts
)
rewritten: Final = _rewritten_event(_as_str_mapping(data), rewrite_event)
return line if rewritten is data else "data: " + json.dumps(rewritten)
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
stream_ended: Final = self._check_streaming_has_ended(responses_so_far)

View file

@ -21,6 +21,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_file_ids_from_messages,
is_encrypted_reasoning_block,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
@ -1201,6 +1202,32 @@ def strip_thinking_blocks_from_anthropic_messages(messages: list[Any]) -> list[A
return out
def _without_encrypted_reasoning_blocks(message: dict) -> dict | None: # mutable-ok: Anthropic message payload shape
if not isinstance(message, Mapping):
return message
content: Final = message.get("content")
if not isinstance(content, list):
return message
kept: Final = [b for b in content if not is_encrypted_reasoning_block(b)] # mutable-ok: API message payload
if len(kept) == len(content):
return message
if not kept:
return None
return {**message, "content": kept} # mutable-ok: API message payload
def strip_encrypted_reasoning_blocks_from_anthropic_messages(
messages: Sequence[dict], # mutable-ok: Anthropic message payload shape
) -> list[dict]: # mutable-ok: AnthropicMessagesRequest.messages is typed list[dict]
"""
Drop thinking / redacted_thinking blocks that carry another provider's encrypted
reasoning (a turn the Responses API bridge served) before the request reaches
Anthropic, which cannot verify them. Anthropic's own signed blocks are kept.
"""
stripped: Final = (_without_encrypted_reasoning_blocks(m) for m in messages)
return [m for m in stripped if m is not None] # mutable-ok: API message payload
def strip_thinking_blocks_from_anthropic_messages_request_dict(
data: dict[str, Any],
) -> None:

View file

@ -113,6 +113,7 @@ from litellm.litellm_core_utils.reasoning_effort_utils import (
from litellm.llms.anthropic.common_utils import (
is_empty_unsigned_thinking_block,
normalize_anthropic_tool_use_id,
strip_encrypted_reasoning_blocks_from_anthropic_messages,
)
from litellm.llms.anthropic.experimental_pass_through.context_management import (
PolyfillResult,
@ -417,7 +418,8 @@ class LiteLLMAnthropicMessagesAdapter:
model: str | None = None,
) -> list:
new_messages: Final[list[AllMessageValues]] = []
for m in messages:
replayable_messages: Final = strip_encrypted_reasoning_blocks_from_anthropic_messages(messages)
for m in replayable_messages:
user_message: ChatCompletionUserMessage | None = None
tool_message_list: list[ChatCompletionToolMessage] = []
new_user_content_list: list[ChatCompletionTextObject | ChatCompletionImageObject] = []
@ -1487,8 +1489,9 @@ class LiteLLMAnthropicMessagesAdapter:
anthropic_content.insert(0, polyfill_result.compaction_block)
## extract finish reason
openai_finish_reason: Final = response.choices[0].finish_reason if response.choices else "stop"
translated_finish_reason: Final = self._translate_openai_finish_reason_to_anthropic(
openai_finish_reason=response.choices[0].finish_reason
openai_finish_reason=openai_finish_reason
)
anthropic_finish_reason: Final = (
"refusal"

View file

@ -25,6 +25,7 @@ from ...common_utils import (
AnthropicModelInfo,
optionally_handle_anthropic_oauth,
strip_advisor_blocks_from_messages,
strip_encrypted_reasoning_blocks_from_anthropic_messages,
)
DEFAULT_ANTHROPIC_API_VERSION: Final = "2023-06-01"
@ -613,7 +614,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
messages = strip_advisor_blocks_from_messages(messages)
anthropic_messages_request: Final[AnthropicMessagesRequest] = AnthropicMessagesRequest(
messages=messages,
messages=strip_encrypted_reasoning_blocks_from_anthropic_messages(messages),
max_tokens=max_tokens,
model=model,
**anthropic_messages_optional_request_params,

View file

@ -19,6 +19,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.utils import ProviderConfigManager
from ..utils import litellm_logging_obj_from_kwargs, local_model_name
from .streaming_iterator import AnthropicResponsesStreamWrapper
@ -34,6 +35,15 @@ def _forwarded_kwargs(extra_kwargs: Mapping[str, object] | None) -> Mapping[str,
return extra_kwargs or {}
def _provider_returns_encrypted_reasoning(model: str, custom_llm_provider: object) -> bool:
provider: Final = (
custom_llm_provider if isinstance(custom_llm_provider, str) else litellm.get_llm_provider(model=model)[1]
)
provider_model: Final = local_model_name(model, provider)
responses_config: Final = ProviderConfigManager.get_provider_responses_api_config(provider, provider_model)
return responses_config is not None and "include" in responses_config.get_supported_openai_params(provider_model)
def _build_responses_kwargs(
*,
max_tokens: int,
@ -85,8 +95,13 @@ def _build_responses_kwargs(
request_data["output_format"] = output_format
anthropic_request: Final = AnthropicMessagesRequest(**request_data)
responses_kwargs: Final = _ADAPTER.translate_request(anthropic_request)
forwarded_kwargs: Final = _forwarded_kwargs(extra_kwargs)
responses_kwargs: Final = _ADAPTER.translate_request(
anthropic_request,
include_encrypted_reasoning=_provider_returns_encrypted_reasoning(
model, forwarded_kwargs.get("custom_llm_provider")
),
)
# Normalize reasoning effort based on model capabilities
# (e.g. "max" → "xhigh"/"high", "minimal" → "low" if unsupported)
@ -111,7 +126,7 @@ def _build_responses_kwargs(
responses_kwargs["stream"] = True
# Forward litellm-specific kwargs (api_key, api_base, logging obj, etc.)
excluded: Final = {"anthropic_messages"}
excluded: Final = frozenset(("anthropic_messages",))
for key, value in forwarded_kwargs.items():
if key == "litellm_logging_obj" and value is not None:
from litellm.litellm_core_utils.litellm_logging import (
@ -132,6 +147,14 @@ def _build_responses_kwargs(
if explicit_prompt_cache_key is not None:
responses_kwargs["prompt_cache_key"] = explicit_prompt_cache_key
deployment_include: Final = forwarded_kwargs.get("include")
bridge_include: Final = responses_kwargs.get("include")
if isinstance(deployment_include, list) and isinstance(bridge_include, list):
responses_kwargs["include"] = [
*bridge_include,
*(item for item in deployment_include if item not in bridge_include),
]
return responses_kwargs

View file

@ -9,13 +9,19 @@ from typing import TYPE_CHECKING, Any, Final
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.prompt_templates.common_utils import (
encrypted_reasoning_signature,
)
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
refusal_stop_details,
responses_output_refusal_text,
)
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
from .transformation import LiteLLMAnthropicToResponsesAPIAdapter
from .transformation import (
REASONING_SUMMARY_PART_SEPARATOR,
LiteLLMAnthropicToResponsesAPIAdapter,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
@ -29,9 +35,10 @@ class AnthropicResponsesStreamWrapper:
response.created -> message_start
response.output_item.added -> content_block_start (if message/function_call)
response.output_text.delta -> content_block_delta (text_delta)
response.reasoning_summary_part.added -> content_block_delta (thinking_delta separator)
response.reasoning_summary_text.delta -> content_block_delta (thinking_delta)
response.function_call_arguments.delta -> content_block_delta (input_json_delta)
response.output_item.done -> content_block_stop
response.output_item.done -> content_block_delta (signature_delta) + content_block_stop
response.completed -> message_delta + message_stop
"""
@ -94,6 +101,38 @@ class AnthropicResponsesStreamWrapper:
)
return block_idx
@staticmethod
def _field(source: object, name: str) -> object:
return source.get(name) if isinstance(source, dict) else getattr(source, name, None)
def _close_reasoning_item(self, item: object, item_id: str | None) -> None:
block_idx: Final = self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index
encrypted_content: Final = self._field(item, "encrypted_content")
signature: Final = (
encrypted_reasoning_signature(encrypted_content)
if isinstance(encrypted_content, str) and encrypted_content
else None
)
if block_idx < 0 and signature is None:
return
if block_idx < 0:
redacted_idx: Final = self._open_block(
item_id,
{"type": "redacted_thinking", "data": signature}, # mutable-ok: API message payload
)
stop: Final = {"type": "content_block_stop", "index": redacted_idx} # mutable-ok: API message payload
self._chunk_queue.append(stop)
return
if signature is not None:
self._chunk_queue.append(
{ # mutable-ok: API message payload
"type": "content_block_delta",
"index": block_idx,
"delta": {"type": "signature_delta", "signature": signature}, # mutable-ok: API message payload
}
)
self._chunk_queue.append({"type": "content_block_stop", "index": block_idx}) # mutable-ok: API message payload
def _process_event(self, event: object) -> None:
"""Convert one Responses API event into zero or more Anthropic chunks queued for emission."""
event_type = getattr(event, "type", None)
@ -175,6 +214,26 @@ class AnthropicResponsesStreamWrapper:
)
return
if event_type == "response.reasoning_summary_part.added":
part_item_id: Final = self._field(event, "item_id")
summary_index: Final = self._field(event, "summary_index")
part_block_idx: Final = (
self._item_id_to_block_index.get(part_item_id, -1) if isinstance(part_item_id, str) else -1
)
if part_block_idx < 0 or not isinstance(summary_index, int) or summary_index == 0:
return
self._chunk_queue.append(
{ # mutable-ok: API message payload
"type": "content_block_delta",
"index": part_block_idx,
"delta": { # mutable-ok: API message payload
"type": "thinking_delta",
"thinking": REASONING_SUMMARY_PART_SEPARATOR,
},
}
)
return
# ---- reasoning summary text delta ----
if event_type == "response.reasoning_summary_text.delta":
item_id = getattr(event, "item_id", None) or (event.get("item_id") if isinstance(event, dict) else None)
@ -220,6 +279,9 @@ class AnthropicResponsesStreamWrapper:
item_id = (
getattr(item, "id", None) or (item.get("id") if isinstance(item, dict) else None) if item else None
)
if self._field(item, "type") == "reasoning":
self._close_reasoning_item(item, item_id)
return
block_idx = self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index
if block_idx < 0:
return

View file

@ -13,7 +13,8 @@ from typing import Any, Final, cast
from litellm.litellm_core_utils.prompt_templates.common_utils import (
TOOL_RESULT_IMAGE_BOUNDARY,
TOOL_RESULT_IMAGE_PLACEHOLDER,
responses_reasoning_item_from_thinking_blocks,
encrypted_reasoning_signature,
responses_reasoning_items_from_thinking_blocks,
with_prompt_cache_breakpoint,
)
from litellm.litellm_core_utils.reasoning_effort_utils import (
@ -33,6 +34,7 @@ from litellm.types.llms.anthropic import (
AnthropicFinishReason,
AnthropicMessagesRequest,
AnthropicMessagesToolChoice,
AnthropicResponseContentBlockRedactedThinking,
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockThinking,
AnthropicResponseContentBlockToolUse,
@ -43,11 +45,13 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicUsage,
)
from litellm.types.llms.openai import (
ChatCompletionThinkingBlock,
ResponseAPIUsage,
ResponsesAPIResponse,
)
REASONING_SUMMARY_PART_SEPARATOR: Final = "\n\n"
RESPONSES_INCLUDE_ENCRYPTED_REASONING: Final = "reasoning.encrypted_content"
class LiteLLMAnthropicToResponsesAPIAdapter:
"""
@ -163,49 +167,55 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
return str(getattr(part, "text", None) or "")
@classmethod
def _thinking_blocks_from_reasoning_item(
def _thinking_block_from_reasoning_item(
cls,
summary: Iterable[object],
) -> tuple[dict[str, Any], ...]: # mutable-ok: API message payload
"""Anthropic thinking blocks for one Responses reasoning item.
encrypted_content: object,
) -> dict[str, Any] | None: # mutable-ok: API message payload
"""The one Anthropic block for a Responses reasoning item.
The signature stays empty: only Anthropic can sign a thinking block, and a stand-in
value would be replayed as a real one and rejected by every backend that verifies it.
The item's encrypted reasoning rides the block's opaque field (`signature`, or
`data` when there is no summary text) so the client echoes it back and the next
turn replays the very item OpenAI produced; without it the signature stays empty,
since only Anthropic can sign a thinking block.
"""
return tuple(
AnthropicResponseContentBlockThinking(
type="thinking",
thinking=text,
signature=None,
).model_dump()
for part in summary
if (text := cls._summary_part_text(part))
text: Final = REASONING_SUMMARY_PART_SEPARATOR.join(
part_text for part in summary if (part_text := cls._summary_part_text(part))
)
if not isinstance(encrypted_content, str) or not encrypted_content:
if not text:
return None
return AnthropicResponseContentBlockThinking(type="thinking", thinking=text, signature=None).model_dump()
signature: Final = encrypted_reasoning_signature(encrypted_content)
if not text:
return AnthropicResponseContentBlockRedactedThinking(type="redacted_thinking", data=signature).model_dump()
return AnthropicResponseContentBlockThinking(type="thinking", thinking=text, signature=signature).model_dump()
@staticmethod
def _assistant_block_group_key(indexed_block: tuple[int, Mapping[str, object]]) -> str:
"""Group a run of consecutive thinking blocks together; keep every other block alone."""
index, block = indexed_block
return "thinking" if block.get("type") == "thinking" else f"block:{index}"
return "thinking" if block.get("type") in ("thinking", "redacted_thinking") else f"block:{index}"
@classmethod
def _assistant_group_to_input_item(
def _assistant_group_to_input_items(
cls, group: tuple[Mapping[str, object], ...]
) -> dict[str, Any] | None: # mutable-ok: API message payload
) -> tuple[dict[str, Any], ...]: # mutable-ok: API message payload
first: Final = group[0]
btype: Final = first.get("type")
if btype == "thinking":
blocks: Final = cast(tuple[ChatCompletionThinkingBlock, ...], group) # cast-ok: untrusted client payload
reasoning_item: Final = responses_reasoning_item_from_thinking_blocks(blocks)
return None if reasoning_item is None else dict(reasoning_item) # mutable-ok: API message payload
if btype in ("thinking", "redacted_thinking"):
replayed: Final = responses_reasoning_items_from_thinking_blocks(group)
return tuple(dict(item) for item in replayed) # mutable-ok: API message payload
if btype == "tool_use":
return { # mutable-ok: API message payload
"type": "function_call",
"call_id": first.get("id", ""),
"name": first.get("name", ""),
"arguments": json.dumps(first.get("input", {})), # mutable-ok: API message payload
}
return None
return (
{ # mutable-ok: API message payload
"type": "function_call",
"call_id": first.get("id", ""),
"name": first.get("name", ""),
"arguments": json.dumps(first.get("input", {})), # mutable-ok: API message payload
},
)
return ()
def translate_messages_to_responses_input(
self,
@ -362,7 +372,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
input_items.extend(
item
for _, group in groupby(enumerate(blocks), key=self._assistant_block_group_key)
if (item := self._assistant_group_to_input_item(tuple(block for _, block in group))) is not None
for item in self._assistant_group_to_input_items(tuple(block for _, block in group))
)
asst_parts: list[dict[str, Any]] = [ # mutable-ok: API message payload
{"type": "output_text", "text": block.get("text", "")} # mutable-ok: API message payload
@ -495,10 +505,16 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
def translate_request(
self,
anthropic_request: AnthropicMessagesRequest,
include_encrypted_reasoning: bool = True,
) -> dict[str, Any]:
"""
Translate a full Anthropic /v1/messages request dict to
litellm.responses() / litellm.aresponses() kwargs.
``include_encrypted_reasoning`` asks the provider for ``reasoning.encrypted_content``
on every call, so a reasoning model's items can be replayed intact next turn even
when the client sent no ``thinking`` block; pass False for a provider whose
Responses API rejects ``include``.
"""
model: Final[str] = anthropic_request["model"]
messages_list: Final = cast(
@ -528,6 +544,8 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
"model": model,
"input": input_items,
}
if include_encrypted_reasoning:
responses_kwargs["include"] = [RESPONSES_INCLUDE_ENCRYPTED_REASONING] # mutable-ok: API request payload
if system and not developer_parts:
if isinstance(system, str):
@ -634,7 +652,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
for item in response.output:
if isinstance(item, ResponseReasoningItem):
content.extend(self._thinking_blocks_from_reasoning_item(item.summary))
reasoning_block = self._thinking_block_from_reasoning_item(item.summary, item.encrypted_content)
if reasoning_block is not None:
content.append(reasoning_block)
elif isinstance(item, ResponseOutputMessage):
for part in item.content:
@ -684,11 +704,12 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
).model_dump()
)
elif item_type == "reasoning":
content.extend(
self._thinking_blocks_from_reasoning_item(
cast(Iterable[object], item.get("summary") or ()), # cast-ok: untyped provider json
)
reasoning_block = self._thinking_block_from_reasoning_item(
cast(Iterable[object], item.get("summary") or ()), # cast-ok: untyped provider json
item.get("encrypted_content"),
)
if reasoning_block is not None:
content.append(reasoning_block)
elif item_type == "function_call":
try:
input_data = json.loads(item.get("arguments", "{}"))

View file

@ -14,6 +14,7 @@ from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.constants import DEFAULT_MAX_RETRIES
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.openai.common_utils import BaseOpenAILLM
from litellm.secret_managers.get_azure_ad_token_provider import (
@ -582,7 +583,8 @@ class BaseAzureLLM(BaseOpenAILLM):
if scope is None:
scope = "https://cognitiveservices.azure.com/.default"
max_retries: Final = litellm_params.get("max_retries")
configured_max_retries: Final = litellm_params.get("max_retries")
max_retries: Final = DEFAULT_MAX_RETRIES if configured_max_retries is None else configured_max_retries
timeout: Final = litellm_params.get("timeout")
if not api_key and azure_ad_token_provider is None and tenant_id and client_id and client_secret:
verbose_logger.debug("Using Azure AD Token Provider from Entra ID for Azure Auth")
@ -642,8 +644,7 @@ class BaseAzureLLM(BaseOpenAILLM):
else:
azure_client_params["http_client"] = self._get_sync_http_client()
if max_retries is not None:
azure_client_params["max_retries"] = max_retries
azure_client_params["max_retries"] = max_retries
if timeout is not None:
azure_client_params["timeout"] = timeout

View file

@ -1,24 +1,102 @@
import re
from collections.abc import Callable, Collection, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Optional
import httpx
from httpx import Response
from pydantic import BaseModel, ValidationError
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import (
BasePassthroughConfig,
RelayShape,
logged_relay_shape,
replace_path_segment,
strip_leading_model_segment,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse, ResponsesTerminalEvent
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CallTypes, EmbeddingResponse, ImageResponse
if TYPE_CHECKING:
from httpx import URL
from litellm.types.utils import CostResponseTypes
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
class RelayedChatRequest(BaseModel):
messages: Sequence[Mapping[str, object]] | None = None
class RelayedCallDetails(BaseModel):
request_data: RelayedChatRequest | None = None
def _relayed_messages(litellm_logging_obj: Logging) -> Sequence[Mapping[str, object]] | None:
try:
details: Final = RelayedCallDetails.model_validate(litellm_logging_obj.model_call_details)
except ValidationError:
return None
return details.request_data.messages if details.request_data else None
RESPONSES_RELAY_SHAPE: Final = RelayShape("/responses", CallTypes.aresponses, ResponsesAPIResponse.model_validate)
OPENAI_RELAY_SHAPES: Final = (
RelayShape("/embeddings", CallTypes.aembedding, EmbeddingResponse.model_validate),
RESPONSES_RELAY_SHAPE,
RelayShape("/images/generations", CallTypes.aimage_generation, ImageResponse.model_validate),
)
def logged_responses_stream(all_chunks: Sequence[str], logging_obj: Logging) -> ResponsesTerminalEvent | None:
"""A streaming logging object assembles the logged response from the terminal event, not from its body."""
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
terminal_event: Final = OpenAIResponsesAPIConfig.parse_terminal_event_from_stream_chunks(all_chunks=all_chunks)
if terminal_event is None:
return None
logging_obj.call_type = (
RESPONSES_RELAY_SHAPE.call_type.value
) # rebind-ok: routes cost calculation to the relayed shape's pricing path
return terminal_event
AZURE_DEPLOYMENT_SEGMENT: Final = re.compile(r"(?<![^/])openai/deployments/([^/]+)")
def azure_router_model_in_endpoint(endpoint: str, router_models: Collection[str]) -> str | None:
parts: Final = endpoint.split("/")
if len(parts) < 2:
return None
return next((part for part in parts if part in router_models), None)
def foreign_azure_deployment(
endpoint: str, model_group: str, served_models: Callable[[], Collection[str]]
) -> str | None:
match: Final = AZURE_DEPLOYMENT_SEGMENT.search(endpoint)
if match is None:
return None
deployment: Final = match.group(1)
if deployment == model_group:
return None
served: Final = frozenset(name.casefold() for name in served_models())
return None if deployment.casefold() in served else deployment
def without_api_version(api_base: str) -> str:
url: Final = httpx.URL(api_base)
kept_params: Final = tuple((key, value) for key, value in url.params.multi_items() if key != "api-version")
return str(url.copy_with(params=httpx.QueryParams(kept_params)))
class AzurePassthroughConfig(BasePassthroughConfig):
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
return "stream" in request_data
return bool(request_data.get("stream"))
def get_complete_url(
self,
@ -36,14 +114,17 @@ class AzurePassthroughConfig(BasePassthroughConfig):
litellm_metadata: Final = litellm_params.get("litellm_metadata") or {}
model_group: Final = litellm_metadata.get("model_group")
if model_group and model_group in endpoint:
endpoint = endpoint.replace(model_group, model)
routed_endpoint: Final = replace_path_segment(endpoint, model_group, model) if model_group else endpoint
native_endpoint: Final = strip_leading_model_segment(routed_endpoint, (model,))
caller_api_version: Final = request_query_params.get("api-version") if request_query_params else None
relay_base: Final = without_api_version(base_target_url) if caller_api_version else base_target_url
complete_url: Final = BaseAzureLLM._get_base_azure_url(
api_base=base_target_url,
litellm_params=litellm_params,
route=endpoint,
default_api_version=litellm_params.get("api_version"),
api_base=relay_base,
litellm_params=MappingProxyType(
{**litellm_params, "api_version": caller_api_version or litellm_params.get("api_version")}
),
route=native_endpoint,
)
return (
httpx.URL(complete_url),
@ -92,13 +173,13 @@ class AzurePassthroughConfig(BasePassthroughConfig):
request_data: dict,
logging_obj: Logging,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> Optional["LoggedRelayResponse"]:
from litellm import encoding
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.types.utils import ModelResponse
if "chat/completions" not in endpoint:
return None
return logged_relay_shape(OPENAI_RELAY_SHAPES, httpx_response, logging_obj, endpoint)
openai_chat_config: Final = OpenAIGPTConfig()
@ -116,3 +197,27 @@ class AzurePassthroughConfig(BasePassthroughConfig):
)
return litellm_model_response
def handle_logging_collected_chunks(
self,
all_chunks: Sequence[str],
litellm_logging_obj: Logging,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["LoggedRelayResponse"]:
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
OpenAIPassthroughLoggingHandler,
)
if f"/{endpoint.strip('/')}".endswith(RESPONSES_RELAY_SHAPE.path_suffix):
return logged_responses_stream(all_chunks, litellm_logging_obj)
if "chat/completions" not in endpoint:
return None
return OpenAIPassthroughLoggingHandler()._build_complete_streaming_response( # pyright: ignore[reportPrivateUsage] # the only OpenAI SSE-to-ModelResponse assembler; reimplementing it would fork the parser
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
messages=_relayed_messages(litellm_logging_obj),
)

View file

@ -2,7 +2,6 @@ import copy
import enum
import re
from typing import TYPE_CHECKING, Final, cast
from urllib.parse import urlparse
import httpx
from httpx import Response
@ -15,7 +14,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
)
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
from litellm.llms.azure_ai.common_utils import (
api_key_header_for_base,
is_foundry_model_inference_base,
)
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
@ -146,11 +148,7 @@ class AzureAIStudioConfig(OpenAIConfig):
"""
Returns True if the request should use `api-key` header for authentication.
"""
parsed_url: Final = urlparse(api_base)
host: Final = parsed_url.hostname
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
return True
return False
return api_key_header_for_base(api_base) == "api-key"
def get_complete_url(
self,

View file

@ -19,6 +19,13 @@ def is_foundry_model_inference_base(api_base: str) -> bool:
return "/openai/deployments" not in parsed.path
def api_key_header_for_base(api_base: str | None) -> AzureAIApiKeyHeader:
host: Final = urlparse(api_base).hostname if api_base else None
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
return "api-key"
return "Authorization"
def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None) -> str | None:
"""
Resolve an Entra ID / OAuth access token for an Azure AI Foundry deployment.

View file

@ -23,7 +23,7 @@ def get_azure_ai_image_edit_config(model: str) -> BaseImageEditConfig:
"""
Get the appropriate image edit config for an Azure AI model.
- MAI models use /mai/v1/images/edits with multipart form data and size
- MAI models use /mai/v1/images/edits with multipart form data
- FLUX 2 models use JSON with base64 image
- FLUX 1 models use multipart/form-data
"""

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Any, Final
import httpx
from httpx._types import RequestFiles
@ -13,7 +13,6 @@ from litellm.llms.azure_ai.image_generation.mai_transformation import (
from litellm.llms.openai.common_utils import OpenAIError
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.llms.openai import FileTypes
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ImageResponse
@ -26,65 +25,8 @@ if TYPE_CHECKING:
class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig):
"""Azure AI Foundry MAI image editing (e.g. MAI-Image-2.5)."""
DEFAULT_SIZE = "1024x1024"
def get_supported_openai_params(self, model: str) -> list:
return ["prompt", "image", "model", "n", "size"]
def map_openai_params(
self,
image_edit_optional_params: ImageEditOptionalRequestParams,
model: str,
drop_params: bool,
) -> dict:
optional_params: Final[dict[str, Any]] = {}
supported_params: Final = self.get_supported_openai_params(model)
for key, value in dict(image_edit_optional_params).items():
if value is None or key in optional_params:
continue
if key in supported_params:
if key == "size" and value:
size_param = cast(str, value)
self._validate_size_param(size_param)
optional_params[key] = size_param
else:
optional_params[key] = value
elif not drop_params:
raise ValueError(
f"Parameter {key} is not supported for model {model}. "
f"Supported parameters are {supported_params}. "
f"Set drop_params=True to drop unsupported parameters."
)
if "size" not in optional_params:
optional_params["size"] = self.DEFAULT_SIZE
return optional_params
def _validate_size_param(self, size: str) -> None:
known_sizes: Final = {
"1024x1024",
"1792x1024",
"1024x1792",
"512x512",
"256x256",
}
if size in known_sizes:
return
if "x" in size:
try:
tuple(map(int, size.lower().split("x", 1)))
return
except ValueError:
raise ValueError(f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024').")
raise ValueError(
f"Unsupported size value: '{size}'. Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string."
)
return ["prompt", "image", "model", "n"]
def validate_environment(
self,

View file

@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm.exceptions import UnsupportedParamsError
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
@ -21,6 +22,10 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
DEFAULT_WIDTH = 1024
DEFAULT_HEIGHT = 1024
MAX_IMAGES_PER_REQUEST: Final = 1
MIN_DIMENSION_PX: Final = 768
MAX_TOTAL_PX: Final = 1_056_768
@staticmethod
def get_mai_image_generation_url(
api_base: str | None,
@ -145,16 +150,27 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
if k in supported_params:
if k == "size" and v:
self._map_size_param(v, optional_params)
self._map_size_param(v, optional_params, model)
elif k == "n" and v is not None and self._image_count(v, model) != self.MAX_IMAGES_PER_REQUEST:
if not drop_params:
raise self._unsupported(
model,
f"n={v} is not supported for model {model}. The Azure AI MAI image "
f"endpoint returns exactly {self.MAX_IMAGES_PER_REQUEST} image per "
"request and ignores any count, so a larger value would silently "
"return fewer images than requested. Send one request per image, or "
"set drop_params=True to drop n.",
)
else:
optional_params[k] = v
elif k in ("width", "height"):
optional_params[k] = v
elif not drop_params:
raise ValueError(
raise self._unsupported(
model,
f"Parameter {k} is not supported for model {model}. "
f"Supported parameters are {supported_params} and width/height. "
f"Set drop_params=True to drop unsupported parameters."
f"Set drop_params=True to drop unsupported parameters.",
)
if "width" not in optional_params:
@ -165,7 +181,19 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
optional_params.pop("size", None)
return optional_params
def _map_size_param(self, size: str, optional_params: dict) -> None:
@staticmethod
def _unsupported(model: str, message: str) -> UnsupportedParamsError:
return UnsupportedParamsError(message=message, llm_provider="azure_ai", model=model)
def _image_count(self, n: object, model: str) -> int:
if isinstance(n, int):
return n
try:
return int(str(n))
except ValueError:
raise self._unsupported(model, f"n={n!r} is not a whole number of images for model {model}.")
def _map_size_param(self, size: str, optional_params: dict, model: str) -> None:
size_mapping: Final = {
"1024x1024": (1024, 1024),
"1792x1024": (1792, 1024),
@ -176,19 +204,36 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
if size in size_mapping:
width, height = size_mapping[size]
optional_params["width"] = width
optional_params["height"] = height
elif "x" in size:
try:
width, height = map(int, size.lower().split("x"))
optional_params["width"] = width
optional_params["height"] = height
except ValueError:
raise ValueError(f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024').")
raise self._unsupported(
model, f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')."
)
else:
raise ValueError(
raise self._unsupported(
model,
f"Unsupported size value: '{size}'. "
f"Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string."
f"Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string.",
)
self._validate_dimensions(model=model, size=size, width=width, height=height)
optional_params["width"] = width
optional_params["height"] = height
def _validate_dimensions(self, model: str, size: str, width: int, height: int) -> None:
if width < self.MIN_DIMENSION_PX or height < self.MIN_DIMENSION_PX:
raise self._unsupported(
model,
f"Unsupported size value: '{size}'. Azure AI MAI image models require width and "
f"height of at least {self.MIN_DIMENSION_PX} pixels.",
)
if width * height > self.MAX_TOTAL_PX:
raise self._unsupported(
model,
f"Unsupported size value: '{size}'. Azure AI MAI image models accept at most "
f"{self.MAX_TOTAL_PX} total pixels ({width}x{height} is {width * height}).",
)
def transform_image_generation_response(

View file

@ -0,0 +1,232 @@
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
import httpx
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.llms.azure_ai.common_utils import (
AzureFoundryModelInfo,
api_key_header_for_base,
get_azure_ai_auth_headers,
)
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
from litellm.llms.base_llm.passthrough.transformation import (
BasePassthroughConfig,
RelayShape,
logged_relay_shape,
strip_leading_model_segment,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.rerank import RerankResponse
from litellm.types.utils import CallTypes, ImageResponse, StandardPassThroughResponseObject
if TYPE_CHECKING:
from httpx import URL, Response
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
EMPTY_QUERY: Final[Mapping[str, object]] = MappingProxyType({})
class PassthroughMetadata(BaseModel):
model_config = ConfigDict(extra="ignore")
model_group: str = ""
def model_group_from(litellm_params: Mapping[str, object]) -> str:
try:
return PassthroughMetadata.model_validate(litellm_params.get("litellm_metadata")).model_group
except ValidationError:
return ""
def api_version_from(litellm_params: Mapping[str, object]) -> str | None:
try:
return TypeAdapter(str | None).validate_python(litellm_params.get("api_version"))
except ValidationError:
return None
def foundry_root(api_base: str) -> str:
url: Final = httpx.URL(api_base)
segments: Final = tuple(segment for segment in url.path.split("/") if segment)
root_segments: Final = segments[: segments.index("models")] if "models" in segments else segments
return str(url.copy_with(path="/" + "/".join(root_segments), query=None)).rstrip("/")
def is_repeated_native_prefix(native_segments: tuple[str, ...], overlap: int) -> bool:
return overlap == len(native_segments) or native_segments[0] == "openai"
def without_repeated_native_prefix(root: str, native_endpoint: str) -> str:
url: Final = httpx.URL(root)
root_segments: Final = tuple(segment for segment in url.path.split("/") if segment)
native_segments: Final = tuple(segment.casefold() for segment in native_endpoint.split("/") if segment)
overlap: Final = next(
(
length
for length in range(min(len(root_segments), len(native_segments)), 0, -1)
if tuple(segment.casefold() for segment in root_segments[-length:]) == native_segments[:length]
and is_repeated_native_prefix(native_segments, length)
),
0,
)
kept_segments: Final = root_segments[: len(root_segments) - overlap]
return str(url.copy_with(path="/" + "/".join(kept_segments), query=None)).rstrip("/")
def relay_query_params(
request_query_params: Mapping[str, object] | None,
deployment_api_version: str | None,
api_base: str,
) -> Mapping[str, object] | None:
if request_query_params and "api-version" in request_query_params:
return request_query_params
api_version: Final = deployment_api_version or httpx.URL(api_base).params.get("api-version")
if api_version is None:
return request_query_params
return MappingProxyType({**(request_query_params or EMPTY_QUERY), "api-version": api_version})
def relayed_body(httpx_response: Response) -> str | dict:
try:
body: Final[object] = httpx_response.json()
except ValueError:
return httpx_response.text
return body if isinstance(body, dict) else httpx_response.text
FOUNDRY_RELAY_SHAPES: Final = (
RelayShape("/rerank", CallTypes.arerank, RerankResponse.model_validate),
RelayShape("/providers/blackforestlabs/v1/flux-2-pro", CallTypes.aimage_generation, ImageResponse.model_validate),
)
class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig):
def __init__(self, ocr_config_for: Callable[[str], BaseOCRConfig | None] = get_azure_ai_ocr_config) -> None:
super().__init__()
self.ocr_config_for: Final = ocr_config_for
def is_streaming_request(self, endpoint: str, request_data: Mapping[str, object]) -> bool:
return bool(request_data.get("stream"))
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
endpoint: str,
request_query_params: Mapping[str, object] | None,
litellm_params: Mapping[str, object],
) -> tuple[URL, str]:
base_target_url: Final = self.get_api_base(api_base)
if base_target_url is None:
raise ValueError("Azure AI api base not found: set `api_base` on the deployment or AZURE_AI_API_BASE")
native_endpoint: Final = strip_leading_model_segment(endpoint, (model, model_group_from(litellm_params)))
root: Final = without_repeated_native_prefix(foundry_root(base_target_url), native_endpoint)
query_params: Final = relay_query_params(
request_query_params, api_version_from(litellm_params), base_target_url
)
return (self.format_url(native_endpoint, root, query_params), root)
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict[str, str]: # mutable-ok: base class contract returns dict for httpx
auth_headers: Final = get_azure_ai_auth_headers(
api_key=api_key,
litellm_params=litellm_params,
api_key_header=api_key_header_for_base(api_base),
)
return {**headers, **auth_headers} # mutable-ok: base class contract returns dict for httpx
def logging_non_streaming_response(
self,
model: str,
custom_llm_provider: str,
httpx_response: Response,
request_data: Mapping[str, object],
logging_obj: Logging,
endpoint: str,
) -> LoggedRelayResponse | OCRResponse | StandardPassThroughResponseObject | None:
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
chat_result: Final = AzurePassthroughConfig().logging_non_streaming_response( # pyright: ignore[reportUnknownMemberType] # the Azure config still types request_data as a bare dict
model=model,
custom_llm_provider=custom_llm_provider,
httpx_response=httpx_response,
request_data=dict(request_data), # mutable-ok: AzurePassthroughConfig wants a dict
logging_obj=logging_obj,
endpoint=endpoint,
)
if chat_result is not None:
return chat_result
ocr_result: Final = self.logged_ocr_response(model, httpx_response, logging_obj, endpoint)
if ocr_result is not None:
return ocr_result
foundry_result: Final = logged_relay_shape(FOUNDRY_RELAY_SHAPES, httpx_response, logging_obj, endpoint)
if foundry_result is not None:
return foundry_result
return StandardPassThroughResponseObject(response=relayed_body(httpx_response))
def logged_ocr_response(
self, model: str, httpx_response: Response, logging_obj: Logging, endpoint: str
) -> OCRResponse | None:
ocr_config: Final = self.ocr_config_for(model)
if ocr_config is None or httpx_response.status_code != 200:
return None
relayed_url: Final = httpx_response.request.url
relayed_origin: Final = str(relayed_url.copy_with(path="/", query=None, fragment=None)).rstrip("/")
ocr_url: Final = httpx.URL(
ocr_config.get_complete_url(
api_base=relayed_origin,
model=model,
optional_params={}, # mutable-ok: BaseOCRConfig wants a dict
)
)
known_prefixes: Final = (model, model_group_from(logging_obj.litellm_params))
native_endpoint: Final = strip_leading_model_segment(endpoint, known_prefixes)
if f"/{native_endpoint.strip('/')}" != ocr_url.path:
return None
try:
ocr_response: Final = ocr_config.transform_ocr_response(
model=model, raw_response=httpx_response, logging_obj=logging_obj
)
except (ValueError, AttributeError) as error:
verbose_logger.warning("azure_ai passthrough: OCR body from %s is not costable: %s", ocr_url, error)
return None
logging_obj.call_type = CallTypes.aocr.value # rebind-ok: routes cost calculation to the per-page OCR path
return ocr_response
def handle_logging_collected_chunks(
self,
all_chunks: Sequence[str],
litellm_logging_obj: Logging,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> LoggedRelayResponse | None:
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
return AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
custom_llm_provider=custom_llm_provider,
endpoint=endpoint,
)

View file

@ -52,13 +52,28 @@ class StreamingScanKey:
class BaseTranslation(ABC):
delivers_ended_stream_text_rewrites: ClassVar[bool] = False
delivers_ended_stream_rewrites: ClassVar[bool] = False
"""Whether ``process_output_streaming_response`` accepts
``deliver_ended_stream_rewrites=True`` and, on an ended (fully buffered)
stream, writes guardrail text rewrites back across ``responses_so_far`` so
a buffered pipeline can release rewritten chunks. Tool-call rewrites, and
text rewrites on every other translation, are undeliverable: the pipeline
executor discards them and releases the original chunks."""
stream, writes guardrail text and tool-call rewrites back across
``responses_so_far`` so a buffered pipeline can release rewritten chunks,
raising ``UndeliverableStreamRewrite`` for a shape it cannot place. Rewrites
on every other translation are undeliverable: the pipeline executor
discards them and releases the original chunks."""
assembles_streamed_response: ClassVar[bool] = False
"""Whether ``process_output_streaming_response`` stores the assembled response of an
ended stream under ``request_data["response"]`` before scanning it, the way the chat,
Responses, and Messages translations do. A streaming pipeline runs a guardrail that only
has the legacy post-call hook against that response, so on a translation without it such
a guardrail keeps running on its own."""
def post_call_hook_response(self, response: object) -> object:
"""The ``response`` this endpoint's non-streaming post-call hooks receive, derived from
the object the translation stores under ``request_data["response"]`` while scanning an
ended stream. Chat and Responses scan that shape already; a translation that scans a
different one (Messages scans an OpenAI-shaped ModelResponse) overrides this."""
return response
@staticmethod
def transform_user_api_key_dict_to_metadata(
@ -175,9 +190,9 @@ class BaseTranslation(ABC):
transformations (see ``StreamTransformSink``); base handlers ignore it.
``deliver_ended_stream_rewrites`` is passed True only when the caller
holds the whole buffered stream and the subclass declares
``delivers_ended_stream_text_rewrites``: the handler then writes
guardrail text rewrites back across ``responses_so_far`` instead of
discarding them.
``delivers_ended_stream_rewrites``: the handler then writes
guardrail text and tool-call rewrites back across ``responses_so_far``
instead of discarding them.
"""
return responses_so_far

View file

@ -1,5 +1,14 @@
from __future__ import annotations
import re
from abc import abstractmethod
from typing import TYPE_CHECKING, Final, Optional, Union
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, TypeAlias
from pydantic import TypeAdapter, ValidationError
from litellm.types.utils import CallTypes
from ..base_utils import BaseLLMModelInfo
@ -7,9 +16,68 @@ if TYPE_CHECKING:
from httpx import URL, Headers, Response
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import CostResponseTypes
from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesTerminalEvent
from litellm.types.rerank import RerankResponse
from litellm.types.utils import CostResponseTypes, StandardPassThroughResponseObject
from ..chat.transformation import BaseLLMException
from ..ocr.transformation import OCRResponse
LoggedRelayResponse: TypeAlias = CostResponseTypes | RerankResponse | ResponsesAPIResponse | ResponsesTerminalEvent
RELAYED_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
def strip_leading_model_segment(endpoint: str, model_names: tuple[str, ...]) -> str:
path: Final = endpoint.lstrip("/")
for model_name in model_names:
if not model_name:
continue
if path == model_name:
return ""
if path.startswith(f"{model_name}/"):
return path[len(model_name) + 1 :]
return path
def replace_path_segment(endpoint: str, segment: str, replacement: str) -> str:
bounded_segment: Final = re.compile(rf"(?<![^/]){re.escape(segment)}(?![^/:])")
return bounded_segment.sub(lambda _: replacement, endpoint)
def relayed_json_object(httpx_response: Response) -> Mapping[str, object] | None:
if httpx_response.status_code != 200:
return None
try:
return RELAYED_JSON_OBJECT.validate_python(httpx_response.json())
except (ValueError, ValidationError):
return None
@dataclass(frozen=True, slots=True)
class RelayShape:
path_suffix: str
call_type: CallTypes
parse: Callable[[Mapping[str, object]], LoggedRelayResponse]
def logged_relay_shape(
shapes: Sequence[RelayShape], httpx_response: Response, logging_obj: LiteLLMLoggingObj, endpoint: str
) -> LoggedRelayResponse | None:
relayed_path: Final = f"/{endpoint.strip('/')}"
shape: Final = next((candidate for candidate in shapes if relayed_path.endswith(candidate.path_suffix)), None)
body: Final = relayed_json_object(httpx_response) if shape else None
if shape is None or body is None:
return None
try:
parsed: Final = shape.parse(body)
except ValidationError:
return None
logging_obj.call_type = (
shape.call_type.value
) # rebind-ok: routes cost calculation to the relayed shape's pricing path
return parsed
class BasePassthroughConfig(BaseLLMModelInfo):
@ -23,8 +91,8 @@ class BasePassthroughConfig(BaseLLMModelInfo):
self,
endpoint: str,
base_target_url: str,
request_query_params: dict | None,
) -> "URL":
request_query_params: Mapping[str, object] | None,
) -> URL:
"""
Helper function to add query params to the url
Args:
@ -58,7 +126,7 @@ class BasePassthroughConfig(BaseLLMModelInfo):
endpoint: str,
request_query_params: dict | None,
litellm_params: dict,
) -> tuple["URL", str]:
) -> tuple[URL, str]:
"""
Get the complete url for the request
Returns:
@ -88,9 +156,7 @@ class BasePassthroughConfig(BaseLLMModelInfo):
"""
return headers, None
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, "Headers"]
) -> "BaseLLMException":
def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException:
from litellm.llms.base_llm.chat.transformation import BaseLLMException
return BaseLLMException(status_code=status_code, message=error_message, headers=headers)
@ -99,21 +165,21 @@ class BasePassthroughConfig(BaseLLMModelInfo):
self,
model: str,
custom_llm_provider: str,
httpx_response: "Response",
httpx_response: Response,
request_data: dict,
logging_obj: "LiteLLMLoggingObj",
logging_obj: LiteLLMLoggingObj,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> LoggedRelayResponse | OCRResponse | StandardPassThroughResponseObject | None:
pass
def handle_logging_collected_chunks(
self,
all_chunks: list[str],
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> LoggedRelayResponse | None:
return None
def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]:

View file

@ -1,13 +1,17 @@
import asyncio
import base64
import contextvars
import hashlib
import json
import os
import re
import urllib.parse
from collections.abc import Callable, Mapping
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
from threading import Lock
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast, get_args, overload
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
from pydantic import BaseModel, ValidationError
@ -16,6 +20,7 @@ from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
AWS_SIGNING_MAX_THREADS,
BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
BEDROCK_IAM_CACHE_FETCH_LOCK_STRIPES,
BEDROCK_IAM_CACHE_MAX_ENTRIES,
@ -80,7 +85,11 @@ class AwsAuthError(Exception):
super().__init__(self.message) # Call the base class constructor with the parameters it needs
class BaseAWSLLM:
class SignsRequestsWithAWS:
pass
class BaseAWSLLM(SignsRequestsWithAWS):
# Process-wide IAM credential cache (shared across instances — Bedrock passthrough is per-request).
# Storage is in-process memory only: no Redis backend unless attached elsewhere. Entry TTL: static
# access-key + secret + region use ``_get_default_ttl_for_boto3_credentials`` (~59 minutes); ambient
@ -1668,3 +1677,52 @@ class BaseAWSLLM:
request_headers_dict["Authorization"] = incoming_authorization
return request_headers_dict, request.body
def sign_aws_json_post(
get_credentials: Callable[[], Credentials],
service_name: str,
aws_region_name: str | None,
url: str,
body: str,
headers: Mapping[str, str],
) -> AWSPreparedRequest:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError(f"Missing boto3 to call {service_name}. Run 'pip install boto3'.")
aws_request: Final = AWSRequest(method="POST", url=url, data=body, headers=headers)
SigV4Auth(get_credentials(), service_name, aws_region_name).add_auth(aws_request)
return aws_request.prepare()
_SignParams = ParamSpec("_SignParams")
_SignedRequest = TypeVar("_SignedRequest")
AWS_SIGNING_EXECUTOR: Final = ThreadPoolExecutor(max_workers=AWS_SIGNING_MAX_THREADS, thread_name_prefix="aws-signing")
async def run_aws_signing(
sign: Callable[_SignParams, _SignedRequest],
/,
*args: _SignParams.args,
**kwargs: _SignParams.kwargs, # kwargs-ok: ParamSpec forwarding keeps the wrapped signing signature
) -> _SignedRequest:
context: Final = contextvars.copy_context()
return await asyncio.get_running_loop().run_in_executor(
AWS_SIGNING_EXECUTOR, partial(context.run, sign, *args, **kwargs)
)
async def sign_request_off_loop_if_aws(
provider_config: object,
sign_request: Callable[_SignParams, _SignedRequest],
/,
*args: _SignParams.args,
**kwargs: _SignParams.kwargs, # kwargs-ok: ParamSpec forwarding keeps the wrapped sign_request signature
) -> _SignedRequest:
if isinstance(provider_config, SignsRequestsWithAWS):
return await run_aws_signing(sign_request, *args, **kwargs)
return sign_request(*args, **kwargs)

View file

@ -21,7 +21,7 @@ from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing
from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
@ -136,7 +136,8 @@ class BedrockConverseLLM(BaseAWSLLM):
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
prepped: Final = await run_aws_signing(
self.get_request_headers,
credentials=credentials,
aws_region_name=litellm_params.get("aws_region_name") or "us-west-2",
extra_headers=headers,
@ -206,7 +207,8 @@ class BedrockConverseLLM(BaseAWSLLM):
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
prepped: Final = await run_aws_signing(
self.get_request_headers,
credentials=credentials,
aws_region_name=litellm_params.get("aws_region_name") or "us-west-2",
extra_headers=headers,

View file

@ -4,6 +4,7 @@ Translating between OpenAI's `/chat/completion` format and Amazon's `/converse`
import copy
import json
import re
import time
import types
from collections.abc import Mapping
@ -293,6 +294,10 @@ class AmazonConverseConfig(BaseConfig):
llm_provider="bedrock",
)
@staticmethod
def _is_openai_gpt_reasoning_model(model: str) -> bool:
return re.search(r"openai\.gpt-\d", model) is not None
def _is_nova_2_model(self, model: str) -> bool:
"""
Check if the model is a Nova 2 model that supports reasoningConfig.
@ -422,15 +427,15 @@ class AmazonConverseConfig(BaseConfig):
"""
Handle the reasoning_effort parameter based on the model type.
- GPT-OSS models: passed through unchanged via additionalModelRequestFields.
- OpenAI GPT-5.x models: mapped to ``reasoning.effort`` via additionalModelRequestFields.
- GPT-OSS and DeepSeek V3 models: passed through unchanged via additionalModelRequestFields.
- OpenAI GPT-5.x and GPT-6 models: mapped to ``reasoning.effort`` via additionalModelRequestFields.
- Nova 2 models: transformed to reasoningConfig.
- Anthropic models: mapped to ``thinking`` (and ``output_config.effort`` on
adaptive Claude 4.6 / 4.7).
"""
if "gpt-oss" in model:
if "gpt-oss" in model or "deepseek" in model:
optional_params["reasoning_effort"] = reasoning_effort
elif "openai.gpt-5" in model:
elif self._is_openai_gpt_reasoning_model(model):
reasoning: Final[BedrockConverseGptReasoningEffortBlock] = {"effort": reasoning_effort}
optional_params["reasoning"] = reasoning
elif self._is_nova_2_model(model):
@ -509,6 +514,36 @@ class AmazonConverseConfig(BaseConfig):
)
thinking["budget_tokens"] = BEDROCK_MIN_THINKING_BUDGET_TOKENS
def _is_deepseek_model(self, model: str, base_model: str) -> bool:
return "deepseek" in model or "deepseek" in base_model
def _is_deepseek_r1_model(self, model: str, base_model: str) -> bool:
return "deepseek.r1" in model or "deepseek.r1" in base_model
def _model_accepts_anthropic_thinking_param(self, model: str, base_model: str) -> bool:
"""Whether the model accepts the Anthropic-shaped ``thinking`` request field.
Only Claude reasoning models accept it. DeepSeek advertises ``supports_reasoning`` but reasons
natively: R1 returns a 400 when the field is sent and V3 silently ignores it.
"""
if self._is_deepseek_model(model=model, base_model=base_model):
return False
return (
"claude-3-7" in model
or "claude-sonnet-4" in model
or "claude-opus-4" in model
or supports_reasoning(model=model, custom_llm_provider=self.custom_llm_provider)
or supports_reasoning(model=base_model, custom_llm_provider=self.custom_llm_provider)
)
def _model_rejects_reasoning_effort_param(self, model: str, base_model: str) -> bool:
"""Whether the model returns a 400 for every ``reasoning_effort`` shape on Converse.
DeepSeek R1 always reasons and rejects any reasoning request field. DeepSeek V3 accepts a raw
``reasoning_effort`` like gpt-oss does, and every other model maps it to a shape it accepts.
"""
return self._is_deepseek_r1_model(model=model, base_model=base_model)
def get_supported_openai_params(self, model: str) -> list[str]:
from litellm.utils import supports_function_calling
@ -564,23 +599,20 @@ class AmazonConverseConfig(BaseConfig):
# only anthropic and mistral support tool choice config. otherwise (E.g. cohere) will fail the call - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html
supported_params.append("tool_choice")
if "gpt-oss" in model or "openai.gpt-5" in model or "openai.gpt-5" in base_model:
if (
"gpt-oss" in model
or self._is_openai_gpt_reasoning_model(model)
or self._is_openai_gpt_reasoning_model(base_model)
):
supported_params.append("reasoning_effort")
elif self._is_deepseek_model(model=model, base_model=base_model):
if not self._is_deepseek_r1_model(model=model, base_model=base_model):
supported_params.append("reasoning_effort")
elif self._is_nova_2_model(model):
# Nova 2 models support reasoning_effort (transformed to reasoningConfig)
# These models use a different reasoning structure than Anthropic's thinking parameter
supported_params.append("reasoning_effort")
elif (
"claude-3-7" in model
or "claude-sonnet-4" in model
or "claude-opus-4" in model
or "deepseek.r1" in model
or supports_reasoning(
model=model,
custom_llm_provider=self.custom_llm_provider,
)
or supports_reasoning(model=base_model, custom_llm_provider=self.custom_llm_provider)
):
elif self._model_accepts_anthropic_thinking_param(model=model, base_model=base_model):
supported_params.append("thinking")
supported_params.append("reasoning_effort")
supported_params.append("output_config")
@ -872,6 +904,11 @@ class AmazonConverseConfig(BaseConfig):
drop_params: bool,
) -> dict:
is_thinking_enabled: Final = self.is_thinking_enabled(non_default_params)
base_model: Final = BedrockModelInfo.get_base_model(model)
drop_thinking_param: Final = self._is_deepseek_model(model=model, base_model=base_model)
drop_reasoning_effort_param: Final = self._model_rejects_reasoning_effort_param(
model=model, base_model=base_model
)
for param, value in non_default_params.items():
if param == "response_format" and isinstance(value, dict):
@ -920,7 +957,12 @@ class AmazonConverseConfig(BaseConfig):
optional_params["_parallel_tool_use_config"] = {
"tool_choice": {"type": "auto", "disable_parallel_tool_use": not value}
}
if param == "thinking" and "openai.gpt-5" not in model:
if param == "thinking" and drop_thinking_param:
verbose_logger.debug(
"Dropping unsupported `thinking` param for Bedrock model=%s; it reasons natively.",
model,
)
elif param == "thinking" and not self._is_openai_gpt_reasoning_model(model):
if (
isinstance(value, dict)
and value.get("type") == "adaptive"
@ -946,6 +988,11 @@ class AmazonConverseConfig(BaseConfig):
AnthropicModelInfo.translate_legacy_thinking_for_adaptive_model(
model=model, optional_params=optional_params, custom_llm_provider="bedrock"
)
elif param == "reasoning_effort" and isinstance(value, str) and drop_reasoning_effort_param:
verbose_logger.debug(
"Dropping unsupported `reasoning_effort` param for Bedrock model=%s; it always reasons and rejects it.",
model,
)
elif param == "reasoning_effort" and isinstance(value, str):
self._handle_reasoning_effort_parameter(
model=model, reasoning_effort=value, optional_params=optional_params
@ -1805,6 +1852,7 @@ class AmazonConverseConfig(BaseConfig):
data=request_data,
messages=messages,
encoding=encoding,
json_mode=json_mode,
)
def _transform_reasoning_content(self, reasoning_content_blocks: list[BedrockConverseReasoningContentBlock]) -> str:
@ -2237,6 +2285,7 @@ class AmazonConverseConfig(BaseConfig):
data: dict | str,
messages: list,
encoding,
json_mode: bool | None = None,
) -> ModelResponse:
## LOGGING
if logging_obj is not None:
@ -2247,7 +2296,9 @@ class AmazonConverseConfig(BaseConfig):
additional_args={"complete_input_dict": data},
)
json_mode: Final[bool | None] = optional_params.get("json_mode", None)
resolved_json_mode: Final[bool | None] = (
json_mode if json_mode is not None else optional_params.get("json_mode", None)
)
## RESPONSE OBJECT
try:
completion_response: Final = ConverseResponseBlock(**response.json())
@ -2339,7 +2390,7 @@ class AmazonConverseConfig(BaseConfig):
chat_completion_message["thinking_blocks"] = self._transform_thinking_blocks(reasoningContentBlocks)
chat_completion_message["content"] = content_str
filtered_tools: Final = self._filter_json_mode_tools(
json_mode=json_mode,
json_mode=resolved_json_mode,
tools=tools,
chat_completion_message=chat_completion_message,
)
@ -2363,7 +2414,7 @@ class AmazonConverseConfig(BaseConfig):
# When json_mode filtered out all synthetic tool calls the response
# is plain content, not a pending tool invocation. Fix finish_reason
# so callers (e.g. OpenAI SDK) don't misinterpret it.
if json_mode and not filtered_tools and tools:
if resolved_json_mode and not filtered_tools and tools:
initial_finish_reason = "stop"
(

View file

@ -340,6 +340,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
json_mode=json_mode,
)
elif provider == "twelvelabs":
return litellm.AmazonTwelveLabsPegasusConfig().transform_response(

View file

@ -10,9 +10,10 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import run_aws_signing
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
class BedrockCountTokensHandler(BedrockCountTokensConfig):
@ -27,6 +28,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
request_data: dict[str, Any],
litellm_params: dict[str, Any],
resolved_model: str,
client: AsyncHTTPHandler | None = None,
) -> dict[str, Any]:
"""
Handle a CountTokens request using existing LiteLLM patterns.
@ -75,7 +77,8 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
# Extract api_key for bearer token auth if provided
api_key: Final = litellm_params.get("api_key", None)
headers: Final = {"Content-Type": "application/json"}
signed_headers, signed_body = self._sign_request(
signed_headers, signed_body = await run_aws_signing(
self._sign_request,
service_name="bedrock",
headers=headers,
optional_params=litellm_params,
@ -85,7 +88,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
api_key=api_key,
)
async_client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
async_client: Final = client or get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
response: Final = await async_client.post(
endpoint_url,

View file

@ -5,7 +5,7 @@ Handles embedding calls to Bedrock's `/invoke` endpoint
import copy
import json
import urllib.parse
from collections.abc import Callable
from collections.abc import Callable, Mapping
from typing import TYPE_CHECKING, Final, get_args, overload
import httpx
@ -26,7 +26,7 @@ from litellm.types.llms.bedrock import (
)
from litellm.types.utils import EmbeddingResponse, LlmProviders
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..base_aws_llm import AWSPreparedRequest, BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing
from ..common_utils import BedrockError
from .amazon_nova_transformation import AmazonNovaEmbeddingConfig
from .amazon_titan_g1_transformation import AmazonTitanG1Config
@ -41,6 +41,20 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
def _sign_get_request(
credentials: Credentials, url: str, headers: Mapping[str, str], aws_region_name: str
) -> AWSPreparedRequest:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
request: Final = AWSRequest(method="GET", url=url, data=None, headers=headers)
SigV4Auth(credentials, "bedrock", aws_region_name).add_auth(request)
return request.prepare()
class BedrockEmbedding(BaseAWSLLM):
@overload
def _load_credentials(
@ -342,7 +356,8 @@ class BedrockEmbedding(BaseAWSLLM):
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
prepped = await run_aws_signing(
self.get_request_headers,
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
@ -600,9 +615,6 @@ class BedrockEmbedding(BaseAWSLLM):
dict: Status response from AWS Bedrock
"""
# Get AWS credentials using the same method as other Bedrock methods
credentials, _ = self._load_credentials(kwargs)
# Get the runtime endpoint
endpoint_url, _ = self.get_runtime_endpoint(
api_base=None,
@ -619,27 +631,13 @@ class BedrockEmbedding(BaseAWSLLM):
# Prepare headers for GET request
headers: Final = {"Content-Type": "application/json"}
# Use AWSRequest directly for GET requests (get_request_headers hardcodes POST)
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
def sign_status_request() -> AWSPreparedRequest:
credentials, _ = self._load_credentials(kwargs)
return _sign_get_request(
credentials=credentials, url=status_url, headers=headers, aws_region_name=aws_region_name
)
# Create AWSRequest with GET method and encoded URL
request: Final = AWSRequest(
method="GET",
url=status_url,
data=None, # GET request, no body
headers=headers,
)
# Sign the request - SigV4Auth will create canonical string from request URL
sigv4: Final = SigV4Auth(credentials, "bedrock", aws_region_name)
sigv4.add_auth(request)
# Prepare the request
prepped: Final = request.prepare()
prepped: Final = await run_aws_signing(sign_status_request)
# LOGGING
if logging_obj is not None:

View file

@ -21,7 +21,7 @@ from litellm.litellm_core_utils.realtime_streaming import DefaultLoggedRealTimeE
from litellm.types.llms.openai import OpenAIRealtimeEvents
from litellm.types.realtime import RealtimeResponseTransformInput
from ..base_aws_llm import BaseAWSLLM
from ..base_aws_llm import BaseAWSLLM, run_aws_signing
from ..common_utils import BedrockError
from .transformation import BedrockRealtimeConfig
@ -149,7 +149,8 @@ class BedrockRealtime(BaseAWSLLM):
verbose_proxy_logger.debug("Bedrock Realtime: Connecting to %s with model %s", endpoint_uri, model)
credentials: Final = self.get_credentials(
credentials: Final = await run_aws_signing(
self.get_credentials,
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
@ -169,7 +170,7 @@ class BedrockRealtime(BaseAWSLLM):
"or configure credentials in the environment"
),
)
frozen_credentials: Final = credentials.get_frozen_credentials()
frozen_credentials: Final = await run_aws_signing(credentials.get_frozen_credentials)
# Initialize Bedrock client with aws_sdk_bedrock_runtime
config: Final = Config(

View file

@ -23,7 +23,7 @@ from botocore.exceptions import (
ProfileNotFound,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, SignsRequestsWithAWS
from litellm.secret_managers.main import get_secret_str
BEDROCK_MANTLE_DEFAULT_REGION: Final = "us-east-1"
@ -55,7 +55,7 @@ def resolve_mantle_region(params: Mapping[str, object]) -> str:
)
class BedrockMantleAuthMixin:
class BedrockMantleAuthMixin(SignsRequestsWithAWS):
_aws_signer: BaseAWSLLM
@staticmethod

View file

@ -77,6 +77,7 @@ from litellm.llms.base_llm.vector_store_files.transformation import (
BaseVectorStoreFilesConfig,
)
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS, run_aws_signing, sign_request_off_loop_if_aws
from litellm.llms.custom_httpx.container_handler import raise_for_error_status
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -637,7 +638,12 @@ class BaseLLMHTTPHandler:
headers=request_headers,
),
)
return await dispatch_async(*await asyncio.to_thread(sign_and_log, transformed))
signed_request: Final = await (
run_aws_signing(sign_and_log, transformed)
if isinstance(provider_config, SignsRequestsWithAWS)
else asyncio.to_thread(sign_and_log, transformed)
)
return await dispatch_async(*signed_request)
return transform_then_dispatch()
@ -1973,7 +1979,9 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
signed_headers, signed_json_body = await sign_request_off_loop_if_aws(
provider_config,
provider_config.sign_request,
headers=headers,
optional_params=optional_params,
request_data=data,
@ -2074,7 +2082,9 @@ class BaseLLMHTTPHandler:
max_attempts,
)
provider_config.transform_anthropic_messages_request_on_http_error(e=e, request_data=request_body)
headers, signed_json_body = provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
provider_config,
provider_config.sign_request,
headers=headers,
optional_params=optional_params_dict,
request_data=request_body,
@ -2234,7 +2244,9 @@ class BaseLLMHTTPHandler:
stream=stream,
)
headers, signed_json_body = anthropic_messages_provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
anthropic_messages_provider_config,
anthropic_messages_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params), # dynamic aws_* params are passed under litellm_params
request_data=request_body,
@ -2910,7 +2922,9 @@ class BaseLLMHTTPHandler:
fake_stream=fake_stream,
)
headers, signed_body = responses_api_provider_config.sign_request(
headers, signed_body = await sign_request_off_loop_if_aws(
responses_api_provider_config,
responses_api_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
@ -4618,7 +4632,9 @@ class BaseLLMHTTPHandler:
)
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
headers, signed_body = responses_api_provider_config.sign_request(
headers, signed_body = await sign_request_off_loop_if_aws(
responses_api_provider_config,
responses_api_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
@ -9845,7 +9861,9 @@ class BaseLLMHTTPHandler:
)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
headers, signed_json_body = vector_store_provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
vector_store_provider_config,
vector_store_provider_config.sign_request,
headers=headers,
optional_params=all_optional_params,
request_data=request_body,

View file

@ -15,6 +15,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 (
_extract_reasoning_content, # pyright: ignore[reportPrivateUsage] # same import as the OpenAI transformation
strip_litellm_internal_message_fields,
strip_name_from_message,
)
@ -23,7 +24,9 @@ from litellm.types.llms.anthropic import AllAnthropicToolsValues
from litellm.types.llms.databricks import (
AllDatabricksContentValues,
DatabricksChoice,
DatabricksDelta,
DatabricksFunction,
DatabricksMessage,
DatabricksResponse,
DatabricksTool,
)
@ -247,8 +250,10 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
litellm_params: dict,
stream: bool | None = None,
) -> str:
api_base = self._get_api_base(api_base)
complete_url: Final = f"{api_base}/chat/completions"
use_ai_gateway: Final = model.removeprefix("databricks/").count(".") >= 2
api_base = self._get_api_base(api_base, use_ai_gateway=use_ai_gateway)
url_base: Final = api_base.rstrip("/") if use_ai_gateway else api_base
complete_url: Final = f"{url_base}/chat/completions"
return complete_url
def get_supported_openai_params(self, model: str | None = None) -> list:
@ -534,6 +539,19 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
thinking_blocks.append(thinking_block)
return reasoning_content, thinking_blocks
@staticmethod
def extract_top_level_reasoning_content(delta: DatabricksDelta) -> str | None:
return delta.get("reasoning_content")
@staticmethod
def resolve_reasoning_and_content(
message: DatabricksMessage, block_reasoning_content: str | None
) -> tuple[str | None, str | None]:
content_str: Final = DatabricksConfig.extract_content_str(message["content"])
if block_reasoning_content is not None:
return block_reasoning_content, content_str
return _extract_reasoning_content({**message, "content": content_str})
@staticmethod
def extract_citations(
content: AllDatabricksContentValues | None,
@ -577,14 +595,13 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
finish_reason = "stop"
if translated_message is None:
## get the content str
content_str = DatabricksConfig.extract_content_str(choice["message"]["content"])
## get the reasoning content
(
reasoning_content,
block_reasoning_content,
thinking_blocks,
) = DatabricksConfig.extract_reasoning_content(choice["message"].get("content"))
reasoning_content, content_str = DatabricksConfig.resolve_reasoning_and_content(
choice["message"], block_reasoning_content
)
citations = DatabricksConfig.extract_citations(choice["message"].get("content"))
@ -738,12 +755,16 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator):
# extract the reasoning content
(
reasoning_content,
block_reasoning_content,
thinking_blocks,
) = DatabricksConfig.extract_reasoning_content(choice["delta"].get("content"))
choice["delta"]["content"] = content_str
choice["delta"]["reasoning_content"] = reasoning_content
choice["delta"]["reasoning_content"] = (
block_reasoning_content
if block_reasoning_content is not None
else DatabricksConfig.extract_top_level_reasoning_content(choice["delta"])
)
choice["delta"]["thinking_blocks"] = thinking_blocks
translated_choices.append(choice)
return ModelResponseStream(

View file

@ -177,19 +177,13 @@ class DatabricksBase:
# Default: just litellm
return f"litellm/{version}"
def _get_api_base(self, api_base: str | None) -> str:
"""
Get the Databricks API base URL.
If not provided, attempts to get it from the Databricks SDK.
"""
def _get_api_base(self, api_base: str | None, use_ai_gateway: bool = False) -> str:
if api_base is None:
try:
from databricks.sdk import WorkspaceClient
databricks_client: Final = WorkspaceClient()
api_base = f"{databricks_client.config.host}/serving-endpoints"
return api_base
except ImportError:
raise DatabricksException(
status_code=400,
@ -198,6 +192,18 @@ class DatabricksBase:
"or install the databricks-sdk Python library."
),
)
if not use_ai_gateway:
return api_base
normalized_api_base: Final = api_base.rstrip("/")
if normalized_api_base.endswith("/ai-gateway/mlflow/v1"):
return normalized_api_base
if normalized_api_base.endswith("/serving-endpoints"):
return f"{normalized_api_base.removesuffix('/serving-endpoints')}/ai-gateway/mlflow/v1"
api_base_parts: Final = urlsplit(normalized_api_base)
if api_base_parts.path in ("", "/"):
return f"{normalized_api_base}/ai-gateway/mlflow/v1"
return api_base
def _get_oauth_m2m_token(

View file

@ -49,6 +49,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
coerce_stream_holdback_value,
)
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Choices,
GenericGuardrailAPIInputs,
ModelResponse,
@ -78,7 +80,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
assembles_streamed_response = True
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
"""
@ -610,13 +613,14 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
deliver_ended_stream_rewrites: bool,
) -> None:
"""Ended-stream path: rebuild the full response, run the non-streaming
output guardrail against it, and (when opted in) write any text rewrite
back across the buffered chunks."""
output guardrail against it, and (when opted in) write any text or
tool-call rewrite back across the buffered chunks."""
model_response: Final = cast(
ModelResponse,
stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj),
)
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
await self.process_output_response(
response=model_response,
guardrail_to_apply=guardrail_to_apply,
@ -624,13 +628,21 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
if deliver_ended_stream_rewrites:
await self._write_ended_stream_text_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_texts=pre_guardrail_texts,
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
if not deliver_ended_stream_rewrites:
return
guardrail_name: Final = guardrail_to_apply.guardrail_name or "unknown"
await self._write_ended_stream_text_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_texts=pre_guardrail_texts,
guardrail_name=guardrail_name,
)
self._write_ended_stream_tool_call_rewrites(
responses_so_far=responses_so_far,
guardrailed_response=model_response,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
guardrail_name=guardrail_name,
)
def build_stream_error_items(
self,
@ -1043,6 +1055,71 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
task_mappings=[(target_choice_index, None) for _ in changed], # mutable-ok: callee takes lists
)
@staticmethod
def _function_tool_call_shapes(response: "ModelResponse") -> tuple[tuple[str | None, str], ...]:
return tuple(
(tool_call.function.name, tool_call.function.arguments)
for choice in response.choices
for tool_call in choice.message.tool_calls or ()
if isinstance(tool_call, ChatCompletionMessageToolCall)
)
@staticmethod
def _function_tool_call_fragments(
responses_so_far: Sequence["ModelResponseStream"],
) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]:
"""Group the stream's function tool-call fragments by their tool-call index, in
the index order ``stream_chunk_builder`` lists the rebuilt tool calls, keeping
only the indices the builder keeps (an id and a name somewhere in the stream)."""
fragments: Final = tuple(
tool_call
for response in responses_so_far
for choice in response.choices
for tool_call in choice.delta.tool_calls or ()
if isinstance(tool_call, ChatCompletionDeltaToolCall)
)
identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id)
named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name)
return tuple(
tuple(fragment for fragment in fragments if fragment.index == index) for index in sorted(identified & named)
)
def _write_ended_stream_tool_call_rewrites(
self,
responses_so_far: list["ModelResponseStream"], # mutable-ok: rewrites the caller's buffered chunks in place
guardrailed_response: "ModelResponse",
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites back across the buffered
chunks: the rewritten name and full arguments land in the tool call's first
fragment and the arguments of its later fragments are blanked, mirroring the
text write-back. A rewrite on a stream carrying more than one distinct choice
index, or whose fragments do not line up with the rebuilt tool calls, is
reported as undeliverable, so the pipeline executor discards it and releases
the original chunks."""
post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response)
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
stream_choice_indices: Final = frozenset(
choice.index for response in responses_so_far for choice in response.choices
)
fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far)
if len(stream_choice_indices) != 1 or len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
for before, (name, arguments), fragments in zip(
pre_guardrail_tool_calls, post_guardrail_tool_calls, fragments_by_tool_call
):
if (name, arguments) == before:
continue
head, *tail = fragments
head.function.name = name
head.function.arguments = arguments
for fragment in tail:
fragment.function.arguments = ""
async def _apply_guardrail_responses_to_output_streaming(
self,
responses: list["ModelResponseStream"],

View file

@ -37,14 +37,14 @@ from itertools import accumulate, chain, repeat
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Union, cast
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
from pydantic import BaseModel, TypeAdapter
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
tool_call_dict_from_output_item,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
@ -84,7 +84,6 @@ from litellm.types.llms.openai import (
)
from litellm.types.responses.main import (
GenericResponseOutputItem,
OutputFunctionToolCall,
OutputText,
)
from litellm.types.utils import GenericGuardrailAPIInputs
@ -101,6 +100,72 @@ if TYPE_CHECKING:
from litellm.types.llms.openai import ResponseInputParam
class _ToolCallShape(NamedTuple):
name: str | None
arguments: str
class _ToolCallFunctionFields(BaseModel):
model_config = ConfigDict(frozen=True)
name: str | None = None
arguments: str = ""
class _ToolCallFields(BaseModel):
model_config = ConfigDict(frozen=True)
function: _ToolCallFunctionFields
def _tool_call_shapes(tool_calls: Sequence[ChatCompletionToolCallChunk]) -> tuple[_ToolCallShape, ...]:
return tuple(
_ToolCallShape(name=tool_call["function"].get("name"), arguments=tool_call["function"].get("arguments", ""))
for tool_call in tool_calls
)
def _returned_tool_call_shape(tool_call: object) -> _ToolCallShape | None:
payload: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call
try:
fields: Final = _ToolCallFields.model_validate(payload)
except ValidationError:
return None
return _ToolCallShape(name=fields.function.name, arguments=fields.function.arguments)
def _post_guardrail_tool_call_shapes(
returned_tool_calls: Sequence[object] | None,
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str | None,
) -> tuple[_ToolCallShape, ...]:
if not pre_guardrail_tool_calls:
return pre_guardrail_tool_calls
if returned_tool_calls is None or len(returned_tool_calls) != len(pre_guardrail_tool_calls):
verbose_proxy_logger.warning(
"OpenAI Responses API: guardrail %s returned %s tool calls for the %d scanned, "
"leaving the tool call output items unchanged",
guardrail_name,
"no" if returned_tool_calls is None else len(returned_tool_calls),
len(pre_guardrail_tool_calls),
)
return pre_guardrail_tool_calls
returned_shapes: Final = tuple(_returned_tool_call_shape(tool_call) for tool_call in returned_tool_calls)
validated_shapes: Final = tuple(shape for shape in returned_shapes if shape is not None)
if len(validated_shapes) != len(returned_shapes):
verbose_proxy_logger.warning(
"OpenAI Responses API: guardrail %s returned tool calls without a function name and arguments, "
"leaving the tool call output items unchanged",
guardrail_name,
)
return pre_guardrail_tool_calls
return validated_shapes
def _tool_call_rewrite(before: _ToolCallShape, after: _ToolCallShape) -> _ToolCallShape:
return _ToolCallShape(name=after.name if after.name != before.name else None, arguments=after.arguments)
class ResponseOutputEnvelope(TypedDict, total=False):
"""Dict form of a Responses API response, as far as guardrail write-back reads it."""
@ -128,6 +193,20 @@ _TERMINAL_ENVELOPE_EVENT_TYPES: Final = frozenset(
)
_TOOL_CALL_ITEM_TYPES: Final = frozenset({"function_call", "custom_tool_call"})
_TOOL_CALL_PAYLOAD_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call": "arguments", "custom_tool_call": "input"}
)
_TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES: Final = frozenset(
{"response.function_call_arguments.delta", "response.custom_tool_call_input.delta"}
)
_TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"response.function_call_arguments.done": "arguments", "response.custom_tool_call_input.done": "input"}
)
_TOOL_CALL_PAYLOAD_EVENT_TYPES: Final = _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES | frozenset(
_TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS
)
_OUTPUT_ITEM_EVENT_TYPES: Final = frozenset({"response.output_item.added", "response.output_item.done"})
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call_output": "output", "message": "content"}
)
@ -164,8 +243,20 @@ def _rewritten_input_item(item: Mapping[str, object], rewritten: object) -> Mapp
return {**item, field: converted_value} # mutable-ok: request input items must stay JSON-plain dicts
def _is_function_call_item(item: object) -> bool:
return isinstance(item, Mapping) and item.get("type") in ("function_call", "custom_tool_call")
def _is_tool_call_item(item: object) -> bool:
return isinstance(item, Mapping) and item.get("type") in _TOOL_CALL_ITEM_TYPES
def _tool_call_output_item_mapping(item: object) -> Mapping[str, object] | None:
if stream_item_field(item, "type") not in _TOOL_CALL_ITEM_TYPES:
return None
if isinstance(item, Mapping):
return cast("Mapping[str, object]", item) # cast-ok: output items are str-keyed JSON objects
return item.model_dump() if isinstance(item, BaseModel) else None
def _is_tool_call_output_item(item: object) -> bool:
return _tool_call_output_item_mapping(item) is not None
def _last_message_role(messages: Sequence[object]) -> str | None:
@ -189,7 +280,7 @@ def _provenance_unit_bounds(
start_indexes: Final = tuple(
index
for index in range(len(raw_input))
if index == 0 or not (_is_function_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
if index == 0 or not (_is_tool_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
)
return tuple(zip(start_indexes, (*start_indexes[1:], len(raw_input))))
@ -340,7 +431,8 @@ class OpenAIResponsesHandler(BaseTranslation):
Methods can be overridden to customize behavior for different message formats.
"""
delivers_ended_stream_text_rewrites = True
delivers_ended_stream_rewrites = True
assembles_streamed_response = True
def get_structured_messages(self, data: dict) -> list[AllMessageValues] | None:
"""
@ -587,7 +679,7 @@ class OpenAIResponsesHandler(BaseTranslation):
- response.output is a list of output items
- Each output item can be:
* GenericResponseOutputItem with a content list of OutputText objects
* ResponseFunctionToolCall with tool call data
* ResponseFunctionToolCall or CustomToolCallOutputItem with tool call data
- Each OutputText object has a text field
"""
@ -652,6 +744,7 @@ class OpenAIResponsesHandler(BaseTranslation):
if response_model:
inputs["model"] = response_model
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
@ -660,6 +753,11 @@ class OpenAIResponsesHandler(BaseTranslation):
)
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
post_guardrail_tool_calls: Final = _post_guardrail_tool_call_shapes(
returned_tool_calls=guardrailed_inputs.get("tool_calls"),
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
guardrail_name=guardrail_to_apply.guardrail_name,
)
# Step 3: Map guardrail responses back to original response structure
await self._apply_guardrail_responses_to_output(
@ -667,6 +765,11 @@ class OpenAIResponsesHandler(BaseTranslation):
responses=guardrailed_texts,
task_mappings=task_mappings,
)
self._write_tool_call_rewrites_to_output(
tool_call_items=tuple(item for item in response_output if _is_tool_call_output_item(item)),
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=post_guardrail_tool_calls,
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed output response: %s", response)
@ -754,6 +857,7 @@ class OpenAIResponsesHandler(BaseTranslation):
if response_model:
inputs["model"] = response_model
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
@ -762,6 +866,11 @@ class OpenAIResponsesHandler(BaseTranslation):
)
guardrailed_texts: Final = guardrailed_inputs.get("texts", [])
post_guardrail_tool_calls: Final = _post_guardrail_tool_call_shapes(
returned_tool_calls=guardrailed_inputs.get("tool_calls"),
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
guardrail_name=guardrail_to_apply.guardrail_name,
)
# Write guardrailed texts back into the output items in-place.
# final_chunk is a reference into responses_so_far so this
@ -784,6 +893,13 @@ class OpenAIResponsesHandler(BaseTranslation):
stream_events=responses_so_far[:-1],
rewrites_by_position=rewrites_by_position,
)
self._deliver_ended_stream_tool_call_rewrites(
responses_so_far=responses_so_far,
outputs=outputs,
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
post_guardrail_tool_calls=post_guardrail_tool_calls,
guardrail_name=guardrail_to_apply.guardrail_name or "unknown",
)
return responses_so_far
# ------------------------------------------------------------------ #
@ -894,6 +1010,148 @@ class OpenAIResponsesHandler(BaseTranslation):
continue
OpenAIResponsesHandler._write_event_field(content[content_idx], "text", rewritten)
def _deliver_ended_stream_tool_call_rewrites(
self,
responses_so_far: Sequence[object],
outputs: Sequence[object],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
guardrail_name: str,
) -> None:
"""Write ended-stream guardrail tool-call rewrites into the completed
envelope's ``function_call`` and ``custom_tool_call`` items and sync the
earlier stream events, keyed by ``call_id``. The guardrail sees the
envelope's tool calls in output order, which is how a rewritten call
finds its ``call_id``; the stream events find their call through the
``call_id`` on ``output_item`` events and the ``item_id`` on argument
and custom-input events, since an
event's ``output_index`` need not match the envelope's (the chat bridge
numbers tool calls from 1 while the envelope lists them after the
message). A rewrite whose calls do not line up with the envelope, or
whose events cannot be found, is reported as undeliverable, so the
pipeline executor discards it and releases the original events."""
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
tool_call_items: Final = tuple(output_item for output_item in outputs if _is_tool_call_output_item(output_item))
call_ids: Final = tuple(
call_id
for output_item in tool_call_items
if isinstance(call_id := stream_item_field(output_item, "call_id"), str) and call_id
)
stream_events: Final = responses_so_far[:-1]
call_id_by_item_id: Final = self._tool_call_ids_by_item_id(stream_events)
event_call_ids: Final = tuple(
self._tool_call_event_call_id(event, call_id_by_item_id) for event in stream_events
)
rewrites_by_call_id: Final = MappingProxyType(
{
call_id: _tool_call_rewrite(before, after)
for call_id, before, after in zip(call_ids, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
}
)
unresolved_argument_event: Final = any(
call_id is None and stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_EVENT_TYPES
for event, call_id in zip(stream_events, event_call_ids)
)
if (
len(call_ids) != len(tool_call_items)
or len(frozenset(call_ids)) != len(call_ids)
or len(call_ids) != len(post_guardrail_tool_calls)
or unresolved_argument_event
or not rewrites_by_call_id.keys() <= frozenset(event_call_ids)
):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
raise UndeliverableStreamRewrite(guardrail_name)
for output_item, rewrite in (
(output_item, rewrites_by_call_id[call_id])
for output_item, call_id in zip(tool_call_items, call_ids)
if call_id in rewrites_by_call_id
):
self._write_tool_call_item(output_item, rewrite.name, rewrite.arguments)
delta_replacements: Final = MappingProxyType(
{call_id: chain((rewrite.arguments,), repeat("")) for call_id, rewrite in rewrites_by_call_id.items()}
)
for event, call_id in zip(stream_events, event_call_ids):
if call_id not in rewrites_by_call_id:
continue
match stream_item_field(event, "type"):
case str() as event_type if event_type in _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES:
self._write_event_field(event, "delta", next(delta_replacements[call_id]))
case str() as event_type if event_type in _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS:
self._write_event_field(
event, _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS[event_type], rewrites_by_call_id[call_id].arguments
)
case "response.output_item.added":
self._write_tool_call_item(
stream_item_field(event, "item"), rewrites_by_call_id[call_id].name, None
)
case "response.output_item.done":
self._write_tool_call_item(
stream_item_field(event, "item"),
rewrites_by_call_id[call_id].name,
rewrites_by_call_id[call_id].arguments,
)
case _:
pass
def _write_tool_call_rewrites_to_output(
self,
tool_call_items: Sequence[object],
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
post_guardrail_tool_calls: tuple[_ToolCallShape, ...],
) -> None:
if len(tool_call_items) != len(post_guardrail_tool_calls):
return
for output_item, rewrite in (
(output_item, _tool_call_rewrite(before, after))
for output_item, before, after in zip(tool_call_items, pre_guardrail_tool_calls, post_guardrail_tool_calls)
if after != before
):
self._write_tool_call_item(output_item, rewrite.name, rewrite.arguments)
@staticmethod
def _tool_call_ids_by_item_id(stream_events: Sequence[object]) -> Mapping[str, str]:
items: Final = tuple(
stream_item_field(event, "item")
for event in stream_events
if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES
)
return MappingProxyType(
{
item_id: call_id
for item in items
if stream_item_field(item, "type") in _TOOL_CALL_ITEM_TYPES
and isinstance(item_id := stream_item_field(item, "id"), str)
and isinstance(call_id := stream_item_field(item, "call_id"), str)
}
)
@staticmethod
def _tool_call_event_call_id(event: object, call_id_by_item_id: Mapping[str, str]) -> str | None:
event_type: Final = stream_item_field(event, "type")
if event_type in _TOOL_CALL_PAYLOAD_EVENT_TYPES:
item_id: Final = stream_item_field(event, "item_id")
return call_id_by_item_id.get(item_id) if isinstance(item_id, str) else None
if event_type not in _OUTPUT_ITEM_EVENT_TYPES:
return None
item: Final = stream_item_field(event, "item")
call_id: Final = stream_item_field(item, "call_id")
return (
call_id if stream_item_field(item, "type") in _TOOL_CALL_ITEM_TYPES and isinstance(call_id, str) else None
)
@staticmethod
def _write_tool_call_item(item: object, name: str | None, payload: str | None) -> None:
if item is None:
return
if name is not None:
OpenAIResponsesHandler._write_event_field(item, "name", name)
item_type: Final = stream_item_field(item, "type")
if payload is not None and isinstance(item_type, str) and item_type in _TOOL_CALL_PAYLOAD_FIELDS:
OpenAIResponsesHandler._write_event_field(item, _TOOL_CALL_PAYLOAD_FIELDS[item_type], payload)
def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool:
"""
Check if the streaming has ended.
@ -920,7 +1178,7 @@ class OpenAIResponsesHandler(BaseTranslation):
def _completed_response_scan_key(response: object) -> StreamingScanKey:
output_items: Final = stream_item_items(response, "output")
message_items: Final = tuple(
item for item in output_items if stream_item_field(item, "type") != "function_call"
item for item in output_items if stream_item_field(item, "type") not in _TOOL_CALL_ITEM_TYPES
)
return StreamingScanKey(
texts=tuple(
@ -932,7 +1190,7 @@ class OpenAIResponsesHandler(BaseTranslation):
tool_calls=tuple(
stream_item_fingerprint(item)
for item in output_items
if stream_item_field(item, "type") == "function_call"
if stream_item_field(item, "type") in _TOOL_CALL_ITEM_TYPES
),
stream_ended=True,
)
@ -1043,34 +1301,10 @@ class OpenAIResponsesHandler(BaseTranslation):
Override this method to customize text/image/tool extraction logic.
"""
# Check if this is a tool call (OutputFunctionToolCall)
if isinstance(output_item, OutputFunctionToolCall) or (
isinstance(output_item, BaseModel)
and hasattr(output_item, "type")
and getattr(output_item, "type") == "function_call"
):
tool_call_item: Final = _tool_call_output_item_mapping(output_item)
if tool_call_item is not None:
if tool_calls_to_check is not None:
tool_call_dict = (
LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
tool_call_item=output_item,
index=output_idx,
)
)
tool_calls_to_check.append(cast(ChatCompletionToolCallChunk, tool_call_dict))
return
elif isinstance(output_item, dict) and output_item.get("type") == "function_call":
# Handle dict representation of tool call
if tool_calls_to_check is not None:
# Convert dict to ResponseFunctionToolCall for processing
try:
tool_call_obj: Final = ResponseFunctionToolCall(**output_item)
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
tool_call_item=tool_call_obj,
index=output_idx,
)
tool_calls_to_check.append(cast(ChatCompletionToolCallChunk, tool_call_dict))
except Exception:
pass
tool_calls_to_check.append(tool_call_dict_from_output_item(tool_call_item, output_idx))
return
# Handle both GenericResponseOutputItem and dict

View file

@ -620,15 +620,20 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
return event_pydantic_model.model_construct(**parsed_chunk)
@staticmethod
def parse_terminal_response_from_stream_chunks(all_chunks: list[str]) -> ResponsesAPIResponse | None:
def parse_terminal_event_from_stream_chunks(all_chunks: Sequence[str]) -> ResponsesTerminalEvent | None:
for chunk_str in reversed(all_chunks):
for event_model in (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent):
try:
return event_model.model_validate_json(chunk_str.removeprefix("data: ")).response
return event_model.model_validate_json(chunk_str.removeprefix("data: "))
except ValueError:
continue
return None
@staticmethod
def parse_terminal_response_from_stream_chunks(all_chunks: list[str]) -> ResponsesAPIResponse | None:
terminal_event: Final = OpenAIResponsesAPIConfig.parse_terminal_event_from_stream_chunks(all_chunks)
return None if terminal_event is None else terminal_event.response
@staticmethod
def get_event_model_class(event_type: str) -> type[BaseLiteLLMOpenAIResponseObject]:
"""

View file

@ -19,7 +19,7 @@ import random
import sys
import time
import traceback
from collections.abc import AsyncIterator, Coroutine, Iterable, Mapping, Sequence
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Mapping, Sequence
from concurrent import futures
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from copy import deepcopy
@ -5398,6 +5398,14 @@ def completion(
if dynamic_api_key is not None:
api_key = dynamic_api_key
# check if user passed in any of the OpenAI optional params
bridges_to_responses_api: Final = (
responses_api_model_info.get("mode") == "responses" and not skip_responses_api_bridge
)
allowed_openai_params: Final[list[str] | None] = (
[*(kwargs.get("allowed_openai_params") or []), "reasoning_effort"]
if bridges_to_responses_api
else kwargs.get("allowed_openai_params")
)
optional_param_args: Final = {
"functions": functions,
"function_call": function_call,
@ -5442,7 +5450,7 @@ def completion(
"service_tier": service_tier,
"store": store,
"prompt_cache_key": prompt_cache_key,
"allowed_openai_params": kwargs.get("allowed_openai_params"),
"allowed_openai_params": allowed_openai_params,
"base_model": base_model,
}
optional_params = get_optional_params(**optional_param_args, **non_default_params)
@ -8587,7 +8595,7 @@ def config_completion(**kwargs):
)
def stream_chunk_builder_text_completion(chunks: list, messages: list | None = None) -> TextCompletionResponse:
def stream_chunk_builder_text_completion(chunks: list, messages: Sequence | None = None) -> TextCompletionResponse:
id: Final = chunks[0]["id"]
object: Final = chunks[0]["object"]
created: Final = chunks[0]["created"]
@ -8704,10 +8712,11 @@ def _stamp_streaming_usage_cost(usage: Usage, response: ModelResponse, logging_o
def stream_chunk_builder(
chunks: list,
messages: list | None = None,
messages: Sequence | None = None,
start_time=None,
end_time=None,
logging_obj: Optional["Logging"] = None,
count_prompt_tokens: Callable[[], int] | None = None,
) -> ModelResponse | TextCompletionResponse | None:
try:
if chunks is None:
@ -8781,6 +8790,7 @@ def stream_chunk_builder(
completion_output=completion_output,
messages=messages,
reasoning_tokens=0,
count_prompt_tokens=count_prompt_tokens,
)
setattr(response, "usage", usage)
@ -8958,6 +8968,7 @@ def stream_chunk_builder(
completion_output=completion_output,
messages=messages,
reasoning_tokens=reasoning_tokens,
count_prompt_tokens=count_prompt_tokens,
)
setattr(response, "usage", usage)

File diff suppressed because it is too large Load diff

View file

@ -179,16 +179,7 @@ def _gateway_dcr_challenge_target(
mcp_servers: list[str] | None,
client_ip: str | None,
) -> str | None:
"""The single path-named server this request targets, iff it resolves to a
gateway-managed oauth2 server the one per-server shape the gateway's own keyless
DCR flow serves end to end, so the 401 challenge may advertise the per-server
protected-resource metadata (whose ``authorization_servers`` names the gateway).
Multi-server CSV paths, header/path mismatches, unknown names, and every
client-forwarded or delegated mode return ``None``: those cells keep their existing
challenge (or absence of one), and a challenge is never emitted for a name the
public discovery routes would 404, so this reveals exactly the server set the
per-server protected-resource metadata already reveals."""
"""Resolve a single path target whose sign-in metadata advertises the gateway."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
@ -217,7 +208,7 @@ def _is_gateway_dcr_challenge_scope(
the caller is not a cold-start DCR client), on the scopes the gateway's keyless
flow serves: the aggregate ``/mcp`` endpoint, an ``x-mcp-servers``-scoped request
(the resource the client configured is still ``/mcp``), or a per-server path whose
single target is a gateway-managed oauth2 server. Every other named target keeps
single target advertises gateway-owned sign-in. Every other named target keeps
its existing behavior, failing closed to the original admission error."""
if not _is_litellm_auth_admission_error(exc):
return False
@ -236,7 +227,7 @@ def _gateway_dcr_challenge(
) -> HTTPException:
"""The RFC 9728 challenge pointing the client at the protected-resource metadata
matching the scope it requested: the per-server document (same URL spelling the
request arrived on) when the single target is a gateway-managed oauth2 server,
request arrived on) when the single target advertises gateway-owned sign-in,
else the gateway's aggregate document. Either way the client discovers the gateway
as its authorization server and starts the same sign-in flow.

View file

@ -2312,8 +2312,7 @@ async def _build_oauth_protected_resource_response(
it. Only the legacy ``is_oauth_passthrough`` opt-in rewrites ``resource`` to
the gateway's own URL so clients present the bearer token back to the gateway.
An explicitly named gateway-managed oauth2 server (interactive with
gateway-vaulted per-user tokens, or M2M) advertises the gateway's own
An explicitly named server with gateway-owned sign-in advertises the gateway's own
authorization server (``{base}/mcp``): a keyless DCR client that configured the
per-server URL completes the same sign-in flow the aggregate ``/mcp`` endpoint
supports and is admitted with a gateway session bearer. The per-server relay
@ -2403,11 +2402,6 @@ async def _build_oauth_protected_resource_response(
if obo_response is not None:
return obo_response
# An OBO server with no configured issuer falls through to the gateway default so discovery still
# returns metadata; every other non-oauth2 named server 404s to avoid enumeration.
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
if explicitly_named and mcp_server is not None and mcp_server.advertises_gateway_authorization_server:
return {
"authorization_servers": [f"{request_base_url}/mcp"],
@ -2415,6 +2409,9 @@ async def _build_oauth_protected_resource_response(
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
}
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
return {
"authorization_servers": [
(f"{request_base_url}/{mcp_server_name}" if mcp_server_name else f"{request_base_url}")

View file

@ -411,7 +411,7 @@ def relative_request_url(request: Request) -> str:
def resolve_scoped_resource_server(request: Request, resource: str | None) -> MCPServer | None:
"""Resolve an RFC 8707 ``resource`` value to the single gateway-managed oauth2 server it
"""Resolve an RFC 8707 ``resource`` value to the single gateway-owned server it
names, or ``None`` for every other shape: absent, the aggregate resource, a foreign
host, an unparseable value, a multi-server path, an unknown name, or any server mode the
keyless gateway flow does not serve (whose protected-resource metadata never directs a
@ -443,7 +443,7 @@ def resolve_scoped_resource_server(request: Request, resource: str | None) -> MC
if len(names) != 1:
return None
server: Final = global_mcp_server_manager.get_mcp_server_by_name(names[0])
if server is None or not server.is_gateway_managed_oauth2:
if server is None or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server):
return None
return server
@ -736,11 +736,15 @@ async def _flow_target(
server: Final = global_mcp_server_manager.get_mcp_server_by_id(flow.resource_server_id)
if (
server is None
or not server.is_gateway_managed_oauth2
or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server)
or not await lookup_server_reachability(flow.user_id, server.server_id)
):
return "stale", None
state: Final = "m2m" if MCPServerManager.effective_oauth2_flow(server) == "client_credentials" else "interactive"
state: Final = (
"interactive"
if server.is_gateway_managed_oauth2 and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
else "m2m"
)
return state, server

View file

@ -22,7 +22,16 @@ Response headers returned (all values are masked for safety):
x-mcp-debug-auth-resolution
Which auth priority was used for the outbound MCP call:
``per-request-header``, ``m2m-client-credentials``, ``static-token``,
``oauth2-passthrough``, or ``no-auth``.
``oauth2-passthrough``, ``stored-user-token``, ``token-exchange``,
``id-jag``, ``aws-sigv4``, ``extra-headers``, or ``no-auth``.
``unresolved`` means no outcome was available before the first response
frame; ``multiple`` means several servers resolved credentials;
``not-applicable`` covers stdio; ``resolution-failed`` is a resolver error.
x-mcp-debug-auth-resolutions
For multiple servers, a JSON map of server IDs to resolution labels.
At most 32 entries are included; x-mcp-debug-auth-resolutions-truncated
is true when additional servers were omitted. No credentials are included.
x-mcp-debug-outbound-url
The upstream MCP server URL that will receive the request.
@ -58,10 +67,16 @@ header is free for OAuth2 discovery::
Symptom: ``x-mcp-debug-oauth2-token`` shows ``(none)`` and
``x-mcp-debug-auth-resolution`` shows ``no-auth``.
This means the client didn't go through the OAuth2 flow. Check that:
1. The ``Authorization`` header is NOT set as a static header in the client config.
2. The ``.well-known/oauth-protected-resource`` endpoint returns valid metadata.
3. The MCP server in LiteLLM config has ``auth_type: oauth2``.
``no-auth`` means the resolved upstream client carries no authentication.
An absent inbound OAuth2 token does not imply the user skipped OAuth: the gateway
can retrieve a stored per-user token, reported as ``stored-user-token``.
``unresolved`` is used when a stream starts before credential resolution, or a
request (such as initialization or a cached tool listing) resolves no credential.
Debug reporting does not fetch credentials or delay a streaming frame to resolve them.
``extra-headers`` identifies supplied headers that won over the resolver or were
the only headers supplied; their values are never inspected to guess a scheme.
``per-request-header`` denotes a legacy credential override, including a BYOK
credential supplied by the gateway; it does not imply a caller-supplied token.
**Common issue: M2M token used instead of user token**
@ -69,8 +84,8 @@ Symptom: ``x-mcp-debug-auth-resolution`` shows ``m2m-client-credentials``.
This means the server has ``client_id``/``client_secret``/``token_url``
configured and LiteLLM is fetching a machine-to-machine token instead of
using the per-user OAuth2 token. If you want per-user tokens, remove the
client credentials from the server config.
using the per-user OAuth2 token. For gateway-stored per-user tokens,
configure ``oauth2_flow: authorization_code``.
Usage from Claude Code::
@ -85,14 +100,16 @@ Usage with curl::
http://localhost:4000/mcp/atlassian_mcp
"""
from typing import TYPE_CHECKING, Final
import json
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import Final
from starlette.requests import HTTPConnection
from starlette.types import Message, Send
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
# Header the client sends to opt into debug mode
MCP_DEBUG_REQUEST_HEADER: Final = "x-litellm-mcp-debug"
@ -101,6 +118,83 @@ MCP_DEBUG_REQUEST_HEADER: Final = "x-litellm-mcp-debug"
_RESPONSE_HEADER_PREFIX: Final = "x-mcp-debug"
MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: Final = "litellm.mcp.auth_diagnostics"
def record_auth_resolution(server_id: str, source: AuthResolution) -> None:
from mcp.server.lowlevel.server import request_ctx
context: Final[object] = request_ctx.get(None)
request: Final[object] = getattr(context, "request", None)
if isinstance(request, HTTPConnection):
diagnostics: Final[object] = request.scope.get(MCP_AUTH_DIAGNOSTICS_SCOPE_KEY)
if isinstance(diagnostics, MCPAuthDiagnostics):
diagnostics.record(server_id, source)
class MCPAuthDiagnostics:
def __init__(self) -> None:
self._outcomes: tuple[tuple[str, AuthResolution], ...] = ()
def record(self, server_id: str, resolution: AuthResolution) -> None:
self._outcomes = tuple(item for item in self._outcomes if item[0] != server_id) + ((server_id, resolution),)
def resolution(self) -> str:
match self._outcomes:
case ():
return AuthResolution.unresolved.value
case ((_, source),):
return source.value
case _:
return AuthResolution.multiple.value
def headers(self) -> Mapping[str, str]:
if len(self._outcomes) <= 1:
return MappingProxyType({"x-mcp-debug-auth-resolution": self.resolution()})
return MappingProxyType(
{
"x-mcp-debug-auth-resolution": AuthResolution.multiple.value,
"x-mcp-debug-auth-resolutions": json.dumps(
{
server_id: source.value for server_id, source in self._outcomes[:32]
}, # mutable-ok: JSON encoder requires a concrete dict
separators=(",", ":"),
ensure_ascii=True,
),
**(
MappingProxyType({"x-mcp-debug-auth-resolutions-truncated": "true"})
if len(self._outcomes) > 32
else MappingProxyType({})
),
}
)
class _DiagnosticSend:
def __init__(self, send: Send, headers: Mapping[str, str], resolution: Callable[[], Mapping[str, str]]) -> None:
self._send = send
self._headers = headers
self._resolution = resolution
self._start: Message | None = None
async def __call__(self, message: Message) -> None:
if message["type"] == "http.response.start":
self._start = message
return
if self._start is not None:
start: Final = self._start
self._start = None
headers: Final = MappingProxyType({**self._headers, **self._resolution()})
await self._send(
{ # mutable-ok: ASGI send consumes a mutable message mapping
**start,
"headers": tuple(start.get("headers", ()))
+ tuple((key.encode(), value.encode()) for key, value in headers.items()),
}
)
await self._send(message)
class MCPDebug:
"""
Static helper class for MCP OAuth2 debug headers.
@ -144,37 +238,6 @@ class MCPDebug:
return val.strip().lower() in ("true", "1", "yes")
return False
@staticmethod
def resolve_auth_resolution(
server: "MCPServer",
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
oauth2_headers: dict[str, str] | None,
) -> str:
"""
Determine which auth priority will be used for the outbound MCP call.
Returns one of: ``per-request-header``, ``m2m-client-credentials``,
``static-token``, ``oauth2-passthrough``, or ``no-auth``.
"""
from litellm.types.mcp import MCPAuth
has_server_specific: Final = bool(
mcp_server_auth_headers
and (
mcp_server_auth_headers.get(server.alias or "") or mcp_server_auth_headers.get(server.server_name or "")
)
)
if has_server_specific or mcp_auth_header:
return "per-request-header"
if server.has_client_credentials:
return "m2m-client-credentials"
if server.authentication_token:
return "static-token"
if oauth2_headers and server.auth_type == MCPAuth.oauth2:
return "oauth2-passthrough"
return "no-auth"
@staticmethod
def build_debug_headers(
*,
@ -244,12 +307,21 @@ class MCPDebug:
return debug
@staticmethod
def wrap_send_with_debug_headers(send: Send, debug_headers: dict[str, str]) -> Send:
def wrap_send_with_debug_headers(
send: Send,
debug_headers: Mapping[str, str],
resolution: Callable[[], Mapping[str, str]] | None = None,
*,
request_method: str | None = None,
) -> Send:
"""
Return a new ASGI ``send`` callable that injects *debug_headers*
into the ``http.response.start`` message.
"""
if resolution is not None and request_method == "POST":
return _DiagnosticSend(send, debug_headers, resolution)
async def _send_with_debug(message: Message) -> None:
if message["type"] == "http.response.start":
headers: Final = list(message.get("headers", []))
@ -266,8 +338,6 @@ class MCPDebug:
raw_headers: dict[str, str] | None,
scope: dict,
mcp_servers: list[str] | None,
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
oauth2_headers: dict[str, str] | None,
client_ip: str | None,
) -> dict[str, str]:
@ -288,16 +358,13 @@ class MCPDebug:
server_url: str | None = None
server_auth_type: str | None = None
auth_resolution = "no-auth"
auth_resolution: Final = AuthResolution.unresolved.value
for server_name in mcp_servers or []:
server = global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
if server:
server_url = server.url
server_auth_type = server.auth_type
auth_resolution = MCPDebug.resolve_auth_resolution(
server, mcp_auth_header, mcp_server_auth_headers, oauth2_headers
)
break
scope_headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)

View file

@ -80,6 +80,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
raise_classified_list_failure,
upstream_auth_challenge,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
@ -108,12 +109,14 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import resolve_credentials_with_source
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
build_token_exchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
AuthorizationCodeConfig,
AuthResolution,
ClientCredentialsConfig,
CredError,
IdJagConfig,
@ -3832,13 +3835,21 @@ class MCPServerManager:
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
"""
match await provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
case Ok(auth):
match await resolve_credentials_with_source(provider, to_subject(user_api_key_auth, subject_token), spec):
case Ok(credential):
auth: Final = credential.auth
# NoOpAuth has no header_name and so never conflicts.
header_name: Final[str | None] = getattr(auth, "header_name", None)
if header_name is None or not extra_headers:
source: Final = (
AuthResolution.extra_headers
if credential.source == AuthResolution.no_auth and extra_headers
else credential.source
)
record_auth_resolution(server.server_id, source)
return auth, extra_headers
if not has_header(extra_headers, header_name):
record_auth_resolution(server.server_id, credential.source)
return auth, extra_headers
if isinstance(
spec.config,
@ -3853,11 +3864,14 @@ class MCPServerManager:
# 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.
record_auth_resolution(server.server_id, credential.source)
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.
record_auth_resolution(server.server_id, AuthResolution.extra_headers)
return None, extra_headers
case Error(err):
record_auth_resolution(server.server_id, AuthResolution.failed)
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
# authorization_code's missing per-user token -> the per-server browser-OAuth
# challenge, built here where the full MCPServer is in hand.
@ -3960,6 +3974,7 @@ class MCPServerManager:
Returns:
Configured MCP client instance.
"""
record_auth_resolution(server.server_id, AuthResolution.unresolved)
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
transport: Final = resolved_server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
@ -4032,6 +4047,7 @@ class MCPServerManager:
env=resolved_env,
)
record_auth_resolution(server.server_id, AuthResolution.not_applicable)
return MCPClient(
server_url="", # Not used for stdio
transport_type=transport,
@ -4086,6 +4102,20 @@ class MCPServerManager:
aws_session_name=resolved_server.aws_session_name,
)
legacy_source: Final = (
AuthResolution.aws_sigv4
if aws_auth is not None
else AuthResolution.extra_headers
if extra_headers and has_header(extra_headers, auth_header_name or "Authorization")
else AuthResolution.per_request_header
if mcp_auth_header
else AuthResolution.static_token
if auth_value
else AuthResolution.extra_headers
if extra_headers
else AuthResolution.no_auth
)
record_auth_resolution(server.server_id, legacy_source)
return MCPClient(
server_url=server_url,
transport_type=transport,

View file

@ -65,6 +65,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
AuthorizationCodeConfig,
AuthResolution,
AuthSpecKind,
AwsSigV4Config,
Byok,
@ -76,6 +77,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
NoneConfig,
PassthroughConfig,
PrivateKeyJwtAuth,
ResolvedCredential,
ServerSpec,
SharedKey,
Subject,
@ -448,3 +450,32 @@ def _client_auth_fingerprint(client_auth: ClientAuth) -> str:
def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
return Error(CredError.of_not_implemented(f"{kind.value}: resolver arm not implemented yet"))
async def resolve_credentials_with_source(
provider: UpstreamCredentialProvider, subject: Subject, server: ServerSpec
) -> Result[ResolvedCredential, CredError]:
match await provider.resolve_credentials(subject, server):
case Error(err):
return Error(err)
case Ok(auth):
if isinstance(auth, NoOpAuth):
return Ok(ResolvedCredential(auth, AuthResolution.no_auth))
match server.config:
case NoneConfig():
return Ok(ResolvedCredential(auth, AuthResolution.no_auth))
case ApiKeyConfig():
return Ok(ResolvedCredential(auth, AuthResolution.static_token))
case PassthroughConfig():
return Ok(ResolvedCredential(auth, AuthResolution.oauth2_passthrough))
case ClientCredentialsConfig():
return Ok(ResolvedCredential(auth, AuthResolution.client_credentials))
case TokenExchangeConfig():
return Ok(ResolvedCredential(auth, AuthResolution.token_exchange))
case IdJagConfig():
return Ok(ResolvedCredential(auth, AuthResolution.id_jag))
case AuthorizationCodeConfig():
return Ok(ResolvedCredential(auth, AuthResolution.stored_user_token))
case AwsSigV4Config():
return Ok(ResolvedCredential(auth, AuthResolution.aws_sigv4))
assert_never(server.config)

View file

@ -26,10 +26,11 @@ union (see `result.py`), not `expression.Result`.
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from dataclasses import dataclass, field
from enum import Enum
from typing import Annotated, Final, Literal
import httpx
from expression import case, tag, tagged_union
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
from typing_extensions import assert_never
@ -46,6 +47,29 @@ from litellm.types.mcp import (
)
class AuthResolution(str, Enum):
no_auth = "no-auth"
stored_user_token = "stored-user-token"
static_token = "static-token"
per_request_header = "per-request-header"
oauth2_passthrough = "oauth2-passthrough"
client_credentials = "m2m-client-credentials"
token_exchange = "token-exchange"
id_jag = "id-jag"
aws_sigv4 = "aws-sigv4"
extra_headers = "extra-headers"
not_applicable = "not-applicable"
unresolved = "unresolved"
failed = "resolution-failed"
multiple = "multiple"
@dataclass(frozen=True, slots=True)
class ResolvedCredential:
auth: httpx.Auth = field(repr=False)
source: AuthResolution
class AuthSpecKind(str, Enum):
"""The server's statically-declared upstream-auth mode — derived from its `config`.

View file

@ -1234,7 +1234,7 @@ if MCP_AVAILABLE:
return client_id, client_secret, scopes
_STAGED_AUTH_VALUE_AUTH_TYPES: Final = frozenset(
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.authorization)
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token)
)
@dataclass(frozen=True, slots=True)
@ -1243,6 +1243,17 @@ if MCP_AVAILABLE:
mcp_auth_header: str | None
oauth2_headers: dict[str, str] | None
def _preview_origin(url: str | None) -> tuple[str, str, int | None] | None:
if not url:
return None
try:
parsed: Final = httpx.URL(url)
except httpx.InvalidURL:
return None
if parsed.scheme not in ("http", "https") or not parsed.host:
return None
return parsed.scheme, parsed.host, parsed.port
def _stage_server_test(new_mcp_server_request: NewMCPServerRequest, headers: Headers) -> _StagedServerTest:
"""
Resolve the credentials a not-yet-saved server config carries for a preview call.
@ -1255,7 +1266,19 @@ if MCP_AVAILABLE:
MCPRequestHandler,
)
request: Final = _inherit_credentials_from_existing_server(new_mcp_server_request)
saved_server: Final = (
global_mcp_server_manager.get_mcp_server_by_id(new_mcp_server_request.server_id)
if new_mcp_server_request.server_id
else None
)
saved_origin: Final = _preview_origin(saved_server.url) if saved_server else None
preview_origin: Final = _preview_origin(new_mcp_server_request.url)
may_inherit: Final = new_mcp_server_request.auth_type not in _STAGED_AUTH_VALUE_AUTH_TYPES or (
saved_origin is not None and saved_origin == preview_origin
)
request: Final = (
_inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request
)
mcp_auth_header: Final = (
request.credentials.get("auth_value")
if request.auth_type in _STAGED_AUTH_VALUE_AUTH_TYPES and isinstance(request.credentials, dict)
@ -1318,8 +1341,15 @@ if MCP_AVAILABLE:
if _oauth2_flow == "client_credentials" and not request.token_url:
_oauth2_flow = None
# Static previews inherit credentials before this step, but must not resolve back to
# the saved record during client creation and discard the edited connection settings.
preview_server_id: Final = (
""
if request.auth_type in _STAGED_AUTH_VALUE_AUTH_TYPES or request.auth_type in (None, MCPAuth.none)
else request.server_id or ""
)
server_model: Final = MCPServer(
server_id=request.server_id or "",
server_id=preview_server_id,
name=request.alias or request.server_name or "",
url=request.url,
transport=request.transport,

View file

@ -49,7 +49,11 @@ from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_server_name,
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode
)
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCP_AUTH_DIAGNOSTICS_SCOPE_KEY,
MCPAuthDiagnostics,
MCPDebug,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
get_passthrough_www_authenticate,
@ -4472,13 +4476,15 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
scope=dict(scope),
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
client_ip=_client_ip,
)
if _debug_headers:
send = MCPDebug.wrap_send_with_debug_headers(send, _debug_headers)
diagnostics: Final = MCPAuthDiagnostics() if _debug_headers else None
if diagnostics is not None:
scope[MCP_AUTH_DIAGNOSTICS_SCOPE_KEY] = diagnostics
send = MCPDebug.wrap_send_with_debug_headers(
send, _debug_headers, diagnostics.headers, request_method=scope.get("method")
)
# Ensure session managers are initialized
if not _SESSION_MANAGERS_INITIALIZED:

View file

@ -27,6 +27,7 @@ from litellm.litellm_core_utils.url_utils import (
provider_url_destination_candidates,
validate_url,
)
from litellm.llms.azure.passthrough.transformation import azure_router_model_in_endpoint
from litellm.proxy._types import *
from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_metadata
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
@ -2003,9 +2004,20 @@ def get_model_from_request(
bedrock_model: Final = _model_from_bedrock_route(route)
return model if bedrock_model is None else bedrock_model
if route.lower().startswith(("/azure/", "/azure_ai/")):
azure_model: Final = _router_model_from_azure_route(route, llm_router)
return model if azure_model is None else azure_model
return model
def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str | None:
if llm_router is None:
return None
endpoint: Final = re.sub(r"^/azure(?:_ai)?/", "", route, flags=re.IGNORECASE)
return azure_router_model_in_endpoint(endpoint, frozenset(llm_router.get_model_names()))
def _model_from_bedrock_route(route: str) -> str | None:
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_extract_model_from_bedrock_endpoint,

View file

@ -510,7 +510,7 @@ If you belong to several teams, `lite login` normally asks which one to attribut
### Route Every Claude Code Session Through the Proxy
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when those keys are missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when those keys are missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
Two things need to already be true: you've run `lite login` (or `lite login --pkce`, whose key the helper renews on its own), since the apiKeyHelper depends on that stored token, and the proxy is already reachable, since `lite up` does not start one for you.
@ -534,12 +534,28 @@ Cursor is not supported: it has no equivalent file-based config to hot-patch thi
lite --base-url https://your-proxy.example.com login --config-claude
```
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY`, and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY`, and `apiKeyHelper`, but persistently: no foreground process to keep alive, and `lite unconfigure claude` restores what it changed (see below). Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
Because the credential is reached through `apiKeyHelper` rather than copied into the file, a later `lite login` refreshes it with no further action: Claude Code re-runs the helper on every request and picks up whatever token the most recent login stored. Nothing secret is written to `settings.json`.
Run it again to point Claude Code at a different proxy; the base URL and the helper are both rewritten. `lite up` and `--config-claude` manage the same file, so the flag refuses to run while a `lite up` session holds a backup, and tells you to run `lite down` first, rather than writing settings that `lite up` would silently revert when it stops.
#### Configuring Claude Code Once, With a Virtual Key or Your Login
`lite configure claude` wires Claude Code up persistently and `lite unconfigure claude` puts things back. It is what `lite login --config-claude` does, plus a pinned model and an undo, and it also takes a long-lived virtual key when that is what you have:
```bash
curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/install.sh | sh
lite --base-url https://your-proxy.example.com configure claude --api-key sk-... --model claude-auto
claude
```
With `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) the key is written into `env.ANTHROPIC_AUTH_TOKEN`. Without one, your `lite login` credential is used the way `--config-claude` uses it, through `apiKeyHelper`, so a later `lite login` (or a `--pkce` renewal) picks up on its own and nothing secret lands in the file; a missing or stale login is refreshed first. Either way the command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (the ones whose id contains `claude` or `anthropic`) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key, which has to be on `/v1/models` for the key. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute up` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window and sends no thinking parameters for it, so either name the group like a Claude model id or append `[1m]` to opt into the 1M window. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control
Plain `lite configure`, with no agent named, asks the same things interactively: which agents to wire (Claude Code today) and which of the proxy's models to start on, picked from `/v1/models` with a type-to-filter prompt
What the command changed is recorded in `~/.litellm/claude_configure_state.json` (previous values plus fingerprints of what was written, never a second copy of the key). `lite unconfigure claude` restores each of those keys only if it still holds what `configure` wrote, so anything you changed since is left alone and named in the output; a `settings.json` or `env` object that only existed because of `configure` is removed again. Ownership moves only by a write: running `configure` again (a re-login is one) refreshes the record only for the keys its merge changed, keeps the original snapshot of a key that still holds what it wrote, and snapshots afresh a key you changed in between, so `unconfigure` brings back whatever the repeat displaced and never adopts your edit as its own. A credential (`env.ANTHROPIC_API_KEY`, `env.ANTHROPIC_AUTH_TOKEN`, `apiKeyHelper`) is put back only when the restored file points at the `ANTHROPIC_BASE_URL` it was captured next to; otherwise it stays removed, the output says which server it belonged to, and the receipt is kept so pointing the URL back and running `unconfigure` again finishes the job. It also undoes `lite login --config-claude`, which writes through the same path. Like `--config-claude`, both refuse to run while a `lite up` or `lite autoroute up` session holds a backup, and that check comes before any login prompt or request
### QA Complexity-Based Auto-Routing Against Your Real Proxy
`lite autoroute` lets you try LiteLLM's complexity-based auto-routing -- picking a cheaper or more expensive model depending on how complex a prompt looks -- against models your key already has access to on your real, running proxy, without editing that proxy's `config.yaml` and without any real request ever bypassing it. It builds a second, throwaway proxy locally that forwards every request back to your real proxy, and points Claude Code at that local proxy for the duration of the session.
@ -586,7 +602,7 @@ An interactive wizard. It runs the same model-group discovery as above, splits t
The wizard writes the result to `~/.litellm/autorouter/config.yaml` with `0600` permissions, since the file embeds your real proxy API key. Every model referenced anywhere in that config -- tier targets, the classifier model, the embedding model -- becomes its own `litellm_proxy/<model-name>` deployment whose `api_base` and `api_key` point back at your real proxy. That is the trick that keeps your real proxy's config untouched: every actual network call this generates, whether it is the routed completion, an LLM-classifier call, or an embedding call, forwards transparently through your real, already-running proxy with your real key.
You do not need to tell Claude Code to request `autorouter` by name yourself: `lite autoroute up` also sets `ANTHROPIC_DEFAULT_SONNET_MODEL`, `ANTHROPIC_DEFAULT_HAIKU_MODEL`, and `ANTHROPIC_DEFAULT_OPUS_MODEL` to `autorouter` in `~/.claude/settings.json`, so every one of Claude Code's own model tiers requests it directly regardless of `/model` or whatever it defaults to otherwise. (A bare `model_name: "*"` deployment looks like the obvious way to catch any request instead, but litellm's Router looks up auto-router deployments by the literal requested model string with no wildcard resolution, so a `"*"` entry would never actually match real traffic -- these env var overrides are what makes it work.)
You do not need to tell Claude Code to request `autorouter` by name yourself: `lite autoroute up` also sets the top-level `model` and `ANTHROPIC_DEFAULT_SONNET_MODEL`, `ANTHROPIC_DEFAULT_HAIKU_MODEL`, `ANTHROPIC_DEFAULT_OPUS_MODEL` and `ANTHROPIC_DEFAULT_FABLE_MODEL` to `autorouter` in `~/.claude/settings.json` (and `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when missing, like every other wiring), so every one of Claude Code's own model tiers requests it directly regardless of `/model` or whatever it defaults to otherwise. (A bare `model_name: "*"` deployment looks like the obvious way to catch any request instead, but litellm's Router looks up auto-router deployments by the literal requested model string with no wildcard resolution, so a `"*"` entry would never actually match real traffic -- these env var overrides are what makes it work.)
You must run `configure` at least once before `up`; running `up` first fails with a clear error telling you to configure first.

View file

@ -4,6 +4,7 @@ import subprocess
import sys
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final, TypeAlias
@ -12,10 +13,12 @@ import requests
from pydantic import BaseModel, TypeAdapter, ValidationError
from .auth import CliContextObj, context_secret_vault, get_stored_api_key, login
from .claude_settings import claude_settings_path, lite_api_key_helper_configured
from .cmd_quoting import quote_for_cmd
from .pi import (
LITELLM_PROXY_API_KEY_ENV,
PI_PROVIDER_NAME,
ListingFailure,
PiSyncError,
fetch_model_ids,
fetch_model_limits,
@ -83,6 +86,8 @@ def build_agent_env(
base_url: str,
api_key: str,
profiles: frozenset[str],
*,
export_anthropic_token: bool = True,
) -> dict[str, str]:
"""Return a copy of base_env wired to route the agent through the proxy.
@ -97,12 +102,19 @@ def build_agent_env(
proxy's /v1/models; likewise left alone when already set.
pi ignores both base URL variables and instead resolves $LITELLM_PROXY_API_KEY
from its synced models.json provider entry.
With export_anthropic_token=False the bearer is left out (and any inherited
one dropped) so Claude Code asks its configured apiKeyHelper instead; Claude
Code prefers ANTHROPIC_AUTH_TOKEN over the helper and warns when both are set.
"""
env: Final = dict(base_env)
root: Final = base_url.rstrip("/")
if PROFILE_ANTHROPIC in profiles:
env[ANTHROPIC_BASE_URL_ENV] = root
env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key
if export_anthropic_token:
env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key
else:
env.pop(ANTHROPIC_AUTH_TOKEN_ENV, None)
env.pop(ANTHROPIC_API_KEY_ENV, None)
if ENABLE_TOOL_SEARCH_ENV not in env:
env[ENABLE_TOOL_SEARCH_ENV] = ENABLE_TOOL_SEARCH_VALUE
@ -165,7 +177,9 @@ def prepare_pi(
"""
ids: Final = fetch_model_ids(base_url, api_key, get=get)
if isinstance(ids, PiSyncError):
raise AgentRunError(ids.message)
raise AgentRunError(
f"{ids.message} pi would have nothing to run." if ids.kind is ListingFailure.EMPTY else ids.message
)
limits: Final = fetch_model_limits(base_url, api_key, get=get)
path: Final = models_json_path(base_env)
error: Final = sync_models_json(path, base_url, ids, limits)
@ -460,6 +474,7 @@ def run_agent(
launcher: Callable[[str, Sequence[str], Mapping[str, str]], None] = _hand_off,
reattach_terminal: Callable[[], None] | None = None,
preparers: Mapping[str, _Preparer] = MappingProxyType(_PREPARERS),
export_anthropic_token: bool = True,
) -> None:
"""Validate, wire the environment, and hand off to the agent.
@ -491,7 +506,9 @@ def run_agent(
env: Final = MappingProxyType(
{
**build_agent_env(env_before_sync, base_url, api_key, profiles),
**build_agent_env(
env_before_sync, base_url, api_key, profiles, export_anthropic_token=export_anthropic_token
),
**(_NO_EXTRA_ENV if isinstance(synced, ModelSyncSkipped) else synced),
}
)
@ -529,14 +546,26 @@ def resolve_api_key(ctx: click.Context) -> str:
_SKIP_VERIFY_HELP: Final = "Skip the pre-launch key check against the proxy."
def _helper_supplies_token(
ctx_obj: CliContextObj, base_url: str, profiles: frozenset[str], settings_path: Path
) -> bool:
if PROFILE_ANTHROPIC not in profiles or not ctx_obj.get("api_key_from_token_file"):
return False
return lite_api_key_helper_configured(base_url, settings_path)
def _launch(ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify: bool) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
started_interactive: Final = _is_interactive()
api_key: Final = resolve_api_key(ctx)
display_name, _ = agent_profile(binary)
display_name, profiles = agent_profile(binary)
settings_path: Final = claude_settings_path(os.environ)
helper_supplies_token: Final = _helper_supplies_token(ctx_obj, base_url, profiles, settings_path)
click.echo(f"litellm: routing {display_name} through proxy at {base_url.rstrip('/')}")
if helper_supplies_token:
click.echo(f"litellm: {display_name} reads its key from the apiKeyHelper in {settings_path}")
try:
run_agent(
@ -545,6 +574,7 @@ def _launch(ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify
[binary, *args],
skip_verify=skip_verify,
reattach_terminal=(_restore_controlling_terminal if started_interactive else None),
export_anthropic_token=not helper_supplies_token,
)
except AgentRunError as e:
raise click.ClickException(str(e))

View file

@ -1,3 +1,4 @@
import os
import sys
import time
import webbrowser
@ -40,10 +41,16 @@ from litellm.litellm_core_utils.cli_token_utils import (
)
from .claude_settings import (
CLAUDE_SETTINGS_PATH,
SETTINGS_FILE_OWNERS,
STARTING_MODEL_ROLE,
ApiKeyHelper,
ClaudeSettingsError,
write_claude_settings,
KeepModel,
claude_settings_path,
configure_claude_settings,
configure_state_path,
refuse_while_owned,
resolve_api_key_helper,
settings_file_owners,
)
from .pkce_login import (
Http,
@ -802,13 +809,24 @@ def _render_and_prompt_for_team_selection(teams: list[CliTeam]) -> str | None:
def _configure_claude_code(base_url: str) -> None:
"""Point Claude Code at base_url by patching ~/.claude/settings.json."""
"""Point Claude Code at base_url by patching the settings.json it reads, undoable with `lite unconfigure claude`."""
settings_path: Final = claude_settings_path(os.environ)
try:
write_claude_settings(base_url, CLAUDE_SETTINGS_PATH, SETTINGS_FILE_OWNERS)
configure_claude_settings(
base_url,
ApiKeyHelper(resolve_api_key_helper(base_url)),
KeepModel(),
settings_path,
configure_state_path(settings_path),
settings_file_owners(settings_path),
)
except ClaudeSettingsError as e:
raise click.ClickException(f"Logged in, but could not configure Claude Code: {e}")
click.echo(f"\nConfigured Claude Code: {CLAUDE_SETTINGS_PATH} now routes through {base_url.rstrip('/')}.")
click.echo("Your other Claude Code settings were left untouched. Restart Claude Code to pick this up.")
click.echo(f"\nConfigured Claude Code: {settings_path} now routes through {base_url.rstrip('/')}.")
click.echo(
"Your other Claude Code settings were left untouched. Restart Claude Code to pick this up. "
f"Undo with `lite unconfigure claude`; `lite configure claude --model` sets {STARTING_MODEL_ROLE}."
)
def _finish_login(base_url: str, api_key: str, config_claude: bool, stored: SecretSave) -> None:
@ -889,6 +907,12 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool, team: str | None)
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
if config_claude:
settings_path: Final = claude_settings_path(os.environ)
try:
refuse_while_owned(settings_path, settings_file_owners(settings_path))
except ClaudeSettingsError as e:
raise click.ClickException(f"Cannot configure Claude Code, so not logging in: {e}")
try:
if pkce:

View file

@ -14,11 +14,13 @@ from ..claude_settings import (
AUTOROUTE_BACKUP_PATH,
CLAUDE_SETTINGS_PATH,
ClaudeSettingsError,
StaticToken,
load_json_or_empty,
merge_claude_settings,
)
from ..up import BackupRecord as ClaudeBackupRecord
from ..up import restore_claude_settings, write_backup
from .config import master_key_from_config
from .config import AUTOROUTER_MODEL_NAME, master_key_from_config
from .process import (
CONFIG_PATH,
DEFAULT_AUTOROUTE_PORT,
@ -37,7 +39,6 @@ from .process import (
terminate,
write_pid_record,
)
from .settings import merge_claude_settings_static_token
from .wizard import run_configure_wizard
_GENERATED_CONFIG_ADAPTER: Final = TypeAdapter(dict[str, JsonValue])
@ -156,7 +157,9 @@ def up(port: int) -> None:
ClaudeBackupRecord(existed=original_existed, content=original_settings if original_existed else None),
AUTOROUTE_BACKUP_PATH,
)
merged: Final = merge_claude_settings_static_token(original_settings, base_url, master_key)
merged: Final = merge_claude_settings(
original_settings, base_url, StaticToken(master_key), AUTOROUTER_MODEL_NAME, AUTOROUTER_MODEL_NAME
)
CLAUDE_SETTINGS_PATH.parent.mkdir(parents=True, exist_ok=True)
with secure_create(CLAUDE_SETTINGS_PATH) as f:
json.dump(merged, f, indent=2)

View file

@ -1,51 +0,0 @@
from typing import Final
from pydantic import JsonValue
from .config import AUTOROUTER_MODEL_NAME
ENV_KEY: Final = "env"
API_KEY_HELPER_KEY: Final = "apiKeyHelper"
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
ANTHROPIC_AUTH_TOKEN_KEY: Final = "ANTHROPIC_AUTH_TOKEN"
ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH"
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
# Force every one of Claude Code's own model tiers to request the auto-router by name.
# Router's auto-router registry is keyed by the literal requested model string
# (litellm/router.py:10711-10717) with no wildcard/pattern resolution, so a bare "*"
# model_name can never work as a catch-all -- these overrides are what actually makes
# Claude Code send "autorouter" regardless of /model or its own version-specific defaults.
ANTHROPIC_DEFAULT_MODEL_ENV_KEYS: Final = (
"ANTHROPIC_DEFAULT_SONNET_MODEL",
"ANTHROPIC_DEFAULT_HAIKU_MODEL",
"ANTHROPIC_DEFAULT_OPUS_MODEL",
)
def merge_claude_settings_static_token(
settings: dict[str, JsonValue], base_url: str, auth_token: str
) -> dict[str, JsonValue]:
"""Return a new settings dict wired to a local ephemeral proxy with a static token.
Unlike up.py's merge_claude_settings (which sets apiKeyHelper for a long-lived, real
remote proxy needing refreshable SSO tokens), this proxy is ephemeral and its key is the
locally persisted autoroute master key, so a plain env var is simpler and correct. Any
existing apiKeyHelper is cleared so it can't fight with the static token.
"""
raw_env: Final = settings.get(ENV_KEY, {})
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
env: Final[dict[str, JsonValue]] = {
ENABLE_TOOL_SEARCH_KEY: ENABLE_TOOL_SEARCH_VALUE,
**base_env,
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
ANTHROPIC_AUTH_TOKEN_KEY: auth_token,
**{key: AUTOROUTER_MODEL_NAME for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS},
}
env.pop(ANTHROPIC_API_KEY_KEY, None)
merged: Final[dict[str, JsonValue]] = {**settings, ENV_KEY: env}
merged.pop(API_KEY_HELPER_KEY, None)
return merged
__all__ = ["merge_claude_settings_static_token"]

View file

@ -1,37 +1,71 @@
"""Shared handling of Claude Code's ~/.claude/settings.json.
`lite up` patches this file temporarily and restores it on exit; `lite login
--config-claude` patches it persistently. Both need the same merge and the same
apiKeyHelper command, and `up` already imports from `auth`, so the shared parts
live here rather than in either command module.
`lite up` and `lite autoroute up` patch this file temporarily and restore it on
exit; `lite login --config-claude` and `lite configure claude` patch it
persistently and record how to undo it. All of them need the same merge and the
same apiKeyHelper command, and `up` already imports from `auth`, so the shared
parts live here rather than in any one command module.
"""
import hashlib
import json
import shlex
import shutil
import sys
from collections.abc import Mapping, Sequence
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from functools import reduce
from itertools import chain
from pathlib import Path
from typing import Final
from types import MappingProxyType
from typing import Final, TypeAlias
from pydantic import JsonValue, TypeAdapter, ValidationError
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
from litellm.litellm_core_utils.private_json import write_private_json
from litellm.litellm_core_utils.private_json import (
commit_staged_json,
discard_staged_json,
ensure_private_dir,
stage_private_json,
)
from .cmd_quoting import quote_for_cmd
ENV_KEY: Final = "env"
API_KEY_HELPER_KEY: Final = "apiKeyHelper"
MODEL_KEY: Final = "model"
ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
ANTHROPIC_AUTH_TOKEN_KEY: Final = "ANTHROPIC_AUTH_TOKEN"
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH"
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: Final = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE: Final = "1"
ANTHROPIC_DEFAULT_MODEL_ENV_KEYS: Final = (
"ANTHROPIC_DEFAULT_SONNET_MODEL",
"ANTHROPIC_DEFAULT_HAIKU_MODEL",
"ANTHROPIC_DEFAULT_OPUS_MODEL",
"ANTHROPIC_DEFAULT_FABLE_MODEL",
)
OWNED_ENV_KEYS: Final = (
ENABLE_TOOL_SEARCH_KEY,
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY,
ANTHROPIC_BASE_URL_KEY,
ANTHROPIC_AUTH_TOKEN_KEY,
ANTHROPIC_API_KEY_KEY,
)
OWNED_TOP_LEVEL_KEYS: Final = (API_KEY_HELPER_KEY, MODEL_KEY)
OWNED_PATHS: Final = (*(f"{ENV_KEY}.{key}" for key in OWNED_ENV_KEYS), *OWNED_TOP_LEVEL_KEYS)
_CREDENTIAL_ENV_KEYS: Final = frozenset((ANTHROPIC_API_KEY_KEY, ANTHROPIC_AUTH_TOKEN_KEY))
_CREDENTIAL_PATHS: Final = (*(f"{ENV_KEY}.{key}" for key in sorted(_CREDENTIAL_ENV_KEYS)), API_KEY_HELPER_KEY)
_BASE_URL_PATH: Final = f"{ENV_KEY}.{ANTHROPIC_BASE_URL_KEY}"
STARTING_MODEL_ROLE: Final = "the /model picker's default row, the model Claude Code starts on"
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
CLAUDE_CONFIG_DIR_ENV: Final = "CLAUDE_CONFIG_DIR"
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
AUTOROUTE_BACKUP_PATH: Final = Path.home() / ".litellm" / "autorouter" / "claude_settings_backup.json"
CONFIGURE_STATE_PATH: Final = Path.home() / ".litellm" / "claude_configure_state.json"
@dataclass(frozen=True, slots=True)
@ -55,6 +89,129 @@ class ClaudeSettingsError(Exception):
"""Raised for any user-actionable failure while reading or writing Claude Code settings."""
def claude_settings_path(environ: Mapping[str, str]) -> Path:
"""The settings.json Claude Code reads: under CLAUDE_CONFIG_DIR when set, else ~/.claude/settings.json."""
config_dir: Final = environ.get(CLAUDE_CONFIG_DIR_ENV, "")
if not config_dir:
return CLAUDE_SETTINGS_PATH
return Path(config_dir).expanduser() / "settings.json"
def _is_default_settings_file(settings_path: Path) -> bool:
return settings_path.resolve() == CLAUDE_SETTINGS_PATH.resolve()
def settings_file_owners(settings_path: Path) -> tuple[SettingsFileOwner, ...]:
"""The commands whose backups guard settings_path: `lite up` and `lite autoroute up` only ever manage the default file."""
return SETTINGS_FILE_OWNERS if _is_default_settings_file(settings_path) else ()
def configure_state_path(settings_path: Path) -> Path:
"""The receipt describing settings_path: the default file keeps CONFIGURE_STATE_PATH, and any other file
(a CLAUDE_CONFIG_DIR) gets its own beside it, keyed by its resolved path, so two settings files never
share one undo record."""
if _is_default_settings_file(settings_path):
return CONFIGURE_STATE_PATH
digest: Final = hashlib.sha256(str(settings_path.resolve()).encode()).hexdigest()
return CONFIGURE_STATE_PATH.parent / CONFIGURE_STATE_PATH.stem / f"{digest}.json"
@dataclass(frozen=True, slots=True)
class StaticToken:
"""A long-lived virtual key, written into env.ANTHROPIC_AUTH_TOKEN."""
token: str
@dataclass(frozen=True, slots=True)
class ApiKeyHelper:
"""A `lite auth print-token` command Claude Code runs per request, so a login renews in place."""
command: str
ClaudeCredential: TypeAlias = StaticToken | ApiKeyHelper
@dataclass(frozen=True, slots=True)
class KeepModel:
"""Leave the top-level `model` as it is, the user's or an earlier configure's (a re-login)."""
@dataclass(frozen=True, slots=True)
class UnpinModel:
"""Let go of a `model` an earlier configure pinned; one the user set themselves stays."""
@dataclass(frozen=True, slots=True)
class StartOn:
"""Pin the top-level `model`, the row Claude Code starts on."""
model: str
ModelChoice: TypeAlias = KeepModel | UnpinModel | StartOn
class OwnedValue(BaseModel):
"""What one key held at a moment in time; `present=False` is an absent key, not a null one."""
model_config = ConfigDict(frozen=True)
present: bool
value: JsonValue = None
class ConfigureReceipt(BaseModel):
"""What `lite configure claude` found and what it owns, keyed by dotted path (`env.X` or a top-level key).
Ownership moves only by a write: `written` fingerprints the keys some configure changed, at the
value it wrote; a repeat configure refreshes a fingerprint only for a key its merge changed and
carries the earlier one otherwise, so a key the user edited in between stops matching and is left
alone. `previous` is what each key held before configure took it over; a repeat keeps the earlier
snapshot while the key still holds our value and snapshots afresh otherwise, so whatever the
repeat displaces is what comes back. `endpoints` is the ANTHROPIC_BASE_URL each credential slot
was captured beside, so a credential is only ever put back next to the server it was issued for.
No fingerprint is a second copy of a token.
"""
model_config = ConfigDict(frozen=True)
file_existed: bool
env_present: bool
env_was_object: bool
previous: Mapping[str, OwnedValue]
written: Mapping[str, str]
endpoints: Mapping[str, OwnedValue]
@dataclass(frozen=True, slots=True)
class WithheldCredential:
"""A credential left removed: captured beside `endpoint`, while the restored file points elsewhere."""
key: str
endpoint: str
@dataclass(frozen=True, slots=True)
class UnconfigureOutcome:
"""Keys whose value unconfigure changed back, keys the user changed since and so were left as they
are, credentials withheld (the receipt is kept for them, so a later unconfigure can finish once the
URL points back), and whether no settings file remains."""
restored: tuple[str, ...]
kept: tuple[str, ...]
withheld: tuple[WithheldCredential, ...] = ()
file_removed: bool = False
@dataclass(frozen=True, slots=True)
class _Claim:
previous: OwnedValue
written: str | None
endpoint: OwnedValue | None
def load_json_or_empty(path: Path) -> dict[str, JsonValue]:
try:
content: Final = path.read_bytes() if path.exists() else b""
@ -70,29 +227,104 @@ def load_json_or_empty(path: Path) -> dict[str, JsonValue]:
)
def merge_claude_settings(
settings: Mapping[str, JsonValue], base_url: str, api_key_helper: str
) -> dict[str, JsonValue]:
"""Return a new settings dict wired to route Claude Code through the proxy.
def _env_object(settings: Mapping[str, JsonValue], path: Path) -> Mapping[str, JsonValue]:
raw_env: Final = settings.get(ENV_KEY)
if raw_env is None:
return MappingProxyType({})
if not isinstance(raw_env, dict):
raise ClaudeSettingsError(
f'{path} has a non-object "{ENV_KEY}" value, which this would discard. Fix or remove it, then retry.'
)
return raw_env
Only env.ANTHROPIC_BASE_URL and the top-level apiKeyHelper are overridden; a
stray env.ANTHROPIC_API_KEY is dropped so it cannot outrank the helper-issued
token (same reasoning as build_agent_env in agents.py). ENABLE_TOOL_SEARCH
defaults to true because Claude Code turns tool search off when
ANTHROPIC_BASE_URL is not a first-party Anthropic host, and
CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY defaults to 1 so the /model picker
is filled from the proxy's /v1/models; existing values of both are left
alone. Every other key is preserved untouched.
def refuse_while_owned(settings_path: Path, owners: Sequence[SettingsFileOwner]) -> None:
"""Refuse while `lite up` or `lite autoroute up` holds a backup it will restore over any write; a
purely local check, so commands run it before any login prompt or request."""
for owner in owners:
if owner.backup_path.exists():
raise ClaudeSettingsError(
f"`{owner.start_command}` is currently managing {settings_path} (backup at "
f"{owner.backup_path}) and will restore it when it stops. "
f"Run `{owner.stop_command}` first, then retry."
)
def _write_target(settings_path: Path) -> Path:
"""Write through a symlinked settings.json rather than replacing the link, which would silently
detach a file symlinked into a dotfiles repo."""
try:
return settings_path.resolve() if settings_path.is_symlink() else settings_path
except OSError as e:
raise ClaudeSettingsError(f"Could not resolve {settings_path}: {e}") from e
def _stage(path: Path, document: Mapping[str, object]) -> str:
try:
return stage_private_json(str(path), document)
except OSError as e:
raise ClaudeSettingsError(f"Could not write {path}: {e}") from e
def _land(
path: Path,
staged: str | None,
also_discard: Sequence[str | None] = (),
commit: Callable[[str, str], None] = commit_staged_json,
) -> None:
"""Commit a staged file to `path`, or remove `path` when nothing is staged for it. The one place a
filesystem error becomes a ClaudeSettingsError; on failure the operation's other staged files are
discarded, so no temp file holding a token is left behind."""
try:
if staged is None:
path.unlink(missing_ok=True)
else:
commit(staged, str(path))
except OSError as e:
for other in also_discard:
if other is not None:
discard_staged_json(other)
raise ClaudeSettingsError(f"Could not {'remove' if staged is None else 'write'} {path}: {e}") from e
def merge_claude_settings(
settings: Mapping[str, JsonValue],
base_url: str,
credential: ClaudeCredential,
default_model: str | None = None,
tier_model: str | None = None,
) -> Mapping[str, JsonValue]:
"""Return a new settings mapping wired to route Claude Code through the proxy.
A StaticToken lands in env.ANTHROPIC_AUTH_TOKEN, an ApiKeyHelper in the top-level apiKeyHelper;
the other credential slots are removed either way, since Claude Code given two credentials may
send the wrong one. ENABLE_TOOL_SEARCH and CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY get their
defaults only when missing. `default_model` is the top-level `model`, the row Claude Code starts
on; `tier_model` is `lite autoroute up`'s knob that points every ANTHROPIC_DEFAULT_*_MODEL at one
group. Apart from those tier keys, exactly OWNED_PATHS are touched.
"""
raw_env: Final = settings.get(ENV_KEY, {})
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
env: Final = {
ENABLE_TOOL_SEARCH_KEY: ENABLE_TOOL_SEARCH_VALUE,
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE,
**{key: value for key, value in base_env.items() if key != ANTHROPIC_API_KEY_KEY},
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
}
return {**settings, ENV_KEY: env, API_KEY_HELPER_KEY: api_key_helper}
current_env: Final = raw_env if isinstance(raw_env, dict) else {}
env: Final = dict( # mutable-ok: JSON document handed to json.dump, which rejects a read-only mapping
chain(
(
(ENABLE_TOOL_SEARCH_KEY, ENABLE_TOOL_SEARCH_VALUE),
(ENABLE_GATEWAY_MODEL_DISCOVERY_KEY, ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE),
),
((key, value) for key, value in current_env.items() if key not in _CREDENTIAL_ENV_KEYS),
((ANTHROPIC_BASE_URL_KEY, base_url.rstrip("/")),),
((ANTHROPIC_AUTH_TOKEN_KEY, credential.token),) if isinstance(credential, StaticToken) else (),
((key, tier_model) for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS if tier_model is not None),
)
)
return dict( # mutable-ok: JSON document handed to json.dump, which rejects a read-only mapping
chain(
((key, value) for key, value in settings.items() if key not in (API_KEY_HELPER_KEY, ENV_KEY)),
((ENV_KEY, env),),
((API_KEY_HELPER_KEY, credential.command),) if isinstance(credential, ApiKeyHelper) else (),
((MODEL_KEY, default_model),) if default_model is not None else (),
)
)
def resolve_api_key_helper(base_url: str, platform: str = sys.platform) -> str:
@ -121,56 +353,273 @@ def resolve_api_key_helper(base_url: str, platform: str = sys.platform) -> str:
return " ".join(quote(token) for token in (lite_path, "--base-url", base_url, "auth", "print-token"))
def write_claude_settings(base_url: str, settings_path: Path, owners: Sequence[SettingsFileOwner]) -> None:
"""Persistently point Claude Code at base_url, preserving every unrelated setting.
def lite_api_key_helper_configured(base_url: str, settings_path: Path) -> bool:
"""Whether settings_path already carries the apiKeyHelper `lite login --config-claude` writes for base_url.
Refuses while any owner holds a backup: each restores its backup when it
stops, which would silently undo this write.
Only an exact match counts: a helper for another proxy, a hand-written one, or
settings that cannot be read leave the caller on the env-token path.
"""
for owner in owners:
if owner.backup_path.exists():
raise ClaudeSettingsError(
f"`{owner.start_command}` is currently managing {settings_path} (backup at "
f"{owner.backup_path}) and will restore it when it stops. "
f"Run `{owner.stop_command}` first, then retry."
)
normalized_base_url: Final = base_url.rstrip("/")
api_key_helper: Final = resolve_api_key_helper(normalized_base_url)
existing: Final = load_json_or_empty(settings_path)
raw_env: Final = existing.get(ENV_KEY)
if raw_env is not None and not isinstance(raw_env, dict):
raise ClaudeSettingsError(
f'{settings_path} has a non-object "{ENV_KEY}" value, which this would discard. '
"Fix or remove it, then retry."
)
merged: Final = merge_claude_settings(existing, normalized_base_url, api_key_helper)
# os.replace() swaps the symlink itself for a regular file, silently detaching a
# settings.json that is symlinked into a dotfiles repo. There is no backup to undo
# that here, unlike `lite up`, so write through to the link's target instead.
target: Final = settings_path.resolve() if settings_path.is_symlink() else settings_path
try:
write_private_json(str(target), merged)
configured_helper: Final = load_json_or_empty(settings_path).get(API_KEY_HELPER_KEY)
return configured_helper == resolve_api_key_helper(base_url.rstrip("/"))
except ClaudeSettingsError:
return False
def _owned(container: Mapping[str, JsonValue], key: str) -> OwnedValue:
return OwnedValue(present=key in container, value=container.get(key))
def _fingerprint(owned: OwnedValue) -> str:
return hashlib.sha256(json.dumps(owned.model_dump(mode="json"), sort_keys=True).encode()).hexdigest()
def _env(settings: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
raw_env: Final = settings.get(ENV_KEY)
return raw_env if isinstance(raw_env, dict) else MappingProxyType({})
def _lookup(settings: Mapping[str, JsonValue], path: str) -> OwnedValue:
section, _, key = path.rpartition(".")
return _owned(_env(settings) if section else settings, key)
def _with_key(container: Mapping[str, JsonValue], key: str, owned: OwnedValue) -> Mapping[str, JsonValue]:
return dict( # mutable-ok: JSON document handed to json.dump, which rejects a read-only mapping
chain(((k, v) for k, v in container.items() if k != key), ((key, owned.value),) if owned.present else ())
)
def _with(settings: Mapping[str, JsonValue], path: str, owned: OwnedValue) -> Mapping[str, JsonValue]:
"""`settings` with the key at `path` set (or removed when `owned` is absent); nothing else changes."""
section, _, key = path.rpartition(".")
if not section:
return _with_key(settings, key, owned)
return _with_key(settings, section, OwnedValue(present=True, value=_with_key(_env(settings), key, owned)))
def _with_all(settings: Mapping[str, JsonValue], updates: Mapping[str, OwnedValue]) -> Mapping[str, JsonValue]:
return reduce(lambda acc, item: _with(acc, *item), updates.items(), settings)
def _ours(settings: Mapping[str, JsonValue], path: str, receipt: ConfigureReceipt) -> bool:
"""Whether the key still holds what a configure wrote (a key no configure ever changed is never ours)."""
return receipt.written.get(path) == _fingerprint(_lookup(settings, path))
def _claim(
path: str,
current: Mapping[str, JsonValue],
merged: Mapping[str, JsonValue],
earlier: ConfigureReceipt | None,
url_now: OwnedValue,
) -> _Claim:
"""What this configure records for one key; see ConfigureReceipt for the rules."""
before, after = _lookup(current, path), _lookup(merged, path)
carried: Final = earlier if earlier is not None and _ours(current, path, earlier) else None
return _Claim(
previous=before if carried is None else carried.previous.get(path, before),
written=_fingerprint(after) if before != after else (None if earlier is None else earlier.written.get(path)),
endpoint=None
if path not in _CREDENTIAL_PATHS
else (url_now if carried is None else carried.endpoints.get(path, url_now)),
)
def _receipt(
current: Mapping[str, JsonValue],
merged: Mapping[str, JsonValue],
earlier: ConfigureReceipt | None,
file_exists: bool,
) -> ConfigureReceipt:
url_now: Final = _lookup(current, _BASE_URL_PATH)
claims: Final = MappingProxyType({path: _claim(path, current, merged, earlier, url_now) for path in OWNED_PATHS})
return ConfigureReceipt(
file_existed=file_exists if earlier is None else earlier.file_existed,
env_present=ENV_KEY in current if earlier is None else earlier.env_present,
env_was_object=isinstance(current.get(ENV_KEY), dict) if earlier is None else earlier.env_was_object,
previous=MappingProxyType({path: claim.previous for path, claim in claims.items()}),
written=MappingProxyType({path: claim.written for path, claim in claims.items() if claim.written is not None}),
endpoints=MappingProxyType(
{path: claim.endpoint for path, claim in claims.items() if claim.endpoint is not None}
),
)
def read_configure_receipt(state_path: Path) -> ConfigureReceipt | None:
if not state_path.exists():
return None
try:
return ConfigureReceipt.model_validate_json(state_path.read_bytes())
except (OSError, ValidationError) as e:
raise ClaudeSettingsError(
f"{state_path} is not a readable `lite configure claude` receipt ({e}). "
"Remove it and edit Claude Code's settings by hand if they still point at the proxy."
) from e
def configure_claude_settings(
base_url: str,
credential: ClaudeCredential,
model: ModelChoice,
settings_path: Path,
state_path: Path,
owners: Sequence[SettingsFileOwner],
commit: Callable[[str, str], None] = commit_staged_json,
) -> None:
"""Persistently route Claude Code through base_url, recording how to undo it.
Both files are staged before either is committed, so a full disk or a read-only directory fails
before anything changes. The two commits are still two renames: a receipt rename that fails
discards the staged settings, and a settings rename that fails after the receipt landed puts the
earlier receipt back (or removes the new one), so the receipt on disk never describes settings
that were not written. `model`: StartOn pins the starting model, UnpinModel lets go of a pin an
earlier configure made (never of the user's own), KeepModel leaves it alone (a re-login).
"""
refuse_while_owned(settings_path, owners)
current: Final = load_json_or_empty(settings_path)
_env_object(current, settings_path)
earlier: Final = read_configure_receipt(state_path)
existing: Final = (
_with(current, MODEL_KEY, earlier.previous[MODEL_KEY])
if isinstance(model, UnpinModel) and earlier is not None and _ours(current, MODEL_KEY, earlier)
else current
)
merged: Final = merge_claude_settings(
existing, base_url, credential, model.model if isinstance(model, StartOn) else None
)
receipt: Final = _receipt(current, merged, earlier, settings_path.exists())
target: Final = _write_target(settings_path)
try:
ensure_private_dir(state_path.parent)
except OSError as e:
raise ClaudeSettingsError(f"Could not write {target}: {e}") from e
raise ClaudeSettingsError(f"Could not write {state_path}: {e}") from e
staged_receipt: Final = _stage(state_path, receipt.model_dump(mode="json"))
try:
staged_settings: Final = _stage(target, merged)
except ClaudeSettingsError:
discard_staged_json(staged_receipt)
raise
_land(state_path, staged_receipt, (staged_settings,), commit)
try:
_land(target, staged_settings, commit=commit)
except ClaudeSettingsError as settings_error:
try:
_land(state_path, None if earlier is None else _stage(state_path, earlier.model_dump(mode="json")))
except ClaudeSettingsError as receipt_error:
raise ClaudeSettingsError(
f"{settings_error} The receipt at {state_path} now describes settings that were not written and "
f"could not be put back either ({receipt_error}); remove it before retrying."
) from settings_error
raise
def _endpoint_text(endpoint: OwnedValue) -> str:
if not endpoint.present:
return f"no {ANTHROPIC_BASE_URL_KEY} (Anthropic's default endpoint)"
return endpoint.value if isinstance(endpoint.value, str) else json.dumps(endpoint.value)
def unconfigure_claude_settings(
settings_path: Path, state_path: Path, owners: Sequence[SettingsFileOwner]
) -> UnconfigureOutcome:
"""Undo `lite configure claude`: put back every key still holding what configure wrote, leave the
rest alone, and withhold a credential the restored file would send to a different server than it
was issued for (the receipt stays, owning only those slots, so a later unconfigure can finish)."""
refuse_while_owned(settings_path, owners)
receipt: Final = read_configure_receipt(state_path)
if receipt is None:
raise ClaudeSettingsError(
f"Claude Code is not configured by `lite configure claude` (no receipt at {state_path}); nothing to undo."
)
current: Final = load_json_or_empty(settings_path)
_env_object(current, settings_path)
ours: Final = tuple(path for path in receipt.written if _ours(current, path, receipt))
kept: Final = tuple(path for path in receipt.written if path not in ours and _lookup(current, path).present)
put_back: Final = _with_all(current, MappingProxyType({path: receipt.previous[path] for path in ours}))
url_after: Final = _lookup(put_back, _BASE_URL_PATH)
withheld: Final = tuple(
WithheldCredential(path, _endpoint_text(receipt.endpoints[path]))
for path in _CREDENTIAL_PATHS
if path in ours and receipt.previous[path].present and receipt.endpoints[path] != url_after
)
absent: Final = OwnedValue(present=False)
trimmed: Final = _with_all(put_back, MappingProxyType({item.key: absent for item in withheld}))
settings: Final = (
trimmed
if _env(trimmed) or receipt.env_was_object
else _with_key(trimmed, ENV_KEY, OwnedValue(present=receipt.env_present, value=None))
)
target: Final = _write_target(settings_path)
file_removed: Final = not settings and not (receipt.file_existed and target.exists())
kept_receipt: Final = ( # mutable-ok: pydantic serializes the update as given and rejects a mappingproxy
receipt.model_copy(update={"written": {item.key: _fingerprint(absent) for item in withheld}})
if withheld
else None
)
staged_settings: Final = None if file_removed else _stage(target, settings)
try:
staged_receipt: Final = (
None if kept_receipt is None else _stage(state_path, kept_receipt.model_dump(mode="json"))
)
except ClaudeSettingsError:
if staged_settings is not None:
discard_staged_json(staged_settings)
raise
_land(target, staged_settings, (staged_receipt,))
_land(state_path, staged_receipt)
return UnconfigureOutcome(
restored=tuple(path for path in ours if _lookup(current, path) != _lookup(settings, path)),
kept=kept,
withheld=withheld,
file_removed=file_removed,
)
__all__ = (
"ANTHROPIC_API_KEY_KEY",
"ANTHROPIC_AUTH_TOKEN_KEY",
"ANTHROPIC_BASE_URL_KEY",
"ANTHROPIC_DEFAULT_MODEL_ENV_KEYS",
"API_KEY_HELPER_KEY",
"AUTOROUTE_BACKUP_PATH",
"BACKUP_PATH",
"CLAUDE_CONFIG_DIR_ENV",
"CLAUDE_SETTINGS_PATH",
"CONFIGURE_STATE_PATH",
"ENABLE_GATEWAY_MODEL_DISCOVERY_KEY",
"ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE",
"ENABLE_TOOL_SEARCH_KEY",
"ENABLE_TOOL_SEARCH_VALUE",
"ENV_KEY",
"MODEL_KEY",
"OWNED_ENV_KEYS",
"OWNED_PATHS",
"OWNED_TOP_LEVEL_KEYS",
"SETTINGS_FILE_OWNERS",
"STARTING_MODEL_ROLE",
"ApiKeyHelper",
"ClaudeCredential",
"ClaudeSettingsError",
"ConfigureReceipt",
"KeepModel",
"ModelChoice",
"OwnedValue",
"SettingsFileOwner",
"StartOn",
"StaticToken",
"UnconfigureOutcome",
"UnpinModel",
"WithheldCredential",
"claude_settings_path",
"configure_claude_settings",
"configure_state_path",
"lite_api_key_helper_configured",
"load_json_or_empty",
"merge_claude_settings",
"read_configure_receipt",
"refuse_while_owned",
"resolve_api_key_helper",
"write_claude_settings",
"settings_file_owners",
"unconfigure_claude_settings",
)

View file

@ -0,0 +1,262 @@
"""`lite configure claude` and `lite unconfigure claude`: persistent Claude Code wiring, undoable."""
import os
import re
import sys
from collections.abc import Callable, Sequence
from pathlib import Path
from typing import Final
import click
from InquirerPy import inquirer
from InquirerPy.base.control import Choice
from .auth import CliContextObj, context_secret_vault, get_stored_api_key
from .claude_settings import (
STARTING_MODEL_ROLE,
ApiKeyHelper,
ClaudeCredential,
ClaudeSettingsError,
ModelChoice,
StartOn,
StaticToken,
UnconfigureOutcome,
UnpinModel,
claude_settings_path,
configure_claude_settings,
configure_state_path,
refuse_while_owned,
resolve_api_key_helper,
settings_file_owners,
unconfigure_claude_settings,
)
from .pi import ListingFailure, PiSyncError, fetch_model_ids
from .up import ensure_fresh_login
_LISTED_MODELS_SHOWN: Final = 20
_CLAUDE_TARGET: Final = "claude"
_TARGETS: Final = ((_CLAUDE_TARGET, "Claude Code (CLI)"),)
_KEEP_DEFAULT_MODEL: Final = "Keep Claude Code's own default"
_CLAUDE_CODE_PICKER_FILTER: Final = re.compile(r"claude|anthropic", re.IGNORECASE)
_MODEL_OPTION_HELP: Final = (
f"Proxy model to set as {STARTING_MODEL_ROLE}. Must be listed on /v1/models for the key; without it, "
"Claude Code keeps its own default and a pin an earlier configure made is let go of. Nothing pins Claude "
"Code's sub-agent or background tiers; `lite autoroute up` is the mode that does."
)
def resolve_credential(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, str]:
"""The credential to write and the key to check the proxy with.
An explicit key (--api-key, `lite --api-key`, LITELLM_PROXY_API_KEY) is long-lived and goes
into settings.json as a static token. Without one, the stored `lite login` credential is used
the way `lite login --config-claude` uses it, through apiKeyHelper, since it expires within a
day and renews in place there; a missing or stale login is refreshed first, as `lite up` does.
"""
ctx_obj: Final[CliContextObj] = ctx.obj
explicit: Final = api_key or (None if ctx_obj.get("api_key_from_token_file") else ctx_obj.get("api_key"))
if explicit:
return StaticToken(explicit), explicit
base_url: Final = ctx_obj["base_url"]
ensure_fresh_login(ctx)
stored: Final = get_stored_api_key(expected_base_url=base_url, vault=context_secret_vault(ctx))
if not stored:
raise ClaudeSettingsError("Login did not produce a usable token.")
return ApiKeyHelper(resolve_api_key_helper(base_url)), stored
def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, tuple[str, ...]]:
"""Every configure path begins the same way: the local ownership check first, so a `lite up`
session is refused before any login prompt or request, then the credential, then the listing."""
settings_path: Final = claude_settings_path(os.environ)
try:
refuse_while_owned(settings_path, settings_file_owners(settings_path))
credential, key = resolve_credential(ctx, api_key)
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
return credential, _listed_models(ctx.obj["base_url"], key)
def _listing_error(base_url: str, error: PiSyncError) -> str:
"""The hint that fits how the listing failed: only an unreachable proxy gets the "is it running" question."""
if error.kind is ListingFailure.REJECTED:
return f"LiteLLM rejected your key (HTTP {error.status}). Run `lite login` to refresh it, or pass a valid --api-key."
if error.kind is ListingFailure.UNREACHABLE:
return f"{error.message} Is the proxy at {base_url} running, and is --base-url (or LITELLM_PROXY_URL) correct?"
if error.kind is ListingFailure.EMPTY:
return f"{error.message} Claude Code would have nothing to run; give the key access to at least one model."
return f"{error.message} The proxy at {base_url} answered, so check that it is a LiteLLM proxy and is healthy."
def _listed_models(base_url: str, key: str) -> tuple[str, ...]:
listed: Final = fetch_model_ids(base_url, key)
if isinstance(listed, PiSyncError):
raise click.ClickException(_listing_error(base_url, listed))
return listed
def _model_choice(model: str | None) -> ModelChoice:
return StartOn(model) if model is not None else UnpinModel()
def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listed: Sequence[str], model: str | None) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
if model is not None and model not in listed:
shown: Final = ", ".join(listed[:_LISTED_MODELS_SHOWN])
more: Final = f", and {len(listed) - _LISTED_MODELS_SHOWN} more" if len(listed) > _LISTED_MODELS_SHOWN else ""
raise click.ClickException(
f"{model!r} is not served by {base_url} for this key. /v1/models lists: {shown}{more}."
)
settings_path: Final = claude_settings_path(os.environ)
try:
configure_claude_settings(
base_url,
credential,
_model_choice(model),
settings_path,
configure_state_path(settings_path),
settings_file_owners(settings_path),
)
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
in_picker: Final = sum(1 for listed_model in listed if _CLAUDE_CODE_PICKER_FILTER.search(listed_model))
click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.")
click.echo(
"Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN."
if isinstance(credential, StaticToken)
else "Credential: your `lite login`, read through apiKeyHelper on every request, so a later login renews it."
)
click.echo(
f"Starting model: {model} ({STARTING_MODEL_ROLE}); switch any time with /model."
if model is not None
else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or "
"pass --model to start on a proxy model."
)
click.echo(
f"/model will list {in_picker} of the proxy's {len(listed)} models (Claude Code shows only ids containing "
"'claude' or 'anthropic')."
)
click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.")
if isinstance(credential, StaticToken) and settings_path.is_symlink():
click.echo(
f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in "
"that file; keep it out of version control.",
err=True,
)
def _pick_targets() -> tuple[str, ...]:
picked: Final = inquirer.checkbox(
message="Which agents should route through LiteLLM?",
choices=[Choice(value, name=label, enabled=True) for value, label in _TARGETS],
validate=lambda chosen: len(chosen) > 0,
invalid_message="Pick at least one.",
).execute()
return tuple(str(value) for value in picked)
def _pick_model(listed: Sequence[str]) -> str | None:
picked: Final = inquirer.fuzzy(
message="Model Claude Code starts on (type to filter; /model switches any time):",
choices=[_KEEP_DEFAULT_MODEL, *listed],
).execute()
return None if picked == _KEEP_DEFAULT_MODEL else str(picked)
def interactive_configure(
ctx: click.Context,
pick_targets: Callable[[], tuple[str, ...]] = _pick_targets,
pick_model: Callable[[Sequence[str]], str | None] = _pick_model,
) -> None:
"""`lite configure` with no agent named: ask which agents to wire and which model to pin."""
targets: Final = pick_targets()
if _CLAUDE_TARGET not in targets:
return
credential, listed = _start(ctx, None)
_apply_claude(ctx, credential, listed, pick_model(listed))
@click.group(name="configure", invoke_without_command=True)
@click.pass_context
def configure_group(ctx: click.Context) -> None:
"""Persistently route a coding agent through your LiteLLM proxy.
With no agent named, asks which agents to wire and which proxy model to pin.
"""
if ctx.invoked_subcommand is not None:
return
if not sys.stdin.isatty():
raise click.ClickException(
"`lite configure` asks questions, so it needs a terminal. Non-interactively, run "
"`lite configure claude --api-key <key> --model <model>`."
)
interactive_configure(ctx)
@click.group(name="unconfigure")
def unconfigure_group() -> None:
"""Undo `lite configure` for a coding agent."""
@configure_group.command(name="claude")
@click.option(
"--api-key",
"api_key",
default=None,
help="Long-lived LiteLLM virtual key written into Claude Code's settings. Defaults to the `lite --api-key` / "
"LITELLM_PROXY_API_KEY value; with neither, your `lite login` credential is used through apiKeyHelper.",
)
@click.option("--model", default=None, help=_MODEL_OPTION_HELP)
@click.pass_context
def configure_claude(ctx: click.Context, api_key: str | None, model: str | None) -> None:
"""Route every Claude Code session through your LiteLLM proxy until `lite unconfigure claude`.
Patches ~/.claude/settings.json in place: the proxy URL, your credential (a virtual key as a
static token, or your `lite login` through apiKeyHelper), and gateway model discovery so
/model lists the proxy's models; --model picks the one Claude Code starts on. Every other
setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back.
Assumes the proxy is already running.
"""
credential, listed = _start(ctx, api_key)
_apply_claude(ctx, credential, listed, model)
@unconfigure_group.command(name="claude")
def unconfigure_claude() -> None:
"""Return Claude Code's settings to what they were before `lite configure claude`.
Also undoes `lite login --config-claude`. Only keys still holding what configure wrote are
put back; anything you changed since is left as it is and named in the output.
"""
settings_path: Final = claude_settings_path(os.environ)
state_path: Final = configure_state_path(settings_path)
try:
outcome: Final = unconfigure_claude_settings(settings_path, state_path, settings_file_owners(settings_path))
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
_report_unconfigure(settings_path, state_path, outcome)
def _report_unconfigure(settings_path: Path, state_path: Path, outcome: UnconfigureOutcome) -> None:
"""Say what unconfigure did, naming only keys whose value it changed."""
if outcome.file_removed:
click.echo(
f"No settings file remains at {settings_path}; it held nothing but `lite configure claude`'s own keys."
)
elif outcome.restored:
click.echo(f"Restored in {settings_path}: {', '.join(outcome.restored)}.")
else:
click.echo(f"Nothing in {settings_path} was still ours to restore.")
if outcome.kept:
click.echo(f"Left as you changed them since: {', '.join(outcome.kept)}.")
if outcome.withheld:
click.echo(
"Left removed, since the file now points at a different server than they were issued for: "
+ "; ".join(f"{item.key} (captured with {item.endpoint})" for item in outcome.withheld)
+ f". They stay in {state_path}: point env.ANTHROPIC_BASE_URL back and run `lite unconfigure claude` "
"again to put them back, or delete that file to drop them."
)
__all__ = ("configure_group", "interactive_configure", "resolve_credential", "unconfigure_group")

View file

@ -10,6 +10,7 @@ import os
import tempfile
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from enum import StrEnum
from pathlib import Path
from types import MappingProxyType
from typing import Final
@ -20,11 +21,28 @@ from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
PI_CONFIG_DIR_ENV: Final = "PI_CODING_AGENT_DIR"
PI_PROVIDER_NAME: Final = "litellm"
LITELLM_PROXY_API_KEY_ENV: Final = "LITELLM_PROXY_API_KEY"
_REJECTED_STATUSES: Final = frozenset((401, 403))
class ListingFailure(StrEnum):
"""Why a proxy could not be listed, decided once where the HTTP outcome is classified.
`unreachable` means no response at all; the other kinds prove the proxy answered, so callers
must not suggest checking whether it is running.
"""
UNREACHABLE = "unreachable"
REJECTED = "rejected"
BAD_BODY = "bad_body"
EMPTY = "empty"
OTHER = "other"
@dataclass(frozen=True, slots=True)
class PiSyncError:
message: str
status: int | None = None
kind: ListingFailure | None = None
@dataclass(frozen=True, slots=True)
@ -65,16 +83,20 @@ def fetch_model_ids(
timeout=10,
)
except requests.RequestException as e:
return PiSyncError(f"Could not list models from the proxy: {e}")
return PiSyncError(f"Could not list models from the proxy: {e}", kind=ListingFailure.UNREACHABLE)
if resp.status_code != 200:
return PiSyncError(f"The proxy returned HTTP {resp.status_code} for /v1/models; cannot build pi's model list.")
return PiSyncError(
f"The proxy returned HTTP {resp.status_code} for /v1/models; cannot list models.",
resp.status_code,
ListingFailure.REJECTED if resp.status_code in _REJECTED_STATUSES else ListingFailure.OTHER,
)
try:
listing: Final = _ModelList.model_validate(resp.json())
except (ValueError, ValidationError) as e:
return PiSyncError(f"Unexpected /v1/models response from the proxy: {e}")
return PiSyncError(f"Unexpected /v1/models response from the proxy: {e}", kind=ListingFailure.BAD_BODY)
ids: Final = tuple(dict.fromkeys(model.id for model in listing.data))
if not ids:
return PiSyncError("The proxy returned no models for your key, so pi would have nothing to run.")
return PiSyncError("The proxy returned no models for your key.", kind=ListingFailure.EMPTY)
return ids
@ -200,6 +222,7 @@ __all__ = (
"LITELLM_PROXY_API_KEY_ENV",
"PI_CONFIG_DIR_ENV",
"PI_PROVIDER_NAME",
"ListingFailure",
"ModelLimits",
"PiSyncError",
"fetch_model_ids",

View file

@ -23,6 +23,7 @@ from .auth import CliContextObj, context_secret_vault, get_stored_api_key, load_
from .claude_settings import (
BACKUP_PATH,
CLAUDE_SETTINGS_PATH,
ApiKeyHelper,
ClaudeSettingsError,
load_json_or_empty,
merge_claude_settings,
@ -123,7 +124,7 @@ def _stored_login_is_pkce(vault: SecretVault) -> bool:
return token_data is not None and token_data.get("refresh_token") is not None
def _ensure_fresh_login(ctx: click.Context) -> None:
def ensure_fresh_login(ctx: click.Context) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"].rstrip("/")
vault: Final = context_secret_vault(ctx)
@ -141,7 +142,7 @@ def _ensure_fresh_login(ctx: click.Context) -> None:
click.echo("No fresh LiteLLM login found for this proxy; starting login...")
ctx.invoke(login, pkce=pkce)
if not _usable_login(get_stored_api_key(expected_base_url=base_url, vault=vault), vault):
raise UpError("Login did not produce a usable token; cannot start `lite up`.")
raise UpError("Login did not produce a usable token.")
def _restore_and_report() -> None:
@ -169,7 +170,7 @@ def up(ctx: click.Context) -> None:
base_url: Final = ctx.obj["base_url"]
try:
_ensure_fresh_login(ctx)
ensure_fresh_login(ctx)
api_key: Final = resolve_api_key(ctx)
verify_proxy_key(base_url, api_key)
@ -190,7 +191,7 @@ def up(ctx: click.Context) -> None:
)
CLAUDE_SETTINGS_PATH.parent.mkdir(exist_ok=True)
merged: Final = merge_claude_settings(original_settings, base_url, api_key_helper)
merged: Final = merge_claude_settings(original_settings, base_url, ApiKeyHelper(api_key_helper))
with open(CLAUDE_SETTINGS_PATH, "w") as f:
json.dump(merged, f, indent=2)
except (AgentRunError, ClaudeSettingsError) as e:

View file

@ -13,6 +13,7 @@ from .commands.auth import auth_group, context_secret_vault, get_stored_api_key,
from .commands.autoroute.commands import autoroute_group
from .commands.chat import chat
from .commands.config import config_commands, get_config_value, hidden_command_names
from .commands.configure import configure_group, unconfigure_group
from .commands.credentials import credentials
from .commands.debug import debug
from .commands.encryption import encryption
@ -162,6 +163,9 @@ cli.add_command(model_groups)
# Add the autoroute command group (QA auto-routing against your real proxy)
cli.add_command(autoroute_group, name="autoroute")
cli.add_command(config_commands)
# Add configure/unconfigure (persistently wire a coding agent to the proxy with a virtual key)
cli.add_command(configure_group)
cli.add_command(unconfigure_group)
if __name__ == "__main__":

View file

@ -3300,9 +3300,10 @@ class ProxyBaseLLMRequestProcessing:
has completed.
Guardrails routed through unified_guardrail are skipped, since they already ran
via its streaming iterator. Guardrails that override
async_post_call_success_hook directly run here, including those that implement
apply_guardrail but keep their native lifecycle hooks.
via its streaming iterator, and so are guardrails a post_call policy pipeline
manages, since the pipeline ran them against the buffered stream. Guardrails
that override async_post_call_success_hook directly run here, including those
that implement apply_guardrail but keep their native lifecycle hooks.
This is audit-only content has already been delivered to the client.
@ -3312,12 +3313,18 @@ class ProxyBaseLLMRequestProcessing:
_response = assembled_response
try:
from litellm.proxy.proxy_server import llm_router as _global_llm_router
from litellm.proxy.utils import _check_and_merge_model_level_guardrails
from litellm.proxy.utils import (
_check_and_merge_model_level_guardrails,
stream_gated_guardrail_names,
)
guardrail_data = _check_and_merge_model_level_guardrails(data=captured_data, llm_router=_global_llm_router)
stream_gated: Final = stream_gated_guardrail_names(captured_data, captured_user_api_key_dict)
for cb in litellm.callbacks:
if not isinstance(cb, CustomGuardrail):
continue
if cb.guardrail_name in stream_gated:
continue
if not cb.should_run_guardrail(
data=guardrail_data,
event_type=GuardrailEventHooks.post_call,

View file

@ -39,7 +39,7 @@ def _is_form_content_type(content_type: str) -> bool:
return _normalize_media_type(content_type) in _FORM_CONTENT_TYPES
def _is_json_content_type(content_type: str) -> bool:
def is_json_content_type(content_type: str) -> bool:
"""True iff the body should be parsed as JSON."""
return _normalize_media_type(content_type) == "application/json"
@ -406,7 +406,7 @@ async def get_request_body(request: Request) -> dict[str, Any]:
"""
if request.method == "POST":
content_type: Final = request.headers.get("content-type", "")
if _is_json_content_type(content_type):
if is_json_content_type(content_type):
return await _read_request_body(request)
elif _is_form_content_type(content_type):
return await get_form_data(request)

View file

@ -80,17 +80,19 @@ def policy_from_litellm_params(litellm_params: Mapping[str, object]) -> AutoRout
def policy_for_model(
llm_router: "Router | None",
model_alias: str,
team_id: str | None,
request_kwargs: Mapping[str, object],
request_tags: Sequence[str],
) -> AutoRouterCompressionPolicy | None:
"""The compression policy of the auto router marker `model_alias` resolves to.
"""The compression policy of the auto router marker `model_alias` resolves to for this caller.
Pre-call arming and the routing hook both resolve through here, so an alias with
several tag-scoped markers cannot suppress under one and then route under another.
Pre-call arming and the routing hook both resolve through here, and here resolves through the
router's own request-scoped deployment lookup, so an alias with several tag-scoped markers
cannot suppress under one and then route under another, and a team router reached by its
public name carries its policy for every principal that can reach it.
"""
if llm_router is None:
return None
deployments: Final = llm_router.get_model_list(model_name=model_alias, team_id=team_id) or ()
deployments: Final = llm_router.deployments_for_request(model_alias, request_kwargs)
markers: Final = tuple(
litellm_params
for deployment in deployments
@ -108,17 +110,6 @@ def policy_for_model(
return next((policy for policy in candidates if policy is not None), None)
def team_id_from_request(request_kwargs: Mapping[str, object]) -> str | None:
"""The caller's team id, from whichever metadata bucket this surface writes to."""
for meta_key in ("metadata", "litellm_metadata"):
meta = request_kwargs.get(meta_key)
if isinstance(meta, Mapping):
team_id = meta.get("user_api_key_team_id")
if isinstance(team_id, str):
return team_id
return None
def _compression_guardrail_classes() -> tuple[type, ...]:
"""The registered guardrail classes whose provider compresses prompts."""
from litellm.proxy.guardrails.guardrail_registry import guardrail_class_registry
@ -172,7 +163,7 @@ async def arm_pre_call(
policy: Final = policy_for_model(
llm_router=llm_router,
model_alias=model_alias,
team_id=team_id_from_request(data),
request_kwargs=data,
request_tags=_get_tags_from_request_kwargs(data),
)
if policy is None:

View file

@ -44,7 +44,7 @@ from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicM
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_scan_only_tool_results_for_guardrail,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -917,7 +917,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
source,
)
return BedrockGuardrailResponse()
credentials, aws_region_name = self._load_credentials(bearer_token=bedrock_bearer_token(api_key))
credentials, aws_region_name = await run_aws_signing(
self._load_credentials, bearer_token=bedrock_bearer_token(api_key)
)
allow_chunking: Final = not self._content_uses_contextual_grounding(content)
completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator
@ -1178,7 +1180,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
**base_request_data,
"content": content,
} # mutable-ok: outbound JSON request body
prepared_request: Final = self._prepare_request(
prepared_request: Final = await run_aws_signing(
self._prepare_request,
credentials=credentials,
data=bedrock_request_data,
optional_params=self.optional_params,
@ -1875,10 +1878,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return BedrockGuardrailResponse()
api_key: Final[str | None] = request_data.get("api_key") if request_data else None
credentials, aws_region_name = self._load_credentials(bearer_token=bedrock_bearer_token(api_key))
credentials, aws_region_name = await run_aws_signing(
self._load_credentials, bearer_token=bedrock_bearer_token(api_key)
)
body: Final[dict[str, object]] = {"messages": checks_messages, "checks": self.checks}
prepared_request: Final = self._prepare_request(
prepared_request: Final = await run_aws_signing(
self._prepare_request,
credentials=credentials,
data=body,
optional_params=self.optional_params,

View file

@ -31,6 +31,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
ESTIMATED_OUTPUT_TOKENS_FIELD,
@ -3307,7 +3308,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
min_configured_tpm_limit=min_configured_otpm_limit,
call_type=call_type,
)
raw_estimated_input_tokens: Final = self._estimate_precise_input_tokens(
raw_estimated_input_tokens: Final = await offload_token_count(self._estimate_precise_input_tokens)(
data=data, model=requested_model, call_type=call_type
)
estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1)

View file

@ -1760,7 +1760,9 @@ class LiteLLMProxyRequestSetup:
callback_vars_dict.pop("success_callback", None)
callback_vars_dict.pop("failure_callback", None)
callback_vars_dict = {
key: (litellm.utils.get_secret(value, default_value=value) or value if isinstance(value, str) else value)
key: (
litellm.utils.get_secret(value, default_value=value) or value if isinstance(value, str) else str(value)
)
for key, value in callback_vars_dict.items()
}

View file

@ -4559,6 +4559,23 @@ async def delete_verification_tokens(
litellm_changed_by=litellm_changed_by,
)
# Snapshot before the delete: the FK cascade drops the mapping rows, but their
# cached jwt_key_mapping entries still resolve to the now-dead token (LIT-5380).
jwt_mapping_cache_keys: Final[tuple[str, ...]] = tuple(
cache_key
for keys_for_token in await asyncio.gather(
*(
get_jwt_key_mapping_cache_keys_for_token(
hashed_token=key.token,
prisma_client=prisma_client,
)
for key in authorized_keys
if key.token is not None
)
)
for cache_key in keys_for_token
)
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
deleted_tokens = await prisma_client.delete_data(tokens=tokens)
if deleted_tokens is not None and len(deleted_tokens) != len(tokens):
@ -4571,6 +4588,8 @@ async def delete_verification_tokens(
if len(deleted_tokens) != len(tokens):
failed_tokens = [token for token in tokens if token not in deleted_tokens]
await evict_and_broadcast(cache_keys=jwt_mapping_cache_keys, user_api_key_cache=user_api_key_cache)
else:
raise Exception("DB not connected. prisma_client is None")
except Exception as e:

View file

@ -15,6 +15,7 @@ import os
import re
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
from dataclasses import dataclass
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
@ -33,6 +34,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -52,6 +54,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_set_request_parsed_body,
get_form_data,
get_request_body,
is_json_content_type,
)
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
@ -77,6 +80,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
from litellm.types.router import LiteLLMParamsTypedDict
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
@ -119,6 +123,24 @@ def is_passthrough_request_using_router_model(request_body: dict, llm_router: li
return False
class RelayRejection(TypedDict):
error: ReadOnly[str]
def _deployment_model_name(litellm_params: LiteLLMParamsTypedDict) -> str:
model: Final = litellm_params.get("model", "")
try:
return get_llm_provider(model=model, custom_llm_provider=litellm_params.get("custom_llm_provider"))[0]
except litellm.BadRequestError:
return model
def _models_served_by_group(llm_router: litellm.Router, model_group: str) -> frozenset[str]:
return frozenset(
_deployment_model_name(row["litellm_params"]) for row in llm_router.get_model_list(model_name=model_group) or ()
)
def is_passthrough_request_streaming(request_body: object) -> bool:
"""
Returns True if the request is streaming.
@ -411,7 +433,7 @@ async def vllm_proxy_route(
content=None,
data=None,
files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None),
json=(request_body if is_json_content_type(request.headers.get("content-type", "")) else None),
params=None,
headers=None,
cookies=None,
@ -1099,13 +1121,6 @@ async def bedrock_proxy_route(
"""
create_request_copy(request)
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
aws_region_name: Final = get_secret_str(secret_name="AWS_REGION_NAME")
if not _is_bedrock_agent_runtime_route(endpoint=endpoint):
return await bedrock_llm_proxy_route(
@ -1136,20 +1151,24 @@ async def bedrock_proxy_route(
)
# Add or update query parameters
from litellm.llms.bedrock.base_aws_llm import run_aws_signing, sign_aws_json_post
from litellm.llms.bedrock.chat import BedrockConverseLLM
bedrock_llm: Final = BedrockConverseLLM()
credentials: Final[Credentials] = bedrock_llm.get_credentials()
sigv4: Final = SigV4Auth(credentials, "bedrock", aws_region_name)
headers: Final = {"Content-Type": "application/json"}
# Assuming the body contains JSON data, parse it
try:
data: Final = await _json_request_body(request)
except Exception as e:
raise HTTPException(status_code=400, detail={"error": e})
_request: Final = AWSRequest(method="POST", url=str(updated_url), data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
prepped: Final = await run_aws_signing(
sign_aws_json_post,
get_credentials=bedrock_llm.get_credentials,
service_name="bedrock",
aws_region_name=aws_region_name,
url=str(updated_url),
body=json.dumps(data),
headers=MappingProxyType({"Content-Type": "application/json"}),
)
## check for streaming
is_streaming_request = False
@ -1207,13 +1226,6 @@ async def comprehend_medical_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)
"""
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call comprehendmedical. Run 'pip install boto3'.")
from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import (
COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS,
)
@ -1244,20 +1256,23 @@ async def comprehend_medical_proxy_route(
if "stream" in data:
raise HTTPException(status_code=400, detail="'stream' is not a Comprehend Medical request member")
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing, sign_aws_json_post
credentials: Final[Credentials] = BaseAWSLLM().get_credentials(aws_region_name=aws_region_name)
sigv4: Final = SigV4Auth(credentials, "comprehendmedical", aws_region_name)
headers: Final = MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
)
target_url: Final = f"https://comprehendmedical.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/"
_request: Final = AWSRequest(method="POST", url=target_url, data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
prepped: Final = await run_aws_signing(
sign_aws_json_post,
get_credentials=partial(BaseAWSLLM().get_credentials, aws_region_name=aws_region_name),
service_name="comprehendmedical",
aws_region_name=aws_region_name,
url=target_url,
body=json.dumps(data),
headers=MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
),
)
endpoint_func: Final = create_pass_through_route(
endpoint=operation,
@ -1505,6 +1520,14 @@ async def _relay_upstream_bytes(upstream: AsyncGenerator[bytes, bytes]) -> Async
await upstream.aclose()
async def _relay_upstream_response(upstream: httpx.Response) -> Response:
return Response(
content=await upstream.aread(),
status_code=upstream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(headers=upstream.headers, custom_headers=None),
)
async def _relay_azure_router_model(
llm_router: litellm.Router,
model: str,
@ -1514,30 +1537,37 @@ async def _relay_azure_router_model(
is_streaming_request: bool,
user_api_key_dict: UserAPIKeyAuth,
) -> Response:
result: Final = await llm_router.allm_passthrough_route(
model=model,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=is_streaming_request,
content=None,
data=None,
files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None),
params=None,
headers=None,
cookies=None,
litellm_metadata=get_passthrough_router_request_metadata(user_api_key_dict),
foreign_deployment: Final = foreign_azure_deployment(
endpoint, model, lambda: _models_served_by_group(llm_router, model)
)
if foreign_deployment is not None:
rejection: Final[RelayRejection] = {
"error": f"deployment '{foreign_deployment}' in the path is not served by model group '{model}'; "
"put the model group name in the deployments segment"
}
raise HTTPException(status_code=400, detail=rejection)
try:
result: Final = await llm_router.allm_passthrough_route(
model=model,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=is_streaming_request,
content=None,
data=None,
files=None,
json=(request_body if is_json_content_type(request.headers.get("content-type", "")) else None),
params=None,
headers=None,
cookies=None,
litellm_metadata=get_passthrough_router_request_metadata(user_api_key_dict),
)
except httpx.HTTPStatusError as upstream_error:
return await _relay_upstream_response(upstream_error.response)
if not is_streaming_request:
upstream: Final = cast(httpx.Response, result)
return Response(
content=await upstream.aread(),
status_code=upstream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(headers=upstream.headers, custom_headers=None),
)
return await _relay_upstream_response(cast(httpx.Response, result))
if inspect.isasyncgen(result):
sse_headers: Final = {"content-type": "text/event-stream"}

View file

@ -4,6 +4,7 @@ OpenAI Passthrough Logging Handler
Handles cost tracking and logging for OpenAI passthrough endpoints, specifically /chat/completions.
"""
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Final
from urllib.parse import urlparse
@ -16,6 +17,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
)
from litellm.litellm_core_utils.token_counter import high_detail_image_token_upper_bound
from litellm.llms.openai.openai import OpenAIConfig
from litellm.llms.openai.openai import OpenAIConfig as OpenAIConfigType
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
@ -96,6 +98,47 @@ def _is_openai_compatible_url(url_route: str | None) -> bool:
return False
def _is_remote_high_detail_image(part: object) -> bool:
if not isinstance(part, Mapping) or part.get("type") != "image_url":
return False
image_url: Final = part.get("image_url")
if not isinstance(image_url, Mapping):
return False
url: Final = image_url.get("url")
return (
isinstance(url, str) and url.lower().startswith(("http://", "https://")) and image_url.get("detail") == "high"
)
def _content_parts(message: Mapping[str, object]) -> Sequence[object]:
content: Final = message.get("content")
return content if isinstance(content, list) else ()
def _without_remote_high_detail_images(message: Mapping[str, object]) -> Mapping[str, object]:
if not isinstance(message.get("content"), list):
return message
kept_parts: Final = [ # mutable-ok: token_counter reads message content only when it is a list
part for part in _content_parts(message) if not _is_remote_high_detail_image(part)
]
return {**message, "content": kept_parts} # mutable-ok: token_counter rejects any message that is not a dict
def count_relayed_prompt_tokens(model: str, messages: Sequence[Mapping[str, object]] | None) -> int:
if messages is None:
return 0
remote_high_detail_images: Final = sum(
1 for message in messages for part in _content_parts(message) if _is_remote_high_detail_image(part)
)
local_messages: Final = [ # mutable-ok: token_counter takes a list of messages
_without_remote_high_detail_images(message) for message in messages
]
return (
litellm.token_counter(model=model, messages=local_messages)
+ high_detail_image_token_upper_bound() * remote_high_detail_images
)
class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""
OpenAI-specific passthrough logging handler that provides cost tracking for /chat/completions endpoints.
@ -512,9 +555,10 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
def _build_complete_streaming_response(
self,
all_chunks: list[str],
all_chunks: Sequence[str],
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
messages: Sequence[Mapping[str, object]] | None = None,
) -> ModelResponse | TextCompletionResponse | None:
"""
Builds complete response from raw chunks for OpenAI streaming responses.
@ -558,7 +602,11 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
return None
# Build complete response from chunks
complete_streaming_response: Final = litellm.stream_chunk_builder(chunks=all_openai_chunks)
complete_streaming_response: Final = litellm.stream_chunk_builder(
chunks=all_openai_chunks,
messages=messages,
count_prompt_tokens=lambda: count_relayed_prompt_tokens(model, messages),
)
return complete_streaming_response

View file

@ -70,6 +70,10 @@ def _text_snapshot(texts: Sequence[str] | None) -> tuple[str, ...] | None:
return None if texts is None else tuple(texts)
def _scanned_texts(texts: Sequence[str] | None) -> tuple[str, ...]:
return tuple(texts or ())
def _tool_call_shapes(tool_calls: Sequence[object] | None) -> tuple[tuple[object, object], ...] | None:
return None if tool_calls is None else tuple(_tool_call_shape(tool_call) for tool_call in tool_calls)
@ -78,6 +82,10 @@ def _rewrote(sent: tuple[object, ...] | None, returned: tuple[object, ...] | Non
return sent is not None and returned is not None and returned != sent
def _changed_count(sent: tuple[object, ...] | None, returned: tuple[object, ...] | None) -> bool:
return sent is not None and returned is not None and len(returned) != len(sent)
_GuardrailMethodT = TypeVar("_GuardrailMethodT", bound=Callable[..., object])
@ -89,10 +97,11 @@ def _logged_by_inner_guardrail(method: _GuardrailMethodT) -> _GuardrailMethodT:
class _StreamRewriteObserver(CustomGuardrail):
"""Stand-in handed to the endpoint translation in place of a streaming pipeline step's
guardrail. It records whether the guardrail returned different output than it was given,
which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text
rewrites are deliverable on translations that write them back across the buffered chunks
(``delivers_ended_stream_text_rewrites``); tool-call rewrites and text rewrites on any
other translation are discarded by the executor, which releases the original chunks.
which for guardrails like Bedrock's ANONYMIZED action is only known at runtime. Text and
tool-call rewrites are deliverable on translations that write them back across the
buffered chunks (``delivers_ended_stream_rewrites``); rewrites on any other translation,
and a rewrite that drops or adds a tool call on any translation, are discarded by the
executor, which releases the original chunks.
The inner guardrail's ``apply_guardrail`` already records the guardrail information
and span, so the observer's stays out of ``log_guardrail_information``."""
@ -101,6 +110,7 @@ class _StreamRewriteObserver(CustomGuardrail):
self.inner: Final = inner
self.rewrote_texts = False
self.rewrote_tool_calls = False
self.changed_tool_call_count = False
def structured_messages_cover_full_request(self) -> bool:
return self.inner.structured_messages_cover_full_request()
@ -118,13 +128,103 @@ class _StreamRewriteObserver(CustomGuardrail):
outputs: Final = await self.inner.apply_guardrail(
inputs=inputs, request_data=request_data, input_type=input_type, logging_obj=logging_obj
)
returned_tool_shapes: Final = _tool_call_shapes(outputs.get("tool_calls"))
self.rewrote_texts = self.rewrote_texts or _rewrote(sent_texts, _text_snapshot(outputs.get("texts")))
self.rewrote_tool_calls = self.rewrote_tool_calls or _rewrote(
sent_tool_shapes, _tool_call_shapes(outputs.get("tool_calls"))
self.rewrote_tool_calls = self.rewrote_tool_calls or _rewrote(sent_tool_shapes, returned_tool_shapes)
self.changed_tool_call_count = self.changed_tool_call_count or _changed_count(
sent_tool_shapes, returned_tool_shapes
)
return outputs
class _ScannedTextRecorder(CustomGuardrail):
def __init__(self, guardrail_name: str) -> None:
super().__init__(guardrail_name=guardrail_name)
self.inputs: GenericGuardrailAPIInputs | None = None
@_logged_by_inner_guardrail
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict, # mutable-ok: matches CustomGuardrail.apply_guardrail
input_type: Literal["request", "response"],
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
self.inputs = inputs
return inputs
class _LegacyHookStreamAdapter(CustomGuardrail):
"""Runs a guardrail that only implements the legacy post-call hook (no unified
``apply_guardrail``, or ``use_native_lifecycle_hooks``) as a streaming pipeline step. The
endpoint translation hands it the texts it scanned plus the assembled response under
``request_data["response"]``; the hook gets that response in the shape its route gives
non-streaming hooks, an exception it raises ends the stream through the executor's
fail/error classification, and the response it hands back, or the one it changed in place
and returned ``None`` for, is re-scanned by the same translation so its texts reach the
client through the translation's ended-stream write-back. A
replacement whose scanned texts do not line up with the originals, or whose tool calls
differ from them, is undeliverable, so the executor releases the original chunks. A stream
that carried no text to scan, such as a tool-only Anthropic message, stays deliverable as
long as the hook left the tool calls alone."""
def __init__(
self,
inner: CustomGuardrail,
endpoint_translation: "BaseTranslation",
user_api_key_dict: "UserAPIKeyAuth",
) -> None:
super().__init__(guardrail_name=inner.guardrail_name)
self.inner: Final = inner
self.endpoint_translation: Final = endpoint_translation
self.user_api_key_dict: Final = user_api_key_dict
def structured_messages_cover_full_request(self) -> bool:
return self.inner.structured_messages_cover_full_request()
@_logged_by_inner_guardrail
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict, # mutable-ok: matches CustomGuardrail.apply_guardrail
input_type: Literal["request", "response"],
logging_obj: "LiteLLMLoggingObj | None" = None,
) -> GenericGuardrailAPIInputs:
hooked: Final = self.endpoint_translation.post_call_hook_response(request_data.get("response"))
replacement: Final = await self.inner.async_post_call_success_hook(
data=request_data,
user_api_key_dict=self.user_api_key_dict,
response=hooked,
)
rewrite: Final = hooked if replacement is None else replacement
if rewrite is None:
return inputs
rescanned: Final = await self._rescan(rewrite, logging_obj)
if rescanned is None:
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
rewritten: Final = rescanned.get("texts")
if len(_scanned_texts(rewritten)) != len(_scanned_texts(inputs.get("texts"))):
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
if _tool_call_shapes(rescanned.get("tool_calls")) != _tool_call_shapes(inputs.get("tool_calls")):
raise UndeliverableStreamRewrite(self.guardrail_name or "unknown")
if not rewritten:
return inputs
rewritten_inputs: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": rewritten}
return rewritten_inputs
async def _rescan(
self, response: object, logging_obj: "LiteLLMLoggingObj | None"
) -> GenericGuardrailAPIInputs | None:
recorder: Final = _ScannedTextRecorder(self.guardrail_name or "unknown")
await self.endpoint_translation.process_output_response(
response=response,
guardrail_to_apply=recorder,
litellm_logging_obj=logging_obj,
user_api_key_dict=self.user_api_key_dict,
)
return recorder.inputs
def _prepare_hook_input(
step: PipelineStep,
callback: CustomGuardrail,
@ -292,18 +392,29 @@ class PipelineExecutor:
endpoint_translation: "BaseTranslation",
streaming_chunks: list[object], # mutable-ok: shared buffered-stream chunks the translation rewrites in place
hook_input: dict[str, object], # mutable-ok: same request-payload shape as data
user_api_key_dict: "UserAPIKeyAuth | None",
user_api_key_dict: "UserAPIKeyAuth",
litellm_logging_obj: "LiteLLMLoggingObj | None",
) -> None:
"""Run one streaming post_call step through the endpoint translation, delivering
text rewrites on translations that support ended-stream write-back. A rewrite that
cannot reach the client yet (a tool-call rewrite, a text rewrite on a translation
without write-back, or one the translation refused with
``UndeliverableStreamRewrite``) is discarded: the buffered chunks go back to the
originals and the step passes, so the client gets the stream the merge base sent."""
observer: Final = _StreamRewriteObserver(callback)
deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_text_rewrites
text and tool-call rewrites on translations that support ended-stream write-back. A
guardrail without the unified interface runs its legacy post-call hook against the
assembled response through ``_LegacyHookStreamAdapter``. A rewrite that cannot reach the
client yet (one on a translation without write-back, one that drops or adds a tool call,
or one the translation or adapter refused with ``UndeliverableStreamRewrite``) is
discarded: the buffered chunks go back to the originals and the step passes, so the
client gets the stream the merge base sent, and the guardrail stays out of the
applied-guardrails header since its output never reached the client. The response an
earlier step's translation stored under ``request_data["response"]`` is dropped first,
so this step's hook sees the stream as the steps before it left it."""
scanner: Final = (
callback
if PipelineExecutor.supports_unified_execution(callback)
else _LegacyHookStreamAdapter(callback, endpoint_translation, user_api_key_dict)
)
observer: Final = _StreamRewriteObserver(scanner)
deliver_rewrites: Final = type(endpoint_translation).delivers_ended_stream_rewrites
originals: Final = copy.deepcopy(streaming_chunks)
hook_input.pop("response", None) # rebind-ok: an earlier step's stored response goes so this step's is stored
try:
if deliver_rewrites:
await endpoint_translation.process_output_streaming_response(
@ -324,9 +435,12 @@ class PipelineExecutor:
)
except UndeliverableStreamRewrite:
_release_original_chunks(step.guardrail, streaming_chunks, originals)
else:
if observer.rewrote_tool_calls or (observer.rewrote_texts and not deliver_rewrites):
_release_original_chunks(step.guardrail, streaming_chunks, originals)
return
if observer.changed_tool_call_count or (
not deliver_rewrites and (observer.rewrote_texts or observer.rewrote_tool_calls)
):
_release_original_chunks(step.guardrail, streaming_chunks, originals)
return
if not callback.records_own_guardrail_information:
add_guardrail_to_applied_guardrails_header(request_data=hook_input, guardrail_name=step.guardrail)
@ -386,11 +500,11 @@ class PipelineExecutor:
if isinstance(response, dict):
callback.mark_pre_call_hook_ran(response)
elif mode == "post_call" and streaming_chunks is not None:
if not use_unified or endpoint_translation is None:
if endpoint_translation is None:
return (
"error",
None,
f"Guardrail '{step.guardrail}' does not support streaming pipeline execution",
f"Guardrail '{step.guardrail}' cannot run on a stream without an endpoint translation",
None,
)
await PipelineExecutor._run_streaming_step(
@ -446,10 +560,22 @@ class PipelineExecutor:
@staticmethod
def supports_unified_execution(callback: CustomGuardrail) -> bool:
"""Whether this guardrail runs through the unified apply_guardrail path,
the interface streaming pipeline execution requires."""
"""Whether this guardrail runs through the unified apply_guardrail path."""
return "apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks
@staticmethod
def supports_streaming_execution(callback: CustomGuardrail) -> bool:
"""Whether a streaming pipeline step can run this guardrail against the buffered
stream: through the unified path, or through its post-call hook on the assembled
response when that hook is its only streaming path. A guardrail with its own
streaming iterator hook, or with neither hook, keeps running on its own."""
callback_type: Final = type(callback)
return PipelineExecutor.supports_unified_execution(callback) or (
callback_type.async_post_call_success_hook is not CustomLogger.async_post_call_success_hook
and callback_type.async_post_call_streaming_iterator_hook
is CustomLogger.async_post_call_streaming_iterator_hook
)
@staticmethod
def find_guardrail_callback(guardrail_name: str) -> CustomGuardrail | None:
"""Look up an initialized guardrail callback by name from litellm.callbacks."""

View file

@ -63,11 +63,13 @@ from litellm.constants import (
LITELLM_UI_SESSION_DURATION,
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
)
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.proxy._types import (
UI_TEAM_ID,
CallbackDelete,
@ -272,7 +274,6 @@ from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.litellm_core_utils.agentic_loop_settings import (
validated_max_agentic_loops,
)
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.audio_utils.utils import resolve_speech_media_type
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
@ -12816,7 +12817,9 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False)
CustomHuggingfaceTokenizer | None,
model_info.get("custom_tokenizer", None),
)
_tokenizer_used: Final = litellm.utils._select_tokenizer(model=model_to_use, custom_tokenizer=custom_tokenizer)
_tokenizer_used: Final = await asyncify(litellm.utils._select_tokenizer)(
model=model_to_use, custom_tokenizer=custom_tokenizer
)
tokenizer_used: Final = str(_tokenizer_used["type"])
system_message: Final = _system_message(system)
@ -12829,7 +12832,7 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False)
counted_tools: Final = cast( # cast-ok: raw OpenAI or Anthropic tool dicts, both of which token_counter formats
list[ChatCompletionToolParam] | None, tools if counted_messages is not None else None
)
total_tokens: Final = await asyncify(litellm.token_counter)(
total_tokens: Final = await offload_token_count(litellm.token_counter)(
model=model_to_use,
text=prompt,
messages=counted_messages,

View file

@ -15,6 +15,7 @@ from typing import TYPE_CHECKING, Final, TypeAlias
from fastapi import Request, Response
from fastapi.responses import StreamingResponse
from starlette.types import Message
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
@ -74,6 +75,20 @@ class _StreamEventParser:
parse: Callable[[str], _StreamEvent] = staticmethod(json.loads)
async def _never_receive() -> Message:
await asyncio.Event().wait()
raise AssertionError("unreachable")
def detach_request_from_client(request: Request) -> Request:
"""Same scope (headers, parsed body, auth) but a receive() that never yields http.disconnect.
The polling client closes its connection right after getting the polling id, so the
upstream call must not be cancelled by the client-disconnect guards.
"""
return Request(request.scope, _never_receive)
async def background_streaming_task(
polling_id: str,
data: dict[str, object],
@ -123,7 +138,7 @@ async def background_streaming_task(
# Pre-call checks (rate limits, guardrails, budget) were already run
# before polling ID creation, so skip them here to avoid double-counting.
response: Final[StreamingResponse] = await processor.base_process_llm_request(
request=request,
request=detach_request_from_client(request),
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="aresponses",

View file

@ -492,7 +492,7 @@ model LiteLLM_JWTKeyMapping {
updated_at DateTime @default(now()) @updatedAt
updated_by String?
litellm_verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token])
litellm_verification_token LiteLLM_VerificationToken @relation(fields: [token], references: [token], onDelete: Cascade)
@@unique([jwt_claim_name, jwt_claim_value])
@@index([jwt_claim_name, jwt_claim_value, is_active])

View file

@ -101,6 +101,7 @@ from litellm.litellm_core_utils.core_helpers import (
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.llms import load_guardrail_translation_mappings
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import (
@ -197,6 +198,7 @@ if TYPE_CHECKING:
from prisma.types import HttpConfig
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.models.team import LiteLLM_TeamTableCachedObj
from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction
from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction
@ -459,7 +461,7 @@ def _pipeline_step_guardrail_names(pipelines: Sequence[tuple[str, "GuardrailPipe
return frozenset(step.guardrail for _policy_name, pipeline in pipelines for step in pipeline.steps)
def _pipeline_managed_guardrail_names(
def pipeline_managed_guardrail_names(
data: Mapping[str, object], mode: Literal["pre_call", "post_call"]
) -> frozenset[str]:
return _pipeline_step_guardrail_names(
@ -522,9 +524,17 @@ def _merge_pipeline_metadata_writes(
_merge_pipeline_metadata_bucket(data, bucket_key, modified_data.get(bucket_key))
def _pipeline_step_supports_unified_streaming(guardrail_name: str) -> bool:
def _pipeline_step_supports_streaming(guardrail_name: str, translation: "BaseTranslation | None") -> bool:
callback: Final = PipelineExecutor.find_guardrail_callback(guardrail_name)
return callback is not None and PipelineExecutor.supports_unified_execution(callback)
if callback is None:
return False
if PipelineExecutor.supports_unified_execution(callback):
return True
return (
translation is not None
and type(translation).assembles_streamed_response
and PipelineExecutor.supports_streaming_execution(callback)
)
def _post_call_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "GuardrailPipeline"], ...]:
@ -581,7 +591,7 @@ def _withdraw_deferred_claims(
outside_by_policy: Final = MappingProxyType(
{policy_name: _guardrails_outside_pipeline(policy_name, pipeline) for policy_name, pipeline in deferred}
)
running_elsewhere: Final = _pipeline_managed_guardrail_names(data, "pre_call").union(
running_elsewhere: Final = pipeline_managed_guardrail_names(data, "pre_call").union(
_guardrails_run_standalone_pre_call(data), *outside_by_policy.values()
)
withdrawn_policies: Final = frozenset(name for name, outside in outside_by_policy.items() if not outside)
@ -656,37 +666,51 @@ def _body_selected_deferrals(
return tuple(policy_name for policy_name, _pipeline in deferred if policy_name not in attributed)
def _pipeline_is_streamable(policy_name: str, pipeline: "GuardrailPipeline") -> bool:
unsupported: Final = tuple(
def _pipeline_unsupported_streaming_guardrails(
pipeline: "GuardrailPipeline", translation: "BaseTranslation | None"
) -> tuple[str, ...]:
return tuple(
dict.fromkeys(
step.guardrail for step in pipeline.steps if not _pipeline_step_supports_unified_streaming(step.guardrail)
step.guardrail
for step in pipeline.steps
if not _pipeline_step_supports_streaming(step.guardrail, translation)
)
)
def _pipeline_is_streamable(
policy_name: str, pipeline: "GuardrailPipeline", translation: "BaseTranslation | None"
) -> bool:
unsupported: Final = _pipeline_unsupported_streaming_guardrails(pipeline, translation)
if not unsupported:
return True
verbose_proxy_logger.warning(
"Policy '%s' has post_call pipeline guardrails without the unified apply_guardrail interface, "
"which streaming pipelines need; the stream skips the pipeline and its guardrails run on their own: %s",
"Policy '%s' has post_call pipeline guardrails a streaming pipeline cannot run on this route yet; they "
"need the unified apply_guardrail interface, or a post-call hook without a streaming iterator hook on a "
"route whose translation assembles the streamed response. The stream skips the pipeline and its "
"guardrails run on their own: %s",
policy_name,
", ".join(unsupported),
)
return False
def _route_supports_streaming_pipelines(user_api_key_dict: UserAPIKeyAuth) -> bool:
return not user_api_key_dict.request_route or resolve_endpoint_translation(user_api_key_dict, None) is not None
def _streaming_pipeline_translation(user_api_key_dict: UserAPIKeyAuth) -> "BaseTranslation | None":
resolved: Final = resolve_endpoint_translation(user_api_key_dict, None)
return None if resolved is None else resolved[1]
def _stream_gated_guardrail_names(
def stream_gated_guardrail_names(
request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth
) -> frozenset[str]:
if not _route_supports_streaming_pipelines(user_api_key_dict):
translation: Final = _streaming_pipeline_translation(user_api_key_dict)
if translation is None:
return frozenset()
return _pipeline_step_guardrail_names(
tuple(
(policy_name, pipeline)
for policy_name, pipeline in _post_call_pipelines(request_data)
if all(_pipeline_step_supports_unified_streaming(step.guardrail) for step in pipeline.steps)
if not _pipeline_unsupported_streaming_guardrails(pipeline, translation)
)
)
@ -698,16 +722,19 @@ def _streamable_post_call_pipelines(
The post_call pipelines a streaming response can be gated through.
Streaming pipelines scan the buffered stream through the endpoint guardrail
translation of the request route, so every step's guardrail needs the
unified apply_guardrail interface and the route needs a translation. A
pipeline that cannot be run that way yet is left out and its guardrails
run on the stream on their own, the way they did before pipelines ran on
streams at all, with a warning naming the pipeline.
translation of the request route, so every step's guardrail needs either the
unified apply_guardrail interface or, on a route whose translation assembles
the streamed response, a post-call hook that is its only streaming path, and
the route needs a translation. A pipeline that
cannot be run that way yet is left out and its guardrails run on the stream
on their own, the way they did before pipelines ran on streams at all, with
a warning naming the pipeline.
"""
post_call_pipelines: Final = _post_call_pipelines(request_data)
if not post_call_pipelines:
return ()
if not _route_supports_streaming_pipelines(user_api_key_dict):
translation: Final = _streaming_pipeline_translation(user_api_key_dict)
if translation is None:
verbose_proxy_logger.warning(
"Policies with post_call guardrail pipelines cannot scan streaming responses on route %s yet "
"(no endpoint guardrail translation); the stream skips the pipelines and their guardrails run "
@ -719,7 +746,7 @@ def _streamable_post_call_pipelines(
return tuple(
(policy_name, pipeline)
for policy_name, pipeline in post_call_pipelines
if _pipeline_is_streamable(policy_name, pipeline)
if _pipeline_is_streamable(policy_name, pipeline, translation)
)
@ -2109,7 +2136,7 @@ class ProxyLogging:
)
# Get pipeline-managed guardrails to skip in normal loop
pipeline_managed: Final = _pipeline_managed_guardrail_names(data, "pre_call")
pipeline_managed: Final = pipeline_managed_guardrail_names(data, "pre_call")
caps: Final = ProxyLogging._callback_capabilities()
# Skip the per-request callback walk entirely when nothing in
@ -2874,7 +2901,7 @@ class ProxyLogging:
original_exception=original_exception,
)
request_data.update(_failure_fields_to_lift(request_data))
request_data.update(await offload_token_count(_failure_fields_to_lift)(request_data))
# Remove before callbacks iterate — not serialisable
request_data.pop("litellm_logging_obj", None)
@ -3113,7 +3140,7 @@ class ProxyLogging:
if pipeline_response is not None:
response = pipeline_response # rebind-ok: adopt the pipeline's replacement response, same contract as the callback loops below
pipeline_managed: Final = _pipeline_managed_guardrail_names(data, "post_call")
pipeline_managed: Final = pipeline_managed_guardrail_names(data, "post_call")
guardrail_callbacks, other_callbacks = _partition_post_call_callbacks()
try:
# Merge model-level guardrails before checking which guardrails to run
@ -3429,7 +3456,7 @@ class ProxyLogging:
_cached_guardrail_data: dict | None = None
_guardrail_data_computed = False
pipeline_gated: Final = (
_stream_gated_guardrail_names(data, user_api_key_dict) if caps.has_guardrail else frozenset()
stream_gated_guardrail_names(data, user_api_key_dict) if caps.has_guardrail else frozenset()
)
for callback in litellm.callbacks:
@ -3564,12 +3591,16 @@ class ProxyLogging:
),
)
if post_call_pipelines:
pipeline_translation: Final = (
resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None
)
if pipeline_translation is not None:
current_response = self._pipeline_gated_stream(
response=current_response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
pipelines=post_call_pipelines,
translation=pipeline_translation,
)
try:
@ -3593,6 +3624,7 @@ class ProxyLogging:
user_api_key_dict: UserAPIKeyAuth,
request_data: dict, # mutable-ok: same request-payload shape the hooks mutate
pipelines: "tuple[tuple[str, GuardrailPipeline], ...]",
translation: "tuple[str, BaseTranslation]",
) -> "AsyncGenerator[Any, None]":
"""
Execute post_call policy pipelines against a streamed response.
@ -3602,14 +3634,13 @@ class ProxyLogging:
assembled output through the endpoint guardrail translation, the same
machinery flat post_call guardrails use at end of stream. An allow
releases the buffered chunks: verbatim when no guardrail rewrote the
output, rewritten in place when one rewrote text and the translation
delivers ended-stream rewrites (later steps then re-scan the rewritten
chunks, so rewrites chain). A rewrite the translation cannot deliver
yet (a tool-call rewrite, or a text rewrite on a route without
write-back) is discarded by the executor and the original chunks are
released, as is a buffered shape no translation resolves; a block or
modify_response terminates with the translation's block chunks or the
raised error.
output, rewritten in place when one rewrote text or a tool call and the
translation delivers ended-stream rewrites (later steps then re-scan the
rewritten chunks, so rewrites chain). A rewrite the translation cannot
deliver yet (one on a route without write-back, or a shape the route
refuses) is discarded by the executor and the original chunks are
released; a block or modify_response terminates with the translation's
block chunks or the raised error.
"""
buffered: Final[list[object]] = [] # mutable-ok: accumulates the stream before the pipeline verdict
async for item in response:
@ -3617,17 +3648,7 @@ class ProxyLogging:
if not buffered:
return
resolved: Final = resolve_endpoint_translation(user_api_key_dict, buffered[0])
if resolved is None:
verbose_proxy_logger.warning(
"Policies with post_call guardrail pipelines cannot scan this streaming response shape yet; "
"the stream is released ungoverned by them: %s",
", ".join(policy_name for policy_name, _pipeline in pipelines),
)
for buffered_item in buffered:
yield buffered_item
return
call_type, endpoint_translation = resolved
call_type, endpoint_translation = translation
for policy_name, pipeline in pipelines:
result: PipelineExecutionResult = await PipelineExecutor.execute_steps(

View file

@ -437,14 +437,11 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
response_created_event_data["temperature"] = self.responses_api_request["temperature"]
if "text" in self.responses_api_request:
response_created_event_data["text"] = self.responses_api_request["text"]
if "tool_choice" in self.responses_api_request:
# Transform tool_choice from dict format (e.g., {"type": "auto"}) to string format
response_created_event_data["tool_choice"] = (
LiteLLMCompletionResponsesConfig._transform_tool_choice(self.responses_api_request["tool_choice"])
or "auto"
response_created_event_data["tool_choice"] = (
LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response(
self.responses_api_request.get("tool_choice")
)
else:
response_created_event_data["tool_choice"] = "auto"
)
if "tools" in self.responses_api_request:
response_created_event_data["tools"] = self.responses_api_request["tools"]
else:

View file

@ -27,8 +27,10 @@ from openai.types.chat.chat_completion_named_tool_choice_param import (
)
from openai.types.responses import ResponseFunctionToolCall
from openai.types.responses.response_create_params import ResponseInputParam
from openai.types.responses.tool_choice_custom_param import ToolChoiceCustomParam
from openai.types.responses.tool_choice_function_param import ToolChoiceFunctionParam
from openai.types.responses.tool_param import FunctionToolParam
from pydantic import TypeAdapter
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
@ -68,6 +70,7 @@ from litellm.types.llms.openai import (
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStatus,
ToolChoice,
ValidChatCompletionMessageContentTypes,
ValidChatCompletionMessageContentTypesLiteral,
)
@ -126,6 +129,7 @@ _STR_KEY_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object])
_DICT_ITEMS_LIST_ADAPTER: Final = TypeAdapter(list[dict[object, object]])
_TEXT_ADAPTER: Final = TypeAdapter(str)
_RESPONSES_API_TOOL_CHOICE_ADAPTER: Final = TypeAdapter(ToolChoice)
@runtime_checkable
@ -267,6 +271,27 @@ class LiteLLMCompletionResponsesConfig:
# Return as-is for unknown formats
return tool_choice
@staticmethod
def _transform_tool_choice_for_responses_api_response(tool_choice: object) -> ToolChoice:
if tool_choice is None:
return "auto"
try:
return _RESPONSES_API_TOOL_CHOICE_ADAPTER.validate_python(tool_choice)
except ValidationError:
return LiteLLMCompletionResponsesConfig._chat_tool_choice_as_responses_api_tool_choice(tool_choice)
@staticmethod
def _chat_tool_choice_as_responses_api_tool_choice(tool_choice: object) -> ToolChoice:
match tool_choice, LiteLLMCompletionResponsesConfig._transform_tool_choice(tool_choice):
case {"type": "custom"}, {"function": {"name": str(custom_name)}}:
return ToolChoiceCustomParam(type="custom", name=custom_name)
case _, {"type": "function", "function": {"name": str(function_name)}}:
return ToolChoiceFunctionParam(type="function", name=function_name)
case _, "none" | "auto" | "required" as normalized:
return normalized
case _, _:
return "auto"
@staticmethod
def _should_drop_derived_web_search_options(model: str, custom_llm_provider: str | None) -> bool:
"""
@ -2263,7 +2288,9 @@ class LiteLLMCompletionResponsesConfig:
),
parallel_tool_calls=getattr(chat_completion_response, "parallel_tool_calls", False),
temperature=getattr(chat_completion_response, "temperature", 0),
tool_choice=getattr(chat_completion_response, "tool_choice", "auto"),
tool_choice=LiteLLMCompletionResponsesConfig._transform_tool_choice_for_responses_api_response(
responses_api_request.get("tool_choice")
),
tools=getattr(chat_completion_response, "tools", []),
top_p=getattr(chat_completion_response, "top_p", None),
max_output_tokens=getattr(chat_completion_response, "max_output_tokens", None),

View file

@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runti
import httpx
from openai._streaming import SSEDecoder
from pydantic import BaseModel, ValidationError
from typing_extensions import TypeIs
import litellm
@ -438,18 +439,7 @@ class BaseResponsesAPIStreamingIterator:
if self._persist_completed_response_before_logging:
self._persist_completed_response_to_cache(is_async=is_async)
# Create a copy for logging to avoid modifying the response object that will be returned to the user
# The logging handlers may transform usage from Responses API format (input_tokens/output_tokens)
# to chat completion format (prompt_tokens/completion_tokens) for internal logging
# Use model_dump + model_validate instead of deepcopy to avoid pickle errors with
# Pydantic ValidatorIterator when response contains tool_choice with allowed_tools (fixes #17192)
logging_response = self.completed_response
if self.completed_response is not None and hasattr(self.completed_response, "model_dump"):
try:
logging_response = type(self.completed_response).model_validate(self.completed_response.model_dump())
except Exception:
# Fallback to original if serialization fails
pass
logging_response: Final[object] = _logging_copy(self.completed_response)
self._restore_provider_response_headers(logging_response)
end_time: Final = datetime.now()
@ -488,10 +478,10 @@ class BaseResponsesAPIStreamingIterator:
def _restore_provider_response_headers(self, logging_response: object) -> None:
"""Re-apply the provider's response headers to the copy handed to logging callbacks.
``model_validate(model_dump())`` above drops pydantic private attributes, so the
``model_validate(model_dump())`` in ``_logging_copy`` drops pydantic private attributes, so the
``_hidden_params`` the provider transform set on the nested response are lost. Returns early
when that copy fell back to the original event, so logging-only state never lands on the
object the caller is iterating.
when the event was not a pydantic model and logging got the original, so logging-only state
never lands on the object the caller is iterating.
"""
if logging_response is self.completed_response:
return
@ -544,7 +534,7 @@ class BaseResponsesAPIStreamingIterator:
def _record_failed_response_usage(self, response_obj: ResponsesAPIResponse | None) -> None:
if response_obj is None or self.logging_obj is None:
return
usage_obj: Final[ResponseAPIUsage | None] = getattr(response_obj, "usage", None)
usage_obj: Final[ResponseAPIUsage | None] = _usage_as_model(getattr(response_obj, "usage", None))
if usage_obj is None:
return
try:
@ -1293,14 +1283,46 @@ def _add_text_like_part_events(
)
def _logging_copy(event: object) -> object:
"""Hand logging callbacks a copy, so their usage rewrite (Responses shape to chat shape) never
reaches the event the caller is iterating. The round trip through ``model_dump`` sidesteps the
deepcopy pickle errors of #17192; when a provider payload fails validation (LIT-7391), shallow
copies of the event and its nested response still keep the caller's ``usage`` attribute separate."""
if not isinstance(event, BaseModel):
return event
try:
return type(event).model_validate(event.model_dump())
except Exception:
return _detached_shallow_copy(event)
def _detached_shallow_copy(event: BaseModel) -> BaseModel:
nested: Final[object] = getattr(event, "response", None)
if isinstance(nested, BaseModel):
return event.model_copy(update={"response": nested.model_copy()})
return event.model_copy()
def _usage_as_model(usage: object) -> ResponseAPIUsage | None:
if isinstance(usage, ResponseAPIUsage):
return usage
if not isinstance(usage, dict):
return None
try:
return ResponseAPIUsage.model_validate(usage)
except ValidationError:
return None
def _stamp_responses_usage_cost(
response_obj: ResponsesAPIResponse | None, logging_obj: LiteLLMLoggingObj | None
) -> None:
if response_obj is None or logging_obj is None:
return
usage_obj: Final[ResponseAPIUsage | None] = getattr(response_obj, "usage", None)
usage_obj: Final[ResponseAPIUsage | None] = _usage_as_model(getattr(response_obj, "usage", None))
if usage_obj is None:
return
response_obj.usage = usage_obj # rebind-ok: the stamped cost has to ride on the response the client receives
if isinstance(getattr(usage_obj, "cost", None), (int, float)):
return
try:

View file

@ -67,7 +67,7 @@ from litellm.constants import (
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.asyncify import asyncify, run_async_function
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
coerce_token_limit,
@ -98,6 +98,8 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
mask_credentials_in_payload,
mask_sensitive_structure,
)
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.llms.base_llm.passthrough.transformation import replace_path_segment
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
vector_store_request_metadata,
@ -148,6 +150,8 @@ from litellm.router_utils.common_utils import (
_is_proxy_admin_request,
filter_team_based_models,
filter_web_search_deployments,
get_request_team_id,
provider_for_generic_call,
resolve_model_group_alias,
truncate_fallback_error_detail,
warn_on_provider_credential_mismatch,
@ -5197,7 +5201,7 @@ class Router:
# If get_llm_provider fails, fall back to using model_name as-is
replacement_model_name = model_name
kwargs["endpoint"] = kwargs["endpoint"].replace(model, replacement_model_name)
kwargs["endpoint"] = replace_path_segment(kwargs["endpoint"], model, replacement_model_name)
return kwargs
async def _ageneric_api_call_with_fallbacks_helper(self, model: str, original_generic_function: Callable, **kwargs):
@ -5233,16 +5237,7 @@ class Router:
kwargs=kwargs, model=model, model_name=model_name
)
# Get custom_llm_provider from deployment params
try:
custom_llm_provider = data.get("custom_llm_provider")
_, inferred_custom_llm_provider, _, _ = get_llm_provider(
model=data["model"],
custom_llm_provider=custom_llm_provider,
)
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
except Exception:
custom_llm_provider = None
custom_llm_provider: Final = provider_for_generic_call(data)
response_kwargs: Final = {
**data,
@ -5753,15 +5748,7 @@ class Router:
# Perform pre-call checks for routing strategy
self.routing_strategy_pre_call_checks(deployment=deployment)
try:
custom_llm_provider = data.get("custom_llm_provider")
_, inferred_custom_llm_provider, _, _ = get_llm_provider(
model=data["model"],
custom_llm_provider=custom_llm_provider,
)
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
except Exception:
custom_llm_provider = None
custom_llm_provider: Final = provider_for_generic_call(data)
response: Final = original_function(
**{
@ -11169,28 +11156,31 @@ class Router:
def get_candidate_model_ids_for_route(self, model: str, team_id: str | None = None) -> frozenset[str]:
"""
Deployment ids that could serve ``model`` for ``team_id``, unioned across the paths
the router resolves a route through: ``model_group_alias``, a routing group, the
``model_name`` and team indexes, and wildcard pattern routes. Read-only and
side-effect-free, unlike ``_common_checks_available_deployment`` which also applies
fallbacks and can raise. Lets a pre-call check tell a genuine cross-group route from
same-group unavailability without re-deriving that precedence at the call site, and
without leaking deployment ids into request kwargs bound for the provider.
Deployment ids that could serve ``model`` for ``team_id``, following the same
precedence ``_common_checks_available_deployment`` uses to build a candidate pool:
``model_group_alias``, then a routing group, then the first matching early-resolve
path for a name that is not a ``model_name`` (team route, wildcard pattern via
``get_deployments_by_pattern``, team pattern router, default deployment), then the
``model_name`` and team indexes. Delegating to the router's own resolvers keeps this
aligned with how a route actually resolves rather than re-deriving it, and unlike
``_common_checks_available_deployment`` it is read-only: it does not apply request
fallbacks and (with ``include_team_models`` left off) does not raise. Lets a pre-call
check tell a genuine cross-group route from same-group unavailability without leaking
deployment ids into request kwargs bound for the provider.
"""
resolved: Final = self._get_model_from_alias(model=model) or model
routing_group_members: Final = self._get_routing_group_deployments(model=resolved, team_id=team_id)
if routing_group_members is not None:
return self._deployment_ids(routing_group_members)
if resolved in self.model_names:
return self._deployment_ids(self._get_all_deployments(model_name=resolved, team_id=team_id))
team_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None
return self._deployment_ids(
(
*self._get_all_deployments(model_name=resolved, team_id=team_id),
*(self.pattern_router.route(resolved) or ()),
*((team_router.route(resolved) or ()) if team_router is not None else ()),
)
early: Final = self._try_early_resolve_deployments_for_model_not_in_names(
model=resolved, request_team_id=team_id
)
if early is not None:
early_deployments: Final = early[1]
return self._deployment_ids(
(early_deployments,) if isinstance(early_deployments, Mapping) else early_deployments
)
return self._deployment_ids(self._get_all_deployments(model_name=resolved, team_id=team_id))
@staticmethod
def _deployment_ids(deployments: Sequence[Mapping[str, object]]) -> frozenset[str]:
@ -12109,7 +12099,7 @@ class Router:
try:
if not self._pre_call_checks_need_token_count(model, healthy_deployments):
return None
return await asyncify(self._count_pre_call_check_tokens)(
return await offload_token_count(self._count_pre_call_check_tokens)(
messages=cast(list[dict[str, str]] | None, messages), # cast-ok: forwarded to the sync counter
input=cast(str | list | None, input), # cast-ok: forwarded to the sync counter
request_kwargs=request_kwargs,
@ -12337,27 +12327,7 @@ class Router:
if team_deployments:
return model, team_deployments
elif include_team_models:
team_deployments = [
self.model_list[index]
for (_, public_model_name), indices in self.team_model_to_deployment_indices.items()
if public_model_name == model
for index in indices
]
team_ids: Final = {
team_id
for deployment in team_deployments
for team_id in [(deployment.get("model_info") or {}).get("team_id")]
if team_id is not None
}
if len(team_ids) > 1:
raise litellm.BadRequestError(
message=(
f"Model name '{model}' matches deployments from multiple teams. "
"Specify the deployment ID directly to disambiguate."
),
model=model,
llm_provider="",
)
team_deployments = self._team_deployments_across_teams(model)
if team_deployments:
return model, team_deployments
@ -12384,6 +12354,45 @@ class Router:
return None
def _team_deployments_across_teams(self, model: str) -> list[DeploymentTypedDict]:
"""Every team's deployments under public name `model`, for a proxy admin calling without a team."""
team_deployments: Final = [
self.model_list[index]
for (_, public_model_name), indices in self.team_model_to_deployment_indices.items()
if public_model_name == model
for index in indices
]
team_ids: Final = {
team_id
for deployment in team_deployments
for team_id in [(deployment.get("model_info") or {}).get("team_id")]
if team_id is not None
}
if len(team_ids) > 1:
raise litellm.BadRequestError(
message=(
f"Model name '{model}' matches deployments from multiple teams. "
"Specify the deployment ID directly to disambiguate."
),
model=model,
llm_provider="",
)
return team_deployments
def deployments_for_request(
self, model: str, request_kwargs: Mapping[str, object]
) -> Sequence[DeploymentTypedDict]:
"""The deployments `model` names for this caller, through the same alias, then team-first, then
global, then admin-across-teams resolution `_common_checks_available_deployment` applies, so
strategy selection and compression policy can never disagree with deployment selection about
which marker a name means."""
registered_name: Final = self._get_model_from_alias(model=model) or model
team_id: Final = get_request_team_id(request_kwargs)
deployments: Final = self._get_all_deployments(model_name=registered_name, team_id=team_id)
if deployments or team_id is not None or not _is_proxy_admin_request(request_kwargs):
return deployments
return self._team_deployments_across_teams(registered_name)
@staticmethod
def _is_strategy_marker_deployment(deployment: Mapping[str, object]) -> bool:
litellm_params: Final = deployment.get("litellm_params")
@ -12411,11 +12420,7 @@ class Router:
- Dict, if specific model chosen
"""
request_team_id: str | None = None
if request_kwargs is not None:
metadata: Final = request_kwargs.get("metadata") or {}
litellm_metadata: Final = request_kwargs.get("litellm_metadata") or {}
request_team_id = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id")
request_team_id: Final = get_request_team_id(request_kwargs)
# check if aliases set on litellm model alias map
if specific_deployment is True:
return model, self._get_deployment_by_litellm_model(model=model)
@ -12440,7 +12445,9 @@ class Router:
include_team_models=_is_proxy_admin_request(request_kwargs),
)
if early is not None:
return early
if not isinstance(early[1], list):
return early
return early[0], self._drop_strategy_markers(early[0], early[1])
## get healthy deployments
### get all deployments
@ -12517,19 +12524,22 @@ class Router:
model
] # update the model to the actual value if an alias has been passed in
marker_flags: Final = tuple(self._is_strategy_marker_deployment(d) for d in healthy_deployments)
if not any(marker_flags):
return model, healthy_deployments
selectable: Final = [ # mutable-ok: matches this function's list contract expected by downstream filters
d for d, is_marker in zip(healthy_deployments, marker_flags, strict=True) if not is_marker
return model, self._drop_strategy_markers(model, healthy_deployments)
def _drop_strategy_markers(
self, model: str, deployments: Sequence[DeploymentTypedDict]
) -> list[DeploymentTypedDict]:
"""A strategy marker is never a callable deployment, whichever resolution arm produced it."""
selectable: Final = [ # mutable-ok: matches _common_checks_available_deployment's list contract
d for d in deployments if not self._is_strategy_marker_deployment(d)
]
if not selectable:
if deployments and not selectable:
raise litellm.BadRequestError(
message=f"You passed in model={model}. {RouterErrors.only_strategy_marker_deployments.value}",
model=model,
llm_provider="",
)
return model, selectable
return selectable
def _filter_deployments_by_model_access_groups(
self,
@ -13219,12 +13229,8 @@ class Router:
return filtered
def _model_name_has_plain_deployments(self, model: str) -> bool:
indices: Final = self.model_name_to_deployment_indices.get(model) or ()
return any(not self._is_strategy_marker_deployment(self.model_list[idx]) for idx in indices)
def _select_pre_routing_strategy(
self, model: str, request_kwargs: dict
self, model: str, request_kwargs: Mapping[str, object]
) -> "TaggedPreRoutingStrategy[PreRoutingStrategy] | None":
"""
Resolve the pre-routing strategy for `model`, disambiguating deployments
@ -13235,6 +13241,12 @@ class Router:
deployment the strategy was registered from via its (model_name, tags)
pair.
The registries are keyed by the marker deployment's own `model_name`, which
for a team-scoped router is the internal `model_name_{team}_{uuid}` while
the caller sends the team's public name. So the names looked up are the
`model_name`s of whatever deployments this caller's request resolves `model`
to, and `model` itself when it resolves to none.
With tag filtering enabled, router-wide or by the request's
enable_tag_filtering (which the proxy sets from key/team
router_settings), strategies that all carry real tags matching none of
@ -13242,12 +13254,14 @@ class Router:
deployments: returning None hands the request to ordinary tag-aware
deployment selection.
"""
candidates: Final[list[TaggedPreRoutingStrategy[PreRoutingStrategy]]] = [
*self.auto_routers.get(model, []),
*self.complexity_routers.get(model, []),
*self.adaptive_routers.get(model, []),
*self.quality_routers.get(model, []),
]
registries: Final = (self.auto_routers, self.complexity_routers, self.adaptive_routers, self.quality_routers)
if not any(registries):
return None
deployments: Final = self.deployments_for_request(model, request_kwargs)
registered_names: Final = tuple(dict.fromkeys(str(d["model_name"]) for d in deployments)) or (model,)
candidates: Final = tuple(
tagged for registry in registries for name in registered_names for tagged in registry.get(name, [])
)
if not candidates:
return None
@ -13265,7 +13279,7 @@ class Router:
if (
(self.enable_tag_filtering or request_scoped_filtering)
and all(tagged.tags for tagged in candidates)
and self._model_name_has_plain_deployments(model)
and any(not self._is_strategy_marker_deployment(d) for d in deployments)
):
return None
return candidates[0]
@ -13377,11 +13391,12 @@ class Router:
Used for the litellm auto-router to modify the request before the routing decision is made.
`model` is whatever the caller asked for, which may be a `model_group_alias` key, while the
strategy registries and the marker deployment are keyed by the marker's own `model_name`, so
every lookup below resolves the alias first. Only the lookups: the caller-facing name stays
the alias, since spend metadata is stamped before routing and the response carries the tier
group the strategy picked.
`model` is whatever the caller asked for, which may be a `model_group_alias` key or a team's
public model name, while the strategy registries and the marker deployment are keyed by the
marker's own `model_name`, so every lookup below resolves the alias first and the team name
through the deployment path. Only the lookups: the caller-facing name stays the alias, since
spend metadata is stamped before routing and the response carries the tier group the
strategy picked.
"""
requested_registered_model_name: Final = self._get_model_from_alias(model=model) or model
registered_model_name: Final = await self._resolve_claude_code_session_router(
@ -13418,7 +13433,6 @@ class Router:
messages_for_routing,
model_hop_compression_armed,
policy_for_model,
team_id_from_request,
)
# Same tag-aware lookup the proxy's pre-call arming used, so an alias with
@ -13426,7 +13440,7 @@ class Router:
compression_policy: Final = policy_for_model(
llm_router=self,
model_alias=registered_model_name,
team_id=team_id_from_request(request_kwargs),
request_kwargs=request_kwargs,
request_tags=_get_tags_from_request_kwargs(request_kwargs),
)
# Shared compression already ran in the pre-call hook, so reuse it rather than
@ -13495,7 +13509,9 @@ class Router:
# Per-tier `litellm_params` on the hook response are deliberate overrides
# the caller applies on top, so those keys are never forwarded here.
marker_params: Final = (
self._forwardable_alias_marker_params(model=registered_model_name, strategy_tags=selected_strategy.tags)
self._forwardable_alias_marker_params(
model=registered_model_name, strategy_tags=selected_strategy.tags, request_kwargs=request_kwargs
)
if pre_routing_hook_response is not None
else ()
)
@ -13513,13 +13529,14 @@ class Router:
return pre_routing_hook_response
def _forwardable_alias_marker_params(
self, model: str, strategy_tags: tuple[str, ...]
self, model: str, strategy_tags: tuple[str, ...], request_kwargs: Mapping[str, object]
) -> tuple[tuple[str, object], ...]:
marker_params: Final = tuple(
litellm_params
for idx in self.model_name_to_deployment_indices.get(model, ())
if isinstance(litellm_params := self.model_list[idx].get("litellm_params", {}), dict)
and str(litellm_params.get("model", "")).startswith(AUTO_ROUTER_MODEL_PREFIX)
for deployment in self.deployments_for_request(model, request_kwargs)
if str((litellm_params := deployment["litellm_params"]).get("model", "")).startswith(
AUTO_ROUTER_MODEL_PREFIX
)
)
tag_matched: Final = tuple(
params for params in marker_params if tuple(params.get("tags") or ()) == strategy_tags

View file

@ -2568,14 +2568,14 @@ class ComplexityRouter(CustomLogger):
"""Real-tokenizer count of the resolved messages plus the out-of-band carriers, off the
event loop; None when counting fails, and the gate then leaves the placement alone."""
import litellm
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.token_counter import offload_token_count
out_of_band: Final = self._out_of_band_request_text(request_kwargs)
try:
counted: Final = await asyncify(litellm.token_counter)(
counted: Final = await offload_token_count(litellm.token_counter)(
messages=cast(list, resolved_messages) # cast-ok: token_counter only iterates the sequence
)
return counted + (await asyncify(litellm.token_counter)(text=out_of_band) if out_of_band else 0)
return counted + (await offload_token_count(litellm.token_counter)(text=out_of_band) if out_of_band else 0)
except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request
verbose_router_logger.debug("ComplexityRouter: context-window token count failed. Got - %s", e)
return None

View file

@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Final
if TYPE_CHECKING:
from litellm.types.llms.openai import OpenAIFileObject
import litellm
from litellm._logging import verbose_logger, verbose_router_logger
from litellm.constants import ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
from litellm.exceptions import BadRequestError
@ -26,6 +27,18 @@ def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool
return getattr(user_api_key_auth, "user_role", None) == "proxy_admin"
def get_request_team_id(request_kwargs: Mapping[str, object] | None) -> str | None:
"""The caller's team id, from whichever metadata bucket this surface writes to."""
if request_kwargs is None:
return None
for bucket_name in ("metadata", "litellm_metadata"):
bucket = request_kwargs.get(bucket_name)
team_id = bucket.get("user_api_key_team_id") if isinstance(bucket, Mapping) else None
if isinstance(team_id, str) and team_id:
return team_id
return None
def resolve_model_group_alias(model_group_alias: object, model: str) -> str | None:
"""
Resolve ``model`` through a ``model_group_alias`` map.
@ -110,7 +123,7 @@ def filter_team_based_models(
metadata: Final = request_kwargs.get("metadata") or {}
litellm_metadata: Final = request_kwargs.get("litellm_metadata") or {}
request_team_id: Final = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id")
request_team_id: Final = get_request_team_id(request_kwargs)
if request_team_id is None and _is_proxy_admin_request(request_kwargs) and isinstance(healthy_deployments, list):
requested_model: Final = (
request_kwargs.get("model") or metadata.get("model_group") or litellm_metadata.get("model_group")
@ -244,6 +257,32 @@ PROVIDER_SCOPED_CREDENTIAL_PARAMS: Final[Mapping[str, frozenset[str]]] = Mapping
)
def provider_for_generic_call(litellm_params: Mapping[str, object]) -> str | None:
"""
The provider the router hands a deployment's generic SDK call, or None when it cannot be resolved.
A model that carries its own provider prefix keeps that prefix even where get_llm_provider
would resolve it to a sibling provider (azure_ai/<openai model> on an Azure OpenAI host
resolves to azure): the SDK call still receives the prefixed model, and an explicit provider
that contradicts the prefix makes get_llm_provider re-prefix it into a deployment name that
does not exist upstream.
"""
declared: Final = litellm_params.get("custom_llm_provider")
if isinstance(declared, str) and declared:
return declared
model: Final = litellm_params.get("model")
if not isinstance(model, str) or not model:
return None
prefix: Final = model.split("/", 1)[0]
if "/" in model and prefix in litellm.provider_list:
return prefix
try:
_, inferred, _, _ = get_llm_provider(model=model)
except BadRequestError:
return None
return inferred
def warn_on_provider_credential_mismatch(model_name: str, litellm_params: Mapping[str, object]) -> str | None:
"""
Warn when a deployment carries one provider's credentials but resolves to another.

View file

@ -37,7 +37,7 @@ Safe to enable globally:
"""
import time
from collections.abc import Mapping
from collections.abc import Iterator, Mapping
from typing import TYPE_CHECKING, Final, Optional, Protocol, cast
import httpx
@ -48,6 +48,10 @@ from litellm.exceptions import (
ServiceUnavailableError,
)
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.prompt_templates.common_utils import (
encrypted_content_of_block,
strip_encrypted_reasoning_from_messages,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.router_utils.cooldown_cache import CooldownCacheValue
from litellm.types.llms.openai import AllMessageValues
@ -138,15 +142,48 @@ class EncryptedContentAffinityCheck(CustomLogger):
# If no encoded ID, check if encrypted_content itself is wrapped
encrypted_content = item.get("encrypted_content")
if encrypted_content and isinstance(encrypted_content, str):
(
model_id,
_,
) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(encrypted_content)
model_id = EncryptedContentAffinityCheck._model_id_from_wrapped_encrypted_content(encrypted_content)
if model_id:
return model_id
return None
@staticmethod
def _anthropic_content_blocks(messages: object) -> Iterator[Mapping[str, object]]:
if not isinstance(messages, list):
return iter(())
return (
cast(Mapping[str, object], block) # cast-ok: narrowed by isinstance
for message in cast(list[object], messages) # cast-ok: narrowed by isinstance
if isinstance(message, Mapping)
for content in (cast(Mapping[str, object], message).get("content"),) # cast-ok: narrowed by isinstance
if isinstance(content, list)
for block in cast(list[object], content) # cast-ok: narrowed by isinstance
if isinstance(block, Mapping)
)
@staticmethod
def _model_id_from_wrapped_encrypted_content(encrypted_content: str) -> str | None:
model_id, _ = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(encrypted_content)
return model_id or None
@staticmethod
def _extract_model_id_from_anthropic_messages(messages: object) -> str | None:
return next(
(
model_id
for block in EncryptedContentAffinityCheck._anthropic_content_blocks(messages)
if (encrypted_content := encrypted_content_of_block(block)) is not None
if (
model_id := EncryptedContentAffinityCheck._model_id_from_wrapped_encrypted_content(
encrypted_content
)
)
is not None
),
None,
)
@staticmethod
def _find_deployment_by_model_id(healthy_deployments: list[dict], model_id: str) -> dict | None:
for deployment in healthy_deployments:
@ -240,8 +277,9 @@ class EncryptedContentAffinityCheck(CustomLogger):
parent_otel_span: Span | None = None,
) -> list[dict]:
"""
If the request ``input`` contains litellm-encoded item IDs, decode the
embedded ``model_id`` and pin the request to that deployment. Raises
If the request ``input`` contains litellm-encoded item IDs, or its Anthropic
``messages`` replay a bridge-tagged thinking block, decode the embedded
``model_id`` and pin the request to that deployment. Raises
``RateLimitError`` / ``ServiceUnavailableError`` when the originating
deployment is a member of the routed model group but currently unavailable
and no encryption-boundary peer exists, rather than dispatching a doomed
@ -270,12 +308,15 @@ class EncryptedContentAffinityCheck(CustomLogger):
request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] = True
request_input: Final = request_kwargs.get("input")
model_id: Final = self._extract_model_id_from_input(request_input)
anthropic_messages: Final = messages or request_kwargs.get("messages")
model_id: Final = self._extract_model_id_from_input(
request_input
) or self._extract_model_id_from_anthropic_messages(anthropic_messages)
if not model_id:
return typed_healthy_deployments
verbose_router_logger.debug(
"EncryptedContentAffinityCheck: decoded model_id=%s from input item IDs",
"EncryptedContentAffinityCheck: decoded model_id=%s from the request's encrypted content markers",
model_id,
)
@ -327,6 +368,7 @@ class EncryptedContentAffinityCheck(CustomLogger):
model,
)
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
strip_encrypted_reasoning_from_messages(anthropic_messages)
return typed_healthy_deployments
# The origin is a member of the routed group but currently unavailable (cooled down); fail fast

View file

@ -21,6 +21,7 @@ import litellm
from litellm import token_counter
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.types.router import RouterCacheEnum, RouterErrors
from litellm.utils import get_utc_datetime
@ -466,7 +467,7 @@ async def async_io_token_pre_call_check(
request_kwargs: Final = get_io_token_rate_limit_request_kwargs()
_model: Final = (deployment.get("litellm_params") or {}).get("model") or ""
estimated_input: Final = _estimate_input_tokens(request_kwargs, model=_model)
estimated_input: Final = await offload_token_count(_estimate_input_tokens)(request_kwargs, model=_model)
max_tokens: Final = _resolve_max_tokens(request_kwargs, deployment)
dt: Final = get_utc_datetime()

View file

@ -14,6 +14,7 @@ from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import CallTypes, StandardLoggingPayload
from litellm.utils import get_prompt_cache_min_tokens, is_prompt_caching_valid_prompt
@ -61,7 +62,7 @@ class PromptCachingDeploymentCheck(CustomLogger):
if request_kwargs is not None and request_kwargs.get("_target_order") is not None:
return healthy_deployments
if messages is not None and is_prompt_caching_valid_prompt(
if messages is not None and await offload_token_count(is_prompt_caching_valid_prompt)(
messages=messages,
model=model,
min_token_count=_get_min_token_count_for_deployments(healthy_deployments),
@ -139,7 +140,7 @@ class PromptCachingDeploymentCheck(CustomLogger):
return
## PROMPT CACHING - cache model id, if prompt caching valid prompt + provider
if is_prompt_caching_valid_prompt(
if await offload_token_count(is_prompt_caching_valid_prompt)(
model=model,
messages=cast(list[AllMessageValues], messages),
):

View file

@ -2,6 +2,7 @@ from typing import Any, Literal
from pydantic import BaseModel
from typing_extensions import (
ReadOnly,
Required,
TypedDict,
)
@ -57,6 +58,14 @@ class DatabricksMessage(TypedDict, total=False):
role: Required[str]
content: Required[AllDatabricksContentValues]
tool_calls: list[DatabricksTool] | None
reasoning_content: ReadOnly[str | None]
reasoning: ReadOnly[str | None]
class DatabricksDelta(TypedDict, total=False):
role: ReadOnly[str]
content: ReadOnly[AllDatabricksContentValues | None]
reasoning_content: ReadOnly[str | None]
class DatabricksChoice(TypedDict, total=False):

View file

@ -1564,6 +1564,9 @@ class ResponseIncompleteEvent(BaseLiteLLMOpenAIResponseObject):
response: ResponsesAPIResponse
ResponsesTerminalEvent: TypeAlias = ResponseCompletedEvent | ResponseIncompleteEvent | ResponseFailedEvent
class ResponsePartAddedEvent(BaseLiteLLMOpenAIResponseObject):
type: Literal[ResponsesAPIStreamEvents.RESPONSE_PART_ADDED]
item_id: str

View file

@ -250,7 +250,23 @@ class MCPServer(BaseModel):
@property
def advertises_gateway_authorization_server(self) -> bool:
"""Whether named discovery should advertise the aggregate gateway authorization server."""
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
if self.auth_type == MCPAuth.oauth2:
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
if self.auth_type not in (
None,
MCPAuth.none,
MCPAuth.api_key,
MCPAuth.bearer_token,
MCPAuth.basic,
MCPAuth.authorization,
MCPAuth.token,
MCPAuth.aws_sigv4,
):
return False
return not any(
header.lower() in ("authorization", "x-api-key", "api-key", "apikey")
for header in (self.extra_headers or ())
)
@property
def is_true_passthrough(self) -> bool:

View file

@ -525,6 +525,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
input_cost_per_second: float | None
output_cost_per_second: float | None
output_cost_per_second_480p: ReadOnly[float | None]
output_cost_per_second_720p: ReadOnly[float | None]
output_cost_per_second_1080p: float | None
output_cost_per_second_4k: ReadOnly[float | None]
num_retries: int | None

View file

@ -318,6 +318,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
float | None
) # video_generation tier: key output_cost_per_second_<resolution> (e.g. 1080p, 720p)
output_cost_per_second_480p: ReadOnly[float | None]
output_cost_per_second_720p: ReadOnly[float | None]
output_cost_per_second_4k: ReadOnly[float | None]
ocr_cost_per_page: float | None # for OCR models
ocr_cost_per_credit: float | None # for OCR models priced by credit
@ -3522,6 +3523,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
output_cost_per_second: float | None = None
output_cost_per_second_1080p: float | None = None
output_cost_per_second_480p: float | None = None
output_cost_per_second_720p: float | None = None
output_cost_per_second_4k: float | None = None
input_cost_per_pixel: float | None = None
output_cost_per_pixel: float | None = None

View file

@ -2293,15 +2293,7 @@ def create_pretrained_tokenizer(identifier: str, revision="main", auth_token: st
dict: A dictionary with the tokenizer and its type.
"""
try:
tokenizer = Tokenizer.from_pretrained(
identifier,
revision=revision,
auth_token=auth_token,
)
except Exception as e:
verbose_logger.error("Error creating pretrained tokenizer: %s. Defaulting to version without 'auth_token'.", e)
tokenizer = Tokenizer.from_pretrained(identifier, revision=revision)
tokenizer: Final = Tokenizer.from_pretrained(identifier, revision=revision, token=auth_token)
return {"type": "huggingface_tokenizer", "tokenizer": tokenizer}
@ -3412,7 +3404,7 @@ def get_optional_params_image_gen(
non_default_params=non_default_params,
optional_params=optional_params,
model=model or "",
drop_params=drop_params if drop_params is not None else False,
drop_params=litellm.drop_params is True or drop_params is True,
)
elif (
custom_llm_provider == "openai"
@ -5913,6 +5905,7 @@ def _get_model_info_helper(
output_cost_per_second=_model_info.get("output_cost_per_second", None),
output_cost_per_second_1080p=_model_info.get("output_cost_per_second_1080p", None),
output_cost_per_second_480p=_model_info.get("output_cost_per_second_480p", None),
output_cost_per_second_720p=_model_info.get("output_cost_per_second_720p", None),
output_cost_per_second_4k=_model_info.get("output_cost_per_second_4k", None),
output_cost_per_video_per_second=_model_info.get("output_cost_per_video_per_second", None),
output_cost_per_image=_model_info.get("output_cost_per_image", None),
@ -8857,6 +8850,12 @@ class ProviderConfigManager:
)
return AzurePassthroughConfig()
elif LlmProviders.AZURE_AI == provider:
from litellm.llms.azure_ai.passthrough.transformation import (
AzureAIPassthroughConfig,
)
return AzureAIPassthroughConfig()
elif LlmProviders.GIGACHAT == provider:
from litellm.llms.gigachat.passthrough.transformation import (
GigaChatPassthroughConfig,

File diff suppressed because it is too large Load diff

View file

@ -478,6 +478,10 @@
"type": "number",
"minimum": 0
},
"output_cost_per_second_720p": {
"type": "number",
"minimum": 0
},
"output_cost_per_token": {
"type": "number",
"minimum": 0,

Some files were not shown because too many files have changed in this diff Show more