mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
chore: merge upstream staging and preserve authentication regressions
This commit is contained in:
commit
7cd2c354d8
188 changed files with 11737 additions and 1442 deletions
|
|
@ -475,6 +475,7 @@ prometheus_metrics_config: Optional[List] = None
|
|||
prometheus_exclude_metrics: Optional[List[str]] = None
|
||||
prometheus_exclude_labels: Optional[List[str]] = None
|
||||
prometheus_emit_stream_label: bool = False
|
||||
prometheus_emit_input_sequence_length_label: bool = False
|
||||
prometheus_deployment_and_latency_caller_identity: Literal[
|
||||
"api_key_alias",
|
||||
"user_email",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
@ -1669,6 +1682,7 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT
|
|||
RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS", "3")))
|
||||
RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2"))
|
||||
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
|
||||
PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25")))
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
|
||||
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -246,6 +246,7 @@ class PrometheusLogger(CustomLogger):
|
|||
# logger so toggling these flags only takes effect after a
|
||||
# restart, keeping init-time and runtime label sets in sync.
|
||||
self._cached_metric_labels: dict[str, list[str]] = {}
|
||||
self._emit_input_sequence_length_label = litellm.prometheus_emit_input_sequence_length_label is True
|
||||
|
||||
_custom_buckets: Final = litellm.prometheus_latency_buckets
|
||||
self.latency_buckets = tuple(_custom_buckets) if _custom_buckets is not None else LATENCY_BUCKETS
|
||||
|
|
@ -1522,6 +1523,11 @@ class PrometheusLogger(CustomLogger):
|
|||
# 2. Pyright does not allow us to run isinstance(standard_logging_payload, StandardLoggingPayload) <- this would be ideal
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
input_sequence_length=(
|
||||
self._get_input_sequence_length(standard_logging_payload, kwargs, response_obj)
|
||||
if self._emit_input_sequence_length_label
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
# set x-ratelimit headers
|
||||
|
|
@ -2192,6 +2198,36 @@ class PrometheusLogger(CustomLogger):
|
|||
)
|
||||
self.litellm_remaining_api_key_tokens_for_model.labels(**tokens_labels).set(remaining_tokens)
|
||||
|
||||
@staticmethod
|
||||
def _get_input_sequence_length(
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
) -> str:
|
||||
prompt_tokens: Final = standard_logging_payload.get("prompt_tokens")
|
||||
if prompt_tokens:
|
||||
return get_input_sequence_length_bucket(prompt_tokens)
|
||||
combined_usage: Final = kwargs.get("combined_usage_object")
|
||||
if (
|
||||
combined_usage is not None
|
||||
and getattr(kwargs.get("_litellm_upstream_reported_usage"), "total_tokens", None) is not None
|
||||
):
|
||||
return get_input_sequence_length_bucket(None)
|
||||
reported_usage: Final = (
|
||||
response_obj.get("usage") if isinstance(response_obj, dict) else getattr(response_obj, "usage", None)
|
||||
)
|
||||
if reported_usage is None and combined_usage is None:
|
||||
return get_input_sequence_length_bucket(None)
|
||||
usage_metadata: Final = standard_logging_payload["metadata"].get("usage_object")
|
||||
if isinstance(usage_metadata, Mapping):
|
||||
return get_input_sequence_length_bucket(usage_metadata.get("prompt_tokens"))
|
||||
if combined_usage is None and isinstance(response_obj, dict):
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
|
||||
normalized_usage: Final[Mapping[str, object]] = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj)
|
||||
return get_input_sequence_length_bucket(normalized_usage.get("prompt_tokens"))
|
||||
return get_input_sequence_length_bucket(prompt_tokens)
|
||||
|
||||
def _set_latency_metrics(
|
||||
self,
|
||||
kwargs: dict,
|
||||
|
|
@ -2202,7 +2238,16 @@ class PrometheusLogger(CustomLogger):
|
|||
user_api_team_alias: str | None,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
label_context: PrometheusLabelFactoryContext | None = None,
|
||||
input_sequence_length: str | None = None,
|
||||
):
|
||||
latency_enum_values: Final = (
|
||||
replace(enum_values, input_sequence_length=input_sequence_length)
|
||||
if input_sequence_length is not None
|
||||
else enum_values
|
||||
)
|
||||
latency_label_context: Final = (
|
||||
PrometheusLabelFactoryContext(latency_enum_values) if input_sequence_length is not None else label_context
|
||||
)
|
||||
# latency metrics
|
||||
end_time: Final[datetime] = kwargs.get("end_time") or datetime.now()
|
||||
start_time: Final[datetime | None] = kwargs.get("start_time")
|
||||
|
|
@ -2220,8 +2265,8 @@ class PrometheusLogger(CustomLogger):
|
|||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_llm_api_time_to_first_token_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
enum_values=latency_enum_values,
|
||||
label_context=latency_label_context,
|
||||
)
|
||||
self.litellm_llm_api_time_to_first_token_metric.labels(**_ttft_labels).observe(time_to_first_token_seconds)
|
||||
self._track_end_user_metric_series(
|
||||
|
|
@ -2241,8 +2286,8 @@ class PrometheusLogger(CustomLogger):
|
|||
if api_call_total_time_seconds is not None:
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_llm_api_latency_metric"),
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
enum_values=latency_enum_values,
|
||||
label_context=latency_label_context,
|
||||
)
|
||||
self.litellm_llm_api_latency_metric.labels(**_labels).observe(api_call_total_time_seconds)
|
||||
self._track_end_user_metric_series(
|
||||
|
|
@ -2272,8 +2317,8 @@ class PrometheusLogger(CustomLogger):
|
|||
)
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_request_total_latency_metric"),
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
enum_values=latency_enum_values,
|
||||
label_context=latency_label_context,
|
||||
)
|
||||
self.litellm_request_total_latency_metric.labels(**_labels).observe(_observed_total_time_seconds)
|
||||
self._track_end_user_metric_series(
|
||||
|
|
|
|||
|
|
@ -10,9 +10,11 @@ import asyncio
|
|||
import time
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final, cast
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
|
||||
|
|
@ -24,7 +26,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,
|
||||
|
|
@ -35,6 +37,9 @@ from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload
|
|||
|
||||
from .custom_batch_logger import CustomBatchLogger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
|
||||
class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
||||
def __init__(
|
||||
|
|
@ -232,6 +237,26 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
f"{get_aws_dns_suffix(self.s3_region_name)}/{encoded_key}"
|
||||
)
|
||||
|
||||
def _sign_put(
|
||||
self, credentials: "Credentials", url: str, json_string: str, headers: Mapping[str, str]
|
||||
) -> dict[str, str]: # mutable-ok: [LIT001] AsyncHTTPHandler.put/HTTPHandler.put only accept dict headers
|
||||
"""
|
||||
``RefreshableCredentials`` (IMDS roles) may refresh between the access key, secret and token
|
||||
reads SigV4 performs, producing a mixed-generation signature that S3 rejects with 403.
|
||||
Freezing first makes the three values one atomic snapshot.
|
||||
"""
|
||||
from botocore.auth import S3SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
from botocore.credentials import RefreshableCredentials
|
||||
|
||||
frozen: Final = (
|
||||
credentials.get_frozen_credentials() if isinstance(credentials, RefreshableCredentials) else credentials
|
||||
)
|
||||
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=dict(headers))
|
||||
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
|
||||
S3SigV4Auth(frozen, "s3", aws_region_name).add_auth(aws_request)
|
||||
return dict(aws_request.headers.items())
|
||||
|
||||
def _sse_headers(self) -> Mapping[str, str]:
|
||||
candidates: Final = {
|
||||
"x-amz-server-side-encryption": self.s3_server_side_encryption,
|
||||
|
|
@ -317,26 +342,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
from botocore.auth import S3SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
try:
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
|
||||
asyncified_get_credentials: Final = asyncify(self.get_credentials)
|
||||
credentials: Final = await asyncified_get_credentials(
|
||||
aws_access_key_id=self.s3_aws_access_key_id,
|
||||
aws_secret_access_key=self.s3_aws_secret_access_key,
|
||||
aws_session_token=self.s3_aws_session_token,
|
||||
aws_region_name=self.s3_region_name,
|
||||
aws_session_name=self.s3_aws_session_name,
|
||||
aws_profile_name=self.s3_aws_profile_name,
|
||||
aws_role_name=self.s3_aws_role_name,
|
||||
aws_web_identity_token=self.s3_aws_web_identity_token,
|
||||
aws_sts_endpoint=self.s3_aws_sts_endpoint,
|
||||
)
|
||||
|
||||
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
|
||||
verbose_logger.debug("s3_v2 logger - s3_verify setting: %s", self.s3_verify)
|
||||
|
|
@ -363,19 +374,28 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
**self._sse_headers(),
|
||||
}
|
||||
|
||||
# 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)
|
||||
async def signed_put() -> httpx.Response:
|
||||
credentials: Final = await asyncified_get_credentials(
|
||||
aws_access_key_id=self.s3_aws_access_key_id,
|
||||
aws_secret_access_key=self.s3_aws_secret_access_key,
|
||||
aws_session_token=self.s3_aws_session_token,
|
||||
aws_region_name=self.s3_region_name,
|
||||
aws_session_name=self.s3_aws_session_name,
|
||||
aws_profile_name=self.s3_aws_profile_name,
|
||||
aws_role_name=self.s3_aws_role_name,
|
||||
aws_web_identity_token=self.s3_aws_web_identity_token,
|
||||
aws_sts_endpoint=self.s3_aws_sts_endpoint,
|
||||
)
|
||||
signed_headers: Final = await run_aws_signing(self._sign_put, credentials, url, json_string, headers)
|
||||
try:
|
||||
return await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
|
||||
# Prepare the signed headers
|
||||
signed_headers: Final = dict(aws_request.headers.items())
|
||||
|
||||
# Make the request with retry for transient S3 errors (500/503)
|
||||
max_retries: Final = 3
|
||||
for attempt in range(max_retries):
|
||||
response = await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
|
||||
if response.status_code in (500, 503) and attempt < max_retries - 1:
|
||||
response = await signed_put()
|
||||
if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
|
||||
wait_time = 2**attempt # 1s, 2s
|
||||
verbose_logger.warning(
|
||||
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
|
||||
|
|
@ -479,20 +499,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
from botocore.auth import S3SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
from botocore.credentials import Credentials
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
try:
|
||||
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
|
||||
credentials: Final[Credentials] = self.get_credentials(
|
||||
aws_access_key_id=self.s3_aws_access_key_id,
|
||||
aws_secret_access_key=self.s3_aws_secret_access_key,
|
||||
aws_session_token=self.s3_aws_session_token,
|
||||
aws_region_name=self.s3_region_name,
|
||||
)
|
||||
|
||||
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
|
||||
|
||||
|
|
@ -516,22 +526,24 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
**self._sse_headers(),
|
||||
}
|
||||
|
||||
# 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)
|
||||
|
||||
# Prepare the signed headers
|
||||
signed_headers: Final = dict(aws_request.headers.items())
|
||||
|
||||
httpx_client: Final = _get_httpx_client(
|
||||
params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None)
|
||||
)
|
||||
# Make the request with retry for transient S3 errors (500/503)
|
||||
|
||||
def signed_put() -> httpx.Response:
|
||||
credentials: Final = self.get_credentials(
|
||||
aws_access_key_id=self.s3_aws_access_key_id,
|
||||
aws_secret_access_key=self.s3_aws_secret_access_key,
|
||||
aws_session_token=self.s3_aws_session_token,
|
||||
aws_region_name=self.s3_region_name,
|
||||
)
|
||||
signed_headers: Final = self._sign_put(credentials, url, json_string, headers)
|
||||
return httpx_client.put(url, data=json_string, headers=signed_headers)
|
||||
|
||||
max_retries: Final = 3
|
||||
for attempt in range(max_retries):
|
||||
response = httpx_client.put(url, data=json_string, headers=signed_headers)
|
||||
if response.status_code in (500, 503) and attempt < max_retries - 1:
|
||||
response = signed_put()
|
||||
if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
|
||||
wait_time = 2**attempt # 1s, 2s
|
||||
verbose_logger.warning(
|
||||
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
|
||||
|
|
@ -597,7 +609,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())
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -441,7 +441,15 @@ class LoggingCallbackManager:
|
|||
|
||||
return result
|
||||
|
||||
def get_callback_objects(self) -> tuple[tuple[str, CustomLogger | Callable], ...]:
|
||||
return tuple(
|
||||
(self._get_callback_string(callback), callback)
|
||||
for callback in self._get_all_callbacks()
|
||||
if not isinstance(callback, str)
|
||||
)
|
||||
|
||||
def _get_callback_string(self, callback: CustomLogger | Callable | str) -> str:
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
from litellm.litellm_core_utils.custom_logger_registry import (
|
||||
CustomLoggerRegistry,
|
||||
)
|
||||
|
|
@ -449,6 +457,8 @@ class LoggingCallbackManager:
|
|||
"""Convert a callback to its string representation"""
|
||||
if isinstance(callback, str):
|
||||
return callback
|
||||
elif isinstance(callback, OpenTelemetry) and callback.callback_name is not None:
|
||||
return callback.callback_name
|
||||
elif isinstance(callback, CustomLogger):
|
||||
# Try to get the string representation from the registry
|
||||
callback_str: Final = CustomLoggerRegistry.get_callback_str_from_class_type(type(callback))
|
||||
|
|
|
|||
|
|
@ -6,8 +6,8 @@ import io
|
|||
import json
|
||||
import mimetypes
|
||||
import re
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from itertools import groupby
|
||||
from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from itertools import groupby, islice
|
||||
from os import PathLike
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
|
|
@ -1320,17 +1320,128 @@ def flatten_top_level_schema_combinators(schema: Mapping[str, object]) -> Mappin
|
|||
return _flatten_schema_against_root(schema, schema, frozenset(), 0, {}) # mutable-ok: fresh per-call $ref memo
|
||||
|
||||
|
||||
def tool_with_flattened_parameters(tool: Mapping[str, object]) -> Mapping[str, object]:
|
||||
_SUBSCHEMA_KEYWORDS: Final = frozenset(
|
||||
{
|
||||
"additionalItems",
|
||||
"additionalProperties",
|
||||
"contains",
|
||||
"else",
|
||||
"if",
|
||||
"items",
|
||||
"not",
|
||||
"propertyNames",
|
||||
"then",
|
||||
"unevaluatedItems",
|
||||
"unevaluatedProperties",
|
||||
}
|
||||
)
|
||||
_SUBSCHEMA_LIST_KEYWORDS: Final = frozenset({"allOf", "anyOf", "items", "oneOf", "prefixItems"})
|
||||
_SUBSCHEMA_MAP_KEYWORDS: Final = frozenset(
|
||||
{"$defs", "definitions", "dependentSchemas", "patternProperties", "properties"}
|
||||
)
|
||||
|
||||
_MAX_SCHEMA_NESTING: Final = 1024
|
||||
|
||||
|
||||
def drop_non_python_regex_patterns(schema: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""Drop every regex in a schema position that Python's ``re`` cannot compile.
|
||||
|
||||
OpenAI validates tool ``parameters`` against the 2020-12 metaschema with
|
||||
``jsonschema``'s format checker, which hands each ``pattern`` value and each
|
||||
``patternProperties`` key to ``re.compile``, so a regex written for an
|
||||
ECMA-262 engine (Unicode property escapes such as ``\\p{Cc}``, as in Claude
|
||||
Code's ``Artifact`` tool) is refused with "'...' is not a 'regex'" by every
|
||||
model family on both the chat and Responses wires. Only schema positions are
|
||||
walked (properties, items, combinators, ``$defs`` and the other applicators),
|
||||
so a ``pattern`` key inside ``default``, ``examples``, ``const`` or vendor
|
||||
extensions is data and stays. Outside strict mode the keyword is only a
|
||||
hint, so dropping it costs the model a constraint and the caller nothing.
|
||||
Compilable regexes and everything else pass through, the input is never
|
||||
mutated, and the same object comes back when nothing was dropped. The walk
|
||||
is level-order rather than recursive, rebuilt deepest level first, and stops
|
||||
at more schema levels than a JSON parser admits, so a cyclic schema built in
|
||||
code cannot spin it.
|
||||
"""
|
||||
rebuilt: dict[int, Mapping[str, object]] = {} # mutable-ok: per-call memo of rewritten nodes, deepest level first
|
||||
for level in reversed(tuple(islice(_schema_levels(schema), _MAX_SCHEMA_NESTING))):
|
||||
rebuilt.update(
|
||||
(id(node), rewritten)
|
||||
for node in level
|
||||
if (rewritten := _node_without_non_python_regex(node, rebuilt)) is not node
|
||||
)
|
||||
return rebuilt.get(id(schema), schema)
|
||||
|
||||
|
||||
def _schema_levels(schema: Mapping[str, object]) -> Iterator[tuple[Mapping[str, object], ...]]:
|
||||
frontier: tuple[Mapping[str, object], ...] = (schema,) # rebind-ok: level-order cursor, one level a round
|
||||
while frontier:
|
||||
yield frontier
|
||||
frontier = tuple(child for node in frontier for child in _subschemas(node))
|
||||
|
||||
|
||||
def _subschemas(node: Mapping[str, object]) -> Iterator[Mapping[str, object]]:
|
||||
for key, value in node.items():
|
||||
if key in _SUBSCHEMA_MAP_KEYWORDS and isinstance(value, dict):
|
||||
yield from (sub for sub in value.values() if isinstance(sub, dict))
|
||||
elif key in _SUBSCHEMA_LIST_KEYWORDS and isinstance(value, list):
|
||||
yield from (sub for sub in value if isinstance(sub, dict))
|
||||
elif key in _SUBSCHEMA_KEYWORDS and isinstance(value, dict):
|
||||
yield value
|
||||
|
||||
|
||||
def _node_without_non_python_regex(
|
||||
node: Mapping[str, object], rebuilt: Mapping[int, Mapping[str, object]]
|
||||
) -> Mapping[str, object]:
|
||||
kept: Final = { # mutable-ok: tool parameters are JSON dicts
|
||||
key: _keyword_value_rebuilt(key, value, rebuilt)
|
||||
for key, value in node.items()
|
||||
if key != "pattern" or not isinstance(value, str) or _is_python_regex(value)
|
||||
}
|
||||
return node if len(kept) == len(node) and all(kept[key] is node[key] for key in kept) else kept
|
||||
|
||||
|
||||
def _keyword_value_rebuilt(key: str, value: object, rebuilt: Mapping[int, Mapping[str, object]]) -> object:
|
||||
if key in _SUBSCHEMA_MAP_KEYWORDS and isinstance(value, dict):
|
||||
kept: Final = { # mutable-ok: tool parameters are JSON dicts
|
||||
name: rebuilt.get(id(sub), sub)
|
||||
for name, sub in value.items()
|
||||
if key != "patternProperties" or not isinstance(name, str) or _is_python_regex(name)
|
||||
}
|
||||
return value if len(kept) == len(value) and all(kept[name] is value[name] for name in kept) else kept
|
||||
if key in _SUBSCHEMA_LIST_KEYWORDS and isinstance(value, list):
|
||||
items: Final = [rebuilt.get(id(sub), sub) for sub in value] # mutable-ok: tool parameters are JSON lists
|
||||
return value if all(new is old for new, old in zip(items, value, strict=True)) else items
|
||||
if key in _SUBSCHEMA_KEYWORDS and isinstance(value, dict):
|
||||
return rebuilt.get(id(value), value)
|
||||
return value
|
||||
|
||||
|
||||
def _is_python_regex(pattern: str) -> bool:
|
||||
try:
|
||||
re.compile(pattern)
|
||||
except (re.error, RecursionError):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def flatten_combinators_and_drop_non_python_regex_patterns(schema: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return flatten_top_level_schema_combinators(drop_non_python_regex_patterns(schema))
|
||||
|
||||
|
||||
def tool_with_sanitized_parameters(
|
||||
tool: Mapping[str, object],
|
||||
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
|
||||
) -> Mapping[str, object]:
|
||||
function: Final = tool.get("function")
|
||||
if not isinstance(function, dict):
|
||||
return tool
|
||||
parameters: Final = function.get("parameters")
|
||||
if not isinstance(parameters, dict):
|
||||
return tool
|
||||
flattened: Final = flatten_top_level_schema_combinators(parameters)
|
||||
if flattened is parameters:
|
||||
sanitized: Final = sanitize(parameters)
|
||||
if sanitized is parameters:
|
||||
return tool
|
||||
return {**tool, "function": {**function, "parameters": flattened}} # mutable-ok: request tools are JSON dicts
|
||||
return {**tool, "function": {**function, "parameters": sanitized}} # mutable-ok: request tools are JSON dicts
|
||||
|
||||
|
||||
def _get_image_mime_type_from_url(url: str) -> str | None:
|
||||
|
|
@ -1823,14 +1934,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 +1951,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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -72,8 +73,13 @@ _CLAUDE_CODE_OBJECT_MAPPING_ADAPTER: Final = TypeAdapter(dict[object, object])
|
|||
_CLAUDE_CODE_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object])
|
||||
|
||||
|
||||
_CLAUDE_CODE_USER_AGENT_PREFIXES: Final = ("claude-cli/", "claude-code/")
|
||||
|
||||
|
||||
def is_claude_code_user_agent(user_agent: str) -> bool:
|
||||
return user_agent.startswith("claude-cli/")
|
||||
"""Claude Code sends its API calls through the Anthropic SDK as `claude-cli/<version>` and its own
|
||||
fetches, such as gateway model discovery, as `claude-code/<version>`"""
|
||||
return user_agent.startswith(_CLAUDE_CODE_USER_AGENT_PREFIXES)
|
||||
|
||||
|
||||
def _validated_claude_code_mapping(value: object) -> dict[object, object] | None:
|
||||
|
|
@ -1201,6 +1207,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:
|
||||
|
|
@ -1629,11 +1661,16 @@ def process_anthropic_headers(headers: httpx.Headers | dict) -> dict:
|
|||
|
||||
|
||||
def _anthropic_model_entry(
|
||||
model: ModelInfoResponse, created_at: str, display_names: Mapping[str, str]
|
||||
model: ModelInfoResponse, created_at: str, display_names: Mapping[str, str], listed_ids: Mapping[str, str]
|
||||
) -> Mapping[str, object]:
|
||||
listed_id: Final = listed_ids.get(model["id"])
|
||||
source: Final[Mapping[str, object]] = (
|
||||
MappingProxyType({"source_model": model["id"]}) if listed_id is not None else MappingProxyType({})
|
||||
)
|
||||
return { # mutable-ok: JSON response body, serialized by the route and never mutated
|
||||
"type": "model",
|
||||
"id": model["id"],
|
||||
"id": listed_id or model["id"],
|
||||
**source,
|
||||
"display_name": display_names.get(model["id"], model["id"]),
|
||||
"created_at": created_at,
|
||||
"max_input_tokens": model.get("max_input_tokens"),
|
||||
|
|
@ -1644,6 +1681,7 @@ def _anthropic_model_entry(
|
|||
def create_anthropic_model_list_response(
|
||||
models: Sequence[ModelInfoResponse],
|
||||
display_names: Mapping[str, str] = MappingProxyType({}),
|
||||
listed_ids: Mapping[str, str] = MappingProxyType({}),
|
||||
) -> Mapping[str, object]:
|
||||
"""Build the Anthropic-native /v1/models envelope.
|
||||
|
||||
|
|
@ -1653,17 +1691,19 @@ def create_anthropic_model_list_response(
|
|||
over from the OpenAI-shaped listing, named as the Messages API names them, and
|
||||
are always present because the vendor shape declares them nullable, not optional.
|
||||
display_names maps a listed model id to a configured human-readable name; ids
|
||||
without an entry fall back to the id itself, matching the vendor behavior
|
||||
without an entry fall back to the id itself, matching the vendor behavior.
|
||||
listed_ids maps a model id to the id the caller should see it under (the Claude
|
||||
Code view); ids without an entry are listed as they are
|
||||
"""
|
||||
created_at: Final = (
|
||||
datetime.fromtimestamp(DEFAULT_MODEL_CREATED_AT_TIME, tz=timezone.utc).isoformat().replace("+00:00", "Z")
|
||||
)
|
||||
data: Final = [ # mutable-ok: JSON response body, serialized by the route and never mutated
|
||||
_anthropic_model_entry(model, created_at, display_names) for model in models
|
||||
_anthropic_model_entry(model, created_at, display_names, listed_ids) for model in models
|
||||
]
|
||||
return { # mutable-ok: JSON response body, serialized by the route and never mutated
|
||||
"data": data,
|
||||
"has_more": False,
|
||||
"first_id": models[0]["id"] if models else None,
|
||||
"last_id": models[-1]["id"] if models else None,
|
||||
"first_id": data[0]["id"] if data else None,
|
||||
"last_id": data[-1]["id"] if data else None,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", "{}"))
|
||||
|
|
|
|||
|
|
@ -7,8 +7,9 @@ from httpx._models import Headers, Response
|
|||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
flatten_combinators_and_drop_non_python_regex_patterns,
|
||||
hoist_images_from_tool_messages,
|
||||
tool_with_flattened_parameters,
|
||||
tool_with_sanitized_parameters,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
convert_to_azure_openai_messages,
|
||||
|
|
@ -39,14 +40,17 @@ else:
|
|||
_NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def flattened_tools_update(optional_params: Mapping[str, object]) -> Mapping[str, object]:
|
||||
def sanitized_tools_update(optional_params: Mapping[str, object]) -> Mapping[str, object]:
|
||||
tools: Final = optional_params.get("tools")
|
||||
if not isinstance(tools, list):
|
||||
return _NO_TOOLS_UPDATE
|
||||
flattened: Final = [ # mutable-ok: request tools are a JSON list
|
||||
tool_with_flattened_parameters(tool) if isinstance(tool, dict) else tool for tool in tools
|
||||
sanitized: Final = [ # mutable-ok: request tools are a JSON list
|
||||
tool_with_sanitized_parameters(tool, flatten_combinators_and_drop_non_python_regex_patterns)
|
||||
if isinstance(tool, dict)
|
||||
else tool
|
||||
for tool in tools
|
||||
]
|
||||
return MappingProxyType({"tools": flattened})
|
||||
return MappingProxyType({"tools": sanitized})
|
||||
|
||||
|
||||
class AzureOpenAIConfig(BaseConfig):
|
||||
|
|
@ -278,7 +282,7 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
"model": model,
|
||||
"messages": azure_messages,
|
||||
**optional_params,
|
||||
**flattened_tools_update(optional_params),
|
||||
**sanitized_tools_update(optional_params),
|
||||
}
|
||||
|
||||
def transform_response(
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.utils import get_model_info, supports_reasoning
|
||||
|
||||
from ...openai.chat.o_series_transformation import OpenAIOSeriesConfig
|
||||
from .gpt_transformation import flattened_tools_update
|
||||
from .gpt_transformation import sanitized_tools_update
|
||||
|
||||
|
||||
class AzureOpenAIO1Config(OpenAIOSeriesConfig):
|
||||
|
|
@ -111,6 +111,6 @@ class AzureOpenAIO1Config(OpenAIOSeriesConfig):
|
|||
model = model.replace("o_series/", "") # handle o_series/my-random-deployment-name
|
||||
flattened_params: Final = { # mutable-ok: transform_request's contract takes a plain JSON params dict
|
||||
**optional_params,
|
||||
**flattened_tools_update(optional_params),
|
||||
**sanitized_tools_update(optional_params),
|
||||
}
|
||||
return super().transform_request(model, messages, flattened_params, litellm_params, headers)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
232
litellm/llms/azure_ai/passthrough/transformation.py
Normal file
232
litellm/llms/azure_ai/passthrough/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -427,13 +427,13 @@ class AmazonConverseConfig(BaseConfig):
|
|||
"""
|
||||
Handle the reasoning_effort parameter based on the model type.
|
||||
|
||||
- GPT-OSS models: passed through unchanged 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 self._is_openai_gpt_reasoning_model(model):
|
||||
reasoning: Final[BedrockConverseGptReasoningEffortBlock] = {"effort": reasoning_effort}
|
||||
|
|
@ -514,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
|
||||
|
||||
|
|
@ -575,21 +605,14 @@ class AmazonConverseConfig(BaseConfig):
|
|||
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")
|
||||
|
|
@ -881,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):
|
||||
|
|
@ -929,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 not self._is_openai_gpt_reasoning_model(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"
|
||||
|
|
@ -955,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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -9867,7 +9883,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,
|
||||
|
|
|
|||
|
|
@ -19,10 +19,12 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
_should_convert_tool_call_to_json_mode,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_non_python_regex_patterns,
|
||||
drop_tool_reference_parts_from_tool_messages,
|
||||
flatten_combinators_and_drop_non_python_regex_patterns,
|
||||
get_tool_call_names,
|
||||
hoist_images_from_tool_messages,
|
||||
tool_with_flattened_parameters,
|
||||
tool_with_sanitized_parameters,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
async_convert_url_to_base64,
|
||||
|
|
@ -432,7 +434,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
custom_llm_provider, api_base
|
||||
)
|
||||
|
||||
def _flattened_tools_update_for_openai(
|
||||
def _sanitized_tools_update_for_openai(
|
||||
self,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
|
|
@ -440,22 +442,26 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
"""
|
||||
OpenAI's chat completions validator rejects tool `parameters` carrying
|
||||
'oneOf'/'anyOf'/'allOf'/'enum'/'const'/'not' at the top level for every
|
||||
model family, unlike the Responses API, where GPT-5+ accepts them.
|
||||
model family, unlike the Responses API, where GPT-5+ accepts them, and
|
||||
a `pattern` Python's `re` cannot compile for every model family on both.
|
||||
A custom api_base on the `openai` provider is usually a proxy in front of
|
||||
the same validator, so regexes are dropped there too, while the lossier
|
||||
combinator flattening stays limited to api.openai.com hosts.
|
||||
"""
|
||||
tools: Final = optional_params.get("tools")
|
||||
if not isinstance(tools, list):
|
||||
return _NO_TOOLS_UPDATE
|
||||
provider: Final = litellm_params.get("custom_llm_provider")
|
||||
raw_api_base: Final = litellm_params.get("api_base")
|
||||
if not self._targets_openai_hosted_endpoint(
|
||||
provider if isinstance(provider, str) else None,
|
||||
raw_api_base if isinstance(raw_api_base, str) else None,
|
||||
):
|
||||
if not isinstance(tools, list) or provider != "openai":
|
||||
return _NO_TOOLS_UPDATE
|
||||
flattened: Final = [ # mutable-ok: request tools are a JSON list
|
||||
tool_with_flattened_parameters(tool) if isinstance(tool, dict) else tool for tool in tools
|
||||
raw_api_base: Final = litellm_params.get("api_base")
|
||||
sanitize: Final = (
|
||||
flatten_combinators_and_drop_non_python_regex_patterns
|
||||
if self._targets_openai_hosted_endpoint(provider, raw_api_base if isinstance(raw_api_base, str) else None)
|
||||
else drop_non_python_regex_patterns
|
||||
)
|
||||
sanitized: Final = [ # mutable-ok: request tools are a JSON list
|
||||
tool_with_sanitized_parameters(tool, sanitize) if isinstance(tool, dict) else tool for tool in tools
|
||||
]
|
||||
return MappingProxyType({"tools": flattened})
|
||||
return MappingProxyType({"tools": sanitized})
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
|
|
@ -489,7 +495,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
"model": model,
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
**self._flattened_tools_update_for_openai(optional_params, litellm_params),
|
||||
**self._sanitized_tools_update_for_openai(optional_params, litellm_params),
|
||||
}
|
||||
|
||||
async def async_transform_request(
|
||||
|
|
@ -521,7 +527,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
"model": model,
|
||||
"messages": transformed_messages,
|
||||
**optional_params,
|
||||
**self._flattened_tools_update_for_openai(optional_params, litellm_params),
|
||||
**self._sanitized_tools_update_for_openai(optional_params, litellm_params),
|
||||
}
|
||||
else:
|
||||
## allow for any object specific behaviour to be handled
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_type_hints
|
||||
|
|
@ -15,6 +15,10 @@ from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
|
|||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_safe_convert_created_field,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
drop_non_python_regex_patterns,
|
||||
flatten_combinators_and_drop_non_python_regex_patterns,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
|
|
@ -40,7 +44,7 @@ else:
|
|||
|
||||
_NO_TOOL_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_MODEL_FAMILIES_REJECTING_TOP_LEVEL_SCHEMA_COMBINATORS: Final = ("gpt-4", "gpt-3.5", "chatgpt-4o", "o1", "o3", "o4")
|
||||
_PROVIDERS_WITH_COMBINATOR_REJECTING_VALIDATOR: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
|
||||
_PROVIDERS_WITH_OPENAI_SCHEMA_VALIDATOR: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
|
||||
_PROVIDERS_VALIDATING_TOOL_CALL_ITEM_IDS: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
|
||||
|
||||
|
||||
|
|
@ -293,7 +297,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
model=model, input=validated_input, tools=tools
|
||||
)
|
||||
object_schema_tools: Final = self._tools_with_object_parameters(model=model, tools=stripped_tools)
|
||||
sanitized_tools: Final = self._flatten_tool_schema_combinators_for_openai(
|
||||
sanitized_tools: Final = self._sanitized_tool_schemas_for_openai(
|
||||
model=model, tools=object_schema_tools, litellm_params=litellm_params
|
||||
)
|
||||
return self._drop_foreign_tool_call_item_ids(stripped_input), sanitized_tools
|
||||
|
|
@ -378,35 +382,35 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
return item
|
||||
return {key: value for key, value in item.items() if key != "id"} # mutable-ok: outgoing JSON request item
|
||||
|
||||
def _flatten_tool_schema_combinators_for_openai(
|
||||
def _sanitized_tool_schemas_for_openai(
|
||||
self,
|
||||
model: str,
|
||||
tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None:
|
||||
"""Flatten top-level schema combinators only where OpenAI's validator rejects them.
|
||||
"""Rewrite tool schemas only where OpenAI's validator rejects them.
|
||||
|
||||
OpenAI-compatible backends reusing this config (and the ChatGPT backend
|
||||
Codex talks to natively) accept them, and so do GPT-5 and later models,
|
||||
which also call tools better with the union intact. Codex wraps MCP tools
|
||||
inside namespace entries, so nested ``tools`` arrays are walked too.
|
||||
Azure OpenAI shares the validator but names deployments arbitrarily, so
|
||||
the router's declared ``model_info.base_model`` wins over the deployment
|
||||
name and an unrecognized name without one is left untouched.
|
||||
Every model family refuses a ``pattern`` Python's ``re`` cannot compile,
|
||||
while top-level schema combinators are flattened only for the families
|
||||
whose validator rejects them: OpenAI-compatible backends reusing this
|
||||
config (and the ChatGPT backend Codex talks to natively) accept them,
|
||||
and so do GPT-5 and later models, which also call tools better with the
|
||||
union intact. Codex wraps MCP tools inside namespace entries, so nested
|
||||
``tools`` arrays are walked too. Azure OpenAI shares the validator but
|
||||
names deployments arbitrarily, so the router's declared
|
||||
``model_info.base_model`` wins over the deployment name and an
|
||||
unrecognized name without one keeps its combinators.
|
||||
"""
|
||||
if tools is None or self.custom_llm_provider not in _PROVIDERS_WITH_COMBINATOR_REJECTING_VALIDATOR:
|
||||
if tools is None or self.custom_llm_provider not in _PROVIDERS_WITH_OPENAI_SCHEMA_VALIDATOR:
|
||||
return tools
|
||||
gate_model: Final = self._combinator_gate_model(model=model, litellm_params=litellm_params)
|
||||
if not self._rejects_top_level_schema_combinators(gate_model):
|
||||
return tools
|
||||
flattened: Final = [ # mutable-ok: request tools are a JSON list
|
||||
self._flattened_tool_or_passthrough(tool) for tool in tools
|
||||
]
|
||||
return cast("Sequence[ALL_RESPONSES_API_TOOL_PARAMS]", flattened) # cast-ok: spread keeps each tool's shape
|
||||
|
||||
@staticmethod
|
||||
def _flattened_tool_or_passthrough(tool: object) -> object:
|
||||
return OpenAIResponsesAPIConfig._flattened_tool_entry(tool) if isinstance(tool, dict) else tool
|
||||
sanitize: Final = (
|
||||
flatten_combinators_and_drop_non_python_regex_patterns
|
||||
if self._rejects_top_level_schema_combinators(gate_model)
|
||||
else drop_non_python_regex_patterns
|
||||
)
|
||||
sanitized: Final = self._sanitized_tools(tools, sanitize)
|
||||
return cast("Sequence[ALL_RESPONSES_API_TOOL_PARAMS]", sanitized) # cast-ok: spread keeps each tool's shape
|
||||
|
||||
@staticmethod
|
||||
def _rejects_top_level_schema_combinators(model: str) -> bool:
|
||||
|
|
@ -421,35 +425,42 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
return base_model if isinstance(base_model, str) and base_model else model
|
||||
|
||||
@staticmethod
|
||||
def _flattened_tool_entry(
|
||||
def _sanitized_tool_entry(
|
||||
entry: Mapping[str, object],
|
||||
) -> dict[str, object]: # mutable-ok: request tools are JSON dicts
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
flatten_top_level_schema_combinators,
|
||||
)
|
||||
|
||||
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
|
||||
) -> Mapping[str, object]:
|
||||
parameters: Final = entry.get("parameters")
|
||||
nested_tools: Final = entry.get("tools")
|
||||
sanitized_parameters: Final = sanitize(parameters) if isinstance(parameters, dict) else parameters
|
||||
sanitized_nested_tools: Final = (
|
||||
OpenAIResponsesAPIConfig._sanitized_tools(nested_tools, sanitize)
|
||||
if isinstance(nested_tools, list)
|
||||
else nested_tools
|
||||
)
|
||||
parameters_update: Final = (
|
||||
MappingProxyType({"parameters": flatten_top_level_schema_combinators(parameters)})
|
||||
if isinstance(parameters, dict)
|
||||
MappingProxyType({"parameters": sanitized_parameters})
|
||||
if sanitized_parameters is not parameters
|
||||
else _NO_TOOL_UPDATE
|
||||
)
|
||||
tools_update: Final = (
|
||||
MappingProxyType({"tools": OpenAIResponsesAPIConfig._flattened_nested_tools(nested_tools)})
|
||||
if isinstance(nested_tools, list)
|
||||
MappingProxyType({"tools": sanitized_nested_tools})
|
||||
if sanitized_nested_tools is not nested_tools
|
||||
else _NO_TOOL_UPDATE
|
||||
)
|
||||
if not parameters_update and not tools_update:
|
||||
return entry
|
||||
return {**entry, **parameters_update, **tools_update} # mutable-ok: request tools are JSON dicts
|
||||
|
||||
@staticmethod
|
||||
def _flattened_nested_tools(
|
||||
nested_tools: Sequence[object],
|
||||
) -> list[object]: # mutable-ok: namespace tools are a JSON list
|
||||
return [ # mutable-ok: namespace tools are a JSON list
|
||||
OpenAIResponsesAPIConfig._flattened_tool_entry(item) if isinstance(item, dict) else item
|
||||
for item in nested_tools
|
||||
def _sanitized_tools(
|
||||
tools: Sequence[object],
|
||||
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
|
||||
) -> Sequence[object]:
|
||||
sanitized: Final = [ # mutable-ok: request tools are a JSON list
|
||||
OpenAIResponsesAPIConfig._sanitized_tool_entry(item, sanitize) if isinstance(item, dict) else item
|
||||
for item in tools
|
||||
]
|
||||
return tools if all(new is old for new, old in zip(sanitized, tools, strict=True)) else sanitized
|
||||
|
||||
def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam:
|
||||
"""
|
||||
|
|
@ -620,15 +631,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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -2310,8 +2310,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
|
||||
|
|
@ -2401,11 +2400,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"],
|
||||
|
|
@ -2413,6 +2407,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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
@ -729,11 +729,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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -94,6 +94,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
team_membership_auth_cache_key,
|
||||
team_membership_reservation_cache_key,
|
||||
)
|
||||
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.guardrails.tool_name_extraction import (
|
||||
TOOL_CAPABLE_CALL_TYPES,
|
||||
|
|
@ -3450,36 +3451,37 @@ async def _fetch_key_object_from_db_with_reconnect(
|
|||
"""
|
||||
Fetch key object from DB and retry once if a DB connection error can be healed.
|
||||
"""
|
||||
try:
|
||||
return await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
if PrismaDBExceptionHandler.is_database_transport_error(e):
|
||||
did_reconnect = False
|
||||
if hasattr(prisma_client, "attempt_db_reconnect"):
|
||||
auth_reconnect_timeout = getattr(prisma_client, "_db_auth_reconnect_timeout_seconds", 2.0)
|
||||
if not isinstance(auth_reconnect_timeout, (int, float)):
|
||||
auth_reconnect_timeout = 2.0
|
||||
auth_reconnect_lock_timeout = getattr(prisma_client, "_db_auth_reconnect_lock_timeout_seconds", 0.1)
|
||||
if not isinstance(auth_reconnect_lock_timeout, (int, float)):
|
||||
auth_reconnect_lock_timeout = 0.1
|
||||
did_reconnect = await prisma_client.attempt_db_reconnect(
|
||||
reason="auth_get_key_object_lookup_failure",
|
||||
timeout_seconds=auth_reconnect_timeout,
|
||||
lock_timeout_seconds=auth_reconnect_lock_timeout,
|
||||
)
|
||||
if did_reconnect:
|
||||
return await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
raise
|
||||
async with db_lookup_gate.current():
|
||||
try:
|
||||
return await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
if PrismaDBExceptionHandler.is_database_transport_error(e):
|
||||
did_reconnect = False
|
||||
if hasattr(prisma_client, "attempt_db_reconnect"):
|
||||
auth_reconnect_timeout = getattr(prisma_client, "_db_auth_reconnect_timeout_seconds", 2.0)
|
||||
if not isinstance(auth_reconnect_timeout, (int, float)):
|
||||
auth_reconnect_timeout = 2.0
|
||||
auth_reconnect_lock_timeout = getattr(prisma_client, "_db_auth_reconnect_lock_timeout_seconds", 0.1)
|
||||
if not isinstance(auth_reconnect_lock_timeout, (int, float)):
|
||||
auth_reconnect_lock_timeout = 0.1
|
||||
did_reconnect = await prisma_client.attempt_db_reconnect(
|
||||
reason="auth_get_key_object_lookup_failure",
|
||||
timeout_seconds=auth_reconnect_timeout,
|
||||
lock_timeout_seconds=auth_reconnect_lock_timeout,
|
||||
)
|
||||
if did_reconnect:
|
||||
return await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def jwt_key_mapping_cache_key(jwt_claim_name: str, jwt_claim_value: str) -> str:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -92,6 +92,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_set_request_parsed_body,
|
||||
populate_request_with_path_params,
|
||||
)
|
||||
from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group
|
||||
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
|
|
@ -183,6 +184,44 @@ def _get_model_from_request_context(
|
|||
)
|
||||
|
||||
|
||||
_CLAUDE_MODEL_ROUTES: Final = frozenset(
|
||||
f"/{prefix}{endpoint}" for prefix in ("", "v1/") for endpoint in ("messages", "chat/completions", "responses")
|
||||
)
|
||||
_CLAUDE_MODEL_NORMALIZED: Final = "litellm.claude_model_normalized"
|
||||
|
||||
|
||||
async def _normalize_claude_model(
|
||||
request_data: dict, valid_token: UserAPIKeyAuth, request: Request | None, route: str
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_config, proxy_logging_obj
|
||||
|
||||
if route not in _CLAUDE_MODEL_ROUTES or llm_router is None:
|
||||
return
|
||||
if request is not None and request.scope.get(_CLAUDE_MODEL_NORMALIZED) is True:
|
||||
return
|
||||
requested: Final = _get_model_from_request_context(request_data, route, request, llm_router)
|
||||
if not isinstance(requested, str) or requested != request_data.get("model"):
|
||||
return
|
||||
if not requested.startswith("claude-router-") and not requested.lower().endswith("[1m]"):
|
||||
return
|
||||
settings: Final = await proxy_config.get_hierarchical_router_settings(
|
||||
user_api_key_dict=valid_token, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
aliases: Final = settings.get("model_group_alias") if isinstance(settings, Mapping) else None
|
||||
source: Final = claude_code_requested_group(
|
||||
requested, llm_router, valid_token.team_id, (valid_token.aliases, valid_token.team_model_aliases, aliases)
|
||||
)
|
||||
if request is not None:
|
||||
request.scope[_CLAUDE_MODEL_NORMALIZED] = True
|
||||
if source is None:
|
||||
return
|
||||
request_data["model"] = source
|
||||
_safe_set_request_parsed_body(request=request, parsed_body=request_data)
|
||||
if request is not None:
|
||||
request._json = request_data
|
||||
request._body = orjson.dumps(request_data)
|
||||
|
||||
|
||||
def _get_model_names_for_budget_checks(
|
||||
model: str | list[str] | None,
|
||||
) -> list[str]:
|
||||
|
|
@ -2793,6 +2832,7 @@ async def _authorize_authenticated_request(
|
|||
"""
|
||||
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
|
||||
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj, request=request)
|
||||
await _normalize_claude_model(request_data, user_api_key_auth_obj, request, route)
|
||||
|
||||
# Single authorization point. Builder paths MUST NOT call common_checks.
|
||||
# Route through the same exception handler the builder uses so
|
||||
|
|
@ -3156,6 +3196,7 @@ async def _enforce_key_and_fallback_model_access(
|
|||
Key-level model allowlist and client fallbacks (same as standard auth).
|
||||
Not included in common_checks — common_checks enforces team/user/project model access only.
|
||||
"""
|
||||
await _normalize_claude_model(request_data, valid_token, request, route)
|
||||
config: Final = valid_token.config
|
||||
|
||||
if config != {}:
|
||||
|
|
|
|||
|
|
@ -490,7 +490,7 @@ lite codex exec "summarize the repo"
|
|||
|
||||
Each command resolves your LiteLLM key (logging in via SSO when none is stored and you are at a terminal; otherwise it expects `LITELLM_PROXY_API_KEY` or `--api-key`), checks the key against the proxy so bad credentials fail immediately instead of deep inside the agent, exports the environment variables the agent reads, then replaces itself with the agent process.
|
||||
|
||||
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. It also gets `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` (again unless you already set it) so Claude Code v2.1.129+ fills its `/model` picker from the proxy's `/v1/models`; Claude Code only lists entries whose id contains `claude` or `anthropic`, and older versions ignore the variable. Export it as `0` to turn discovery off. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol). OpenCode additionally gets `OPENCODE_CONFIG_CONTENT` holding a generated `litellm` provider (`@ai-sdk/openai-compatible`, the proxy `/v1` URL, `{env:OPENAI_API_KEY}`) with one model entry per chat model your key can see on `/v1/models`, so its model picker mirrors the proxy without a hand-maintained `opencode.json`; OpenCode merges that over your own config files, and if you already export `OPENCODE_CONFIG_CONTENT` yours is left alone. When the list cannot be fetched, `lite opencode` says so on stderr and launches anyway.
|
||||
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. It also gets `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` (again unless you already set it) so Claude Code v2.1.129+ fills its `/model` picker from the proxy's `/v1/models`; Claude Code only lists entries whose id contains `claude` or `anthropic`, so the proxy lists every other group to Claude Code as `claude-router-<UTF-8 hex of the group name>` and marks a group whose input window reaches 1M with `[1m]`, and a request on such an id is served by the group. Older Claude Code versions ignore the variable. Export it as `0` to turn discovery off. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol). OpenCode additionally gets `OPENCODE_CONFIG_CONTENT` holding a generated `litellm` provider (`@ai-sdk/openai-compatible`, the proxy `/v1` URL, `{env:OPENAI_API_KEY}`) with one model entry per chat model your key can see on `/v1/models`, so its model picker mirrors the proxy without a hand-maintained `opencode.json`; OpenCode merges that over your own config files, and if you already export `OPENCODE_CONFIG_CONTENT` yours is left alone. When the list cannot be fetched, `lite opencode` says so on stderr and launches anyway.
|
||||
|
||||
pi ignores base-URL environment variables entirely, so `lite pi` (kept out of the `lite --help` command listing for now, but fully functional) wires it up differently: before handoff it fetches the models your key can use from the proxy's `/v1/models` (plus each model's context window and output cap from `/model_group/info`, when available) and syncs them into a `litellm` provider entry in pi's `~/.pi/agent/models.json` (honoring `PI_CODING_AGENT_DIR`), then starts pi on that provider's first model via an injected `--model litellm/<id>`. Only that one provider entry is rewritten; the rest of the file, including any other custom providers, is left alone. The entry references the key as `$LITELLM_PROXY_API_KEY`, which the wrapper exports for the session, so the token itself never lands on disk and plain `pi` outside the wrapper simply shows the litellm models as unavailable. Your own flags come after the injected pin, so `lite pi --model litellm/<other-id>` wins, and inside the TUI the `/model` picker lists every synced litellm model.
|
||||
|
||||
|
|
@ -548,7 +548,7 @@ lite --base-url https://your-proxy.example.com configure claude --api-key sk-...
|
|||
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
|
||||
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 (under `claude-router-<UTF-8 hex of the group name>` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) 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 (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. 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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,6 +13,7 @@ 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,
|
||||
|
|
@ -84,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.
|
||||
|
||||
|
|
@ -98,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
|
||||
|
|
@ -463,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.
|
||||
|
||||
|
|
@ -494,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),
|
||||
}
|
||||
)
|
||||
|
|
@ -532,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(
|
||||
|
|
@ -548,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))
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import os
|
||||
import sys
|
||||
import time
|
||||
import webbrowser
|
||||
|
|
@ -40,16 +41,16 @@ from litellm.litellm_core_utils.cli_token_utils import (
|
|||
)
|
||||
|
||||
from .claude_settings import (
|
||||
CLAUDE_SETTINGS_PATH,
|
||||
CONFIGURE_STATE_PATH,
|
||||
SETTINGS_FILE_OWNERS,
|
||||
STARTING_MODEL_ROLE,
|
||||
ApiKeyHelper,
|
||||
ClaudeSettingsError,
|
||||
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,
|
||||
|
|
@ -784,19 +785,20 @@ 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, undoable with `lite unconfigure claude`."""
|
||||
"""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:
|
||||
configure_claude_settings(
|
||||
base_url,
|
||||
ApiKeyHelper(resolve_api_key_helper(base_url)),
|
||||
KeepModel(),
|
||||
CLAUDE_SETTINGS_PATH,
|
||||
CONFIGURE_STATE_PATH,
|
||||
SETTINGS_FILE_OWNERS,
|
||||
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(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}."
|
||||
|
|
@ -870,8 +872,9 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool) -> 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(CLAUDE_SETTINGS_PATH, SETTINGS_FILE_OWNERS)
|
||||
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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ _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"
|
||||
|
|
@ -88,6 +89,33 @@ 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."""
|
||||
|
|
@ -325,6 +353,19 @@ 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 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.
|
||||
|
||||
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.
|
||||
"""
|
||||
try:
|
||||
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))
|
||||
|
||||
|
|
@ -543,6 +584,7 @@ __all__ = (
|
|||
"API_KEY_HELPER_KEY",
|
||||
"AUTOROUTE_BACKUP_PATH",
|
||||
"BACKUP_PATH",
|
||||
"CLAUDE_CONFIG_DIR_ENV",
|
||||
"CLAUDE_SETTINGS_PATH",
|
||||
"CONFIGURE_STATE_PATH",
|
||||
"ENABLE_GATEWAY_MODEL_DISCOVERY_KEY",
|
||||
|
|
@ -569,11 +611,15 @@ __all__ = (
|
|||
"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",
|
||||
"settings_file_owners",
|
||||
"unconfigure_claude_settings",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,20 +1,25 @@
|
|||
"""`lite configure claude` and `lite unconfigure claude`: persistent Claude Code wiring, undoable."""
|
||||
|
||||
import re
|
||||
import os
|
||||
import sys
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import click
|
||||
from InquirerPy import inquirer
|
||||
from InquirerPy.base.control import Choice
|
||||
|
||||
from litellm.proxy.common_utils.model_listing_utils import (
|
||||
CLAUDE_CODE_CLIENT,
|
||||
CLAUDE_CODE_PICKER_PATTERN,
|
||||
GATEWAY_CLIENT_HEADER,
|
||||
)
|
||||
|
||||
from .auth import CliContextObj, context_secret_vault, get_stored_api_key
|
||||
from .claude_settings import (
|
||||
CLAUDE_SETTINGS_PATH,
|
||||
CONFIGURE_STATE_PATH,
|
||||
SETTINGS_FILE_OWNERS,
|
||||
STARTING_MODEL_ROLE,
|
||||
ApiKeyHelper,
|
||||
ClaudeCredential,
|
||||
|
|
@ -24,19 +29,24 @@ from .claude_settings import (
|
|||
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 .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing
|
||||
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)
|
||||
_CLAUDE_CODE_VIEW: Final = MappingProxyType(
|
||||
{"anthropic-version": "2023-06-01", GATEWAY_CLIENT_HEADER: CLAUDE_CODE_CLIENT}
|
||||
)
|
||||
_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 "
|
||||
|
|
@ -64,11 +74,21 @@ def resolve_credential(ctx: click.Context, api_key: str | None) -> tuple[ClaudeC
|
|||
return ApiKeyHelper(resolve_api_key_helper(base_url)), stored
|
||||
|
||||
|
||||
def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, tuple[str, ...]]:
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Listing:
|
||||
models: tuple[ListedModel, ...]
|
||||
|
||||
@property
|
||||
def ids(self) -> tuple[str, ...]:
|
||||
return tuple(model.id for model in self.models)
|
||||
|
||||
|
||||
def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, _Listing]:
|
||||
"""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(CLAUDE_SETTINGS_PATH, SETTINGS_FILE_OWNERS)
|
||||
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))
|
||||
|
|
@ -86,53 +106,69 @@ def _listing_error(base_url: str, error: PiSyncError) -> str:
|
|||
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)
|
||||
def _listed_models(base_url: str, key: str) -> _Listing:
|
||||
listed: Final = fetch_model_listing(base_url, key, headers=_CLAUDE_CODE_VIEW)
|
||||
if isinstance(listed, PiSyncError):
|
||||
raise click.ClickException(_listing_error(base_url, listed))
|
||||
return listed
|
||||
return _Listing(listed)
|
||||
|
||||
|
||||
def _starting_model(model: str, listing: _Listing) -> str | None:
|
||||
source: Final = next((listed.id for listed in listing.models if listed.source_model == model), None)
|
||||
return source or next((listed.id for listed in listing.models if listed.id == model), None)
|
||||
|
||||
|
||||
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:
|
||||
def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listing: _Listing, 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:
|
||||
listed: Final = listing.ids
|
||||
starting: Final = _starting_model(model, listing) if model is not None else None
|
||||
if model is not None and starting is None:
|
||||
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), CLAUDE_SETTINGS_PATH, CONFIGURE_STATE_PATH, SETTINGS_FILE_OWNERS
|
||||
base_url,
|
||||
credential,
|
||||
_model_choice(starting),
|
||||
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: {CLAUDE_SETTINGS_PATH} now routes through {base_url}.")
|
||||
in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.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
|
||||
f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model."
|
||||
if starting 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')."
|
||||
f"/model will list all {len(listed)} of the proxy's models."
|
||||
if in_picker == len(listed)
|
||||
else f"/model will list {in_picker} of the proxy's {len(listed)} models: Claude Code shows only ids containing "
|
||||
"'claude' or 'anthropic', and this proxy does not list the rest under such names."
|
||||
)
|
||||
click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.")
|
||||
if isinstance(credential, StaticToken) and CLAUDE_SETTINGS_PATH.is_symlink():
|
||||
if isinstance(credential, StaticToken) and settings_path.is_symlink():
|
||||
click.echo(
|
||||
f"Note: {CLAUDE_SETTINGS_PATH} is a symlink to {CLAUDE_SETTINGS_PATH.resolve()}, so your key now lives in "
|
||||
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,
|
||||
)
|
||||
|
|
@ -165,8 +201,10 @@ def interactive_configure(
|
|||
targets: Final = pick_targets()
|
||||
if _CLAUDE_TARGET not in targets:
|
||||
return
|
||||
credential, listed = _start(ctx, None)
|
||||
_apply_claude(ctx, credential, listed, pick_model(listed))
|
||||
credential, listing = _start(ctx, None)
|
||||
_apply_claude(
|
||||
ctx, credential, listing, pick_model(tuple(model.source_model or model.id for model in listing.models))
|
||||
)
|
||||
|
||||
|
||||
@click.group(name="configure", invoke_without_command=True)
|
||||
|
|
@ -210,8 +248,8 @@ def configure_claude(ctx: click.Context, api_key: str | None, model: str | None)
|
|||
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)
|
||||
credential, listing = _start(ctx, api_key)
|
||||
_apply_claude(ctx, credential, listing, model)
|
||||
|
||||
|
||||
@unconfigure_group.command(name="claude")
|
||||
|
|
@ -221,11 +259,13 @@ def unconfigure_claude() -> None:
|
|||
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(CLAUDE_SETTINGS_PATH, CONFIGURE_STATE_PATH, SETTINGS_FILE_OWNERS)
|
||||
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(CLAUDE_SETTINGS_PATH, CONFIGURE_STATE_PATH, outcome)
|
||||
_report_unconfigure(settings_path, state_path, outcome)
|
||||
|
||||
|
||||
def _report_unconfigure(settings_path: Path, state_path: Path, outcome: UnconfigureOutcome) -> None:
|
||||
|
|
|
|||
|
|
@ -13,10 +13,11 @@ from dataclasses import dataclass
|
|||
from enum import StrEnum
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Annotated, Final
|
||||
|
||||
import requests
|
||||
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, model_validator
|
||||
from pydantic.types import StringConstraints
|
||||
|
||||
PI_CONFIG_DIR_ENV: Final = "PI_CODING_AGENT_DIR"
|
||||
PI_PROVIDER_NAME: Final = "litellm"
|
||||
|
|
@ -51,12 +52,25 @@ class ModelLimits:
|
|||
max_tokens: int | None
|
||||
|
||||
|
||||
class _Model(BaseModel):
|
||||
id: str
|
||||
_NonEmptyString = Annotated[str, StringConstraints(min_length=1)]
|
||||
|
||||
|
||||
class ListedModel(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
id: _NonEmptyString
|
||||
source_model: _NonEmptyString | None = None
|
||||
|
||||
|
||||
class _ModelList(BaseModel):
|
||||
data: tuple[_Model, ...]
|
||||
data: tuple[ListedModel, ...]
|
||||
|
||||
@model_validator(mode="after")
|
||||
def unique_id_mappings(self) -> "_ModelList":
|
||||
mappings: Final = frozenset((model.id, model.source_model or model.id) for model in self.data)
|
||||
if len(frozenset(model.id for model in self.data)) != len(mappings):
|
||||
raise ValueError("model ids must not map to multiple source models")
|
||||
return self
|
||||
|
||||
|
||||
class _ModelGroup(BaseModel):
|
||||
|
|
@ -69,17 +83,18 @@ class _ModelGroupList(BaseModel):
|
|||
data: tuple[_ModelGroup, ...]
|
||||
|
||||
|
||||
def fetch_model_ids(
|
||||
def fetch_model_listing(
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
*,
|
||||
get: Callable[..., requests.Response] = requests.get,
|
||||
) -> tuple[str, ...] | PiSyncError:
|
||||
headers: Mapping[str, str] = MappingProxyType({}),
|
||||
) -> tuple[ListedModel, ...] | PiSyncError:
|
||||
url: Final = base_url.rstrip("/") + "/v1/models"
|
||||
try:
|
||||
resp: Final = get(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {api_key}"}, # mutable-ok: requests headers require a dict
|
||||
headers={"Authorization": f"Bearer {api_key}", **headers}, # mutable-ok: requests headers require a dict
|
||||
timeout=10,
|
||||
)
|
||||
except requests.RequestException as e:
|
||||
|
|
@ -94,10 +109,21 @@ def fetch_model_ids(
|
|||
listing: Final = _ModelList.model_validate(resp.json())
|
||||
except (ValueError, ValidationError) as 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:
|
||||
models: Final = tuple(dict.fromkeys(listing.data))
|
||||
if not models:
|
||||
return PiSyncError("The proxy returned no models for your key.", kind=ListingFailure.EMPTY)
|
||||
return ids
|
||||
return models
|
||||
|
||||
|
||||
def fetch_model_ids(
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
*,
|
||||
get: Callable[..., requests.Response] = requests.get,
|
||||
headers: Mapping[str, str] = MappingProxyType({}),
|
||||
) -> tuple[str, ...] | PiSyncError:
|
||||
listed: Final = fetch_model_listing(base_url, api_key, get=get, headers=headers)
|
||||
return listed if isinstance(listed, PiSyncError) else tuple(dict.fromkeys(model.id for model in listed))
|
||||
|
||||
|
||||
_NO_LIMITS: Final[Mapping[str, ModelLimits]] = MappingProxyType({})
|
||||
|
|
@ -222,11 +248,13 @@ __all__ = (
|
|||
"LITELLM_PROXY_API_KEY_ENV",
|
||||
"PI_CONFIG_DIR_ENV",
|
||||
"PI_PROVIDER_NAME",
|
||||
"ListedModel",
|
||||
"ListingFailure",
|
||||
"ModelLimits",
|
||||
"PiSyncError",
|
||||
"fetch_model_ids",
|
||||
"fetch_model_limits",
|
||||
"fetch_model_listing",
|
||||
"models_json_path",
|
||||
"provider_block",
|
||||
"sync_models_json",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -10,12 +10,24 @@ legacy internal names with `general_settings.use_team_public_model_name: false`.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
import re
|
||||
from collections.abc import Container, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
||||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
from litellm.types.proxy.model_listing import ModelInfoResponse
|
||||
|
||||
CLAUDE_CODE_PICKER_PATTERN: Final = re.compile(r"claude|anthropic", re.IGNORECASE)
|
||||
GATEWAY_CLIENT_HEADER: Final = "x-gateway-client"
|
||||
CLAUDE_CODE_CLIENT: Final = "claude-code"
|
||||
_CLAUDE_CODE_ALIAS_PREFIX: Final = "claude-router-"
|
||||
_ONE_MILLION_SUFFIX: Final = "[1m]"
|
||||
_ONE_MILLION_TOKENS: Final = 1_000_000
|
||||
|
||||
|
||||
def configured_display_names(
|
||||
|
|
@ -40,6 +52,115 @@ def configured_display_names(
|
|||
)
|
||||
|
||||
|
||||
def _unmarked(name: str) -> str:
|
||||
return name[: -len(_ONE_MILLION_SUFFIX)] if name.lower().endswith(_ONE_MILLION_SUFFIX) else name
|
||||
|
||||
|
||||
def _compatibility_id(model_id: str) -> str:
|
||||
return f"{_CLAUDE_CODE_ALIAS_PREFIX}{model_id.encode().hex()}"
|
||||
|
||||
|
||||
def _decoded_compatibility_id(view_id: str) -> str | None:
|
||||
encoded: Final = _unmarked(view_id).removeprefix(_CLAUDE_CODE_ALIAS_PREFIX)
|
||||
if encoded == _unmarked(view_id):
|
||||
return None
|
||||
try:
|
||||
model_id: Final = bytes.fromhex(encoded).decode()
|
||||
except (ValueError, UnicodeDecodeError):
|
||||
return None
|
||||
return model_id if _compatibility_id(model_id) == _unmarked(view_id) else None
|
||||
|
||||
|
||||
def claude_code_model_id(
|
||||
model_id: str,
|
||||
max_input_tokens: float | None,
|
||||
routing_names: Container[str],
|
||||
) -> str:
|
||||
"""The collision-free id Claude Code's picker lists a model under."""
|
||||
if "*" in model_id:
|
||||
return model_id
|
||||
shaped: Final = model_id if CLAUDE_CODE_PICKER_PATTERN.search(model_id) else _compatibility_id(model_id)
|
||||
one_million: Final = max_input_tokens is not None and max_input_tokens >= _ONE_MILLION_TOKENS
|
||||
marked: Final = (
|
||||
f"{shaped}{_ONE_MILLION_SUFFIX}" if one_million and not shaped.lower().endswith(_ONE_MILLION_SUFFIX) else shaped
|
||||
)
|
||||
return next(
|
||||
(
|
||||
name
|
||||
for name in (marked, shaped)
|
||||
if name == model_id or claude_code_group_name(name, routing_names) == model_id
|
||||
),
|
||||
model_id,
|
||||
)
|
||||
|
||||
|
||||
def claude_code_group_name(view_id: str, routing_names: Container[str]) -> str | None:
|
||||
"""Decode a canonical compatibility id only when no configured route claims it."""
|
||||
if view_id in routing_names:
|
||||
return None
|
||||
unmarked: Final = _unmarked(view_id)
|
||||
if unmarked != view_id and unmarked in routing_names:
|
||||
return unmarked
|
||||
model_id: Final = _decoded_compatibility_id(view_id)
|
||||
return model_id if model_id and model_id in routing_names else None
|
||||
|
||||
|
||||
def is_claude_code_client(headers: Mapping[str, str]) -> bool:
|
||||
"""Claude Code itself, or a client asking for its view of the listing the way Ramp Router's does"""
|
||||
from litellm.llms.anthropic.common_utils import is_claude_code_user_agent
|
||||
|
||||
return (
|
||||
is_claude_code_user_agent(headers.get("user-agent", ""))
|
||||
or headers.get(GATEWAY_CLIENT_HEADER, "").lower() == CLAUDE_CODE_CLIENT
|
||||
)
|
||||
|
||||
|
||||
def claude_code_view_ids(
|
||||
rows: Sequence[ModelInfoResponse],
|
||||
headers: Mapping[str, str],
|
||||
routing_names: Container[str],
|
||||
) -> Mapping[str, str]:
|
||||
"""served id -> Claude Code id for the requested listing view"""
|
||||
if not is_claude_code_client(headers):
|
||||
return MappingProxyType({})
|
||||
return MappingProxyType(
|
||||
{row["id"]: claude_code_model_id(row["id"], row.get("max_input_tokens"), routing_names) for row in rows}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ClaudeCodeRoutingNames:
|
||||
"""Existing routes always own their names, including aliases and wildcard routes."""
|
||||
|
||||
llm_router: Router | None
|
||||
team_id: str | None = None
|
||||
alias_maps: tuple[object, ...] = ()
|
||||
|
||||
def __contains__(self, name: object) -> bool:
|
||||
if not isinstance(name, str):
|
||||
return False
|
||||
if name in litellm.model_alias_map or any(
|
||||
isinstance(aliases, Mapping) and name in aliases for aliases in self.alias_maps
|
||||
):
|
||||
return True
|
||||
if self.llm_router is None:
|
||||
return False
|
||||
return (
|
||||
name in self.llm_router.model_group_alias
|
||||
or self.llm_router.has_model_id(name)
|
||||
or bool(self.llm_router.get_candidate_model_ids_for_route(name, self.team_id))
|
||||
)
|
||||
|
||||
|
||||
def claude_code_requested_group(
|
||||
requested: str,
|
||||
llm_router: Router,
|
||||
team_id: str | None,
|
||||
alias_maps: tuple[object, ...] = (),
|
||||
) -> str | None:
|
||||
return claude_code_group_name(requested, ClaudeCodeRoutingNames(llm_router, team_id, alias_maps))
|
||||
|
||||
|
||||
class TeamModelNameTranslator:
|
||||
"""Translates internal team routing keys to their public names for the model
|
||||
listing/retrieve responses. Stateless; the live router and general_settings
|
||||
|
|
|
|||
23
litellm/proxy/db/db_lookup_gate.py
Normal file
23
litellm/proxy/db/db_lookup_gate.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY
|
||||
|
||||
|
||||
class LoopBoundSemaphore:
|
||||
__slots__ = ("_loop", "_semaphore", "_value")
|
||||
|
||||
def __init__(self, value: int) -> None:
|
||||
self._value: Final = value
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._semaphore: asyncio.Semaphore | None = None
|
||||
|
||||
def current(self) -> asyncio.Semaphore:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
if self._semaphore is None or self._loop is not loop:
|
||||
self._semaphore = asyncio.Semaphore(self._value)
|
||||
self._loop = loop
|
||||
return self._semaphore
|
||||
|
||||
|
||||
db_lookup_gate: Final = LoopBoundSemaphore(PROXY_DB_LOOKUP_MAX_CONCURRENCY)
|
||||
|
|
@ -23,6 +23,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy._types import Litellm_EntityType
|
||||
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
BudgetWindowSpendRepository,
|
||||
|
|
@ -121,30 +122,33 @@ class SpendCounterReseed:
|
|||
if SpendCounterReseed._is_key_or_team_window_counter(counter_key):
|
||||
return None
|
||||
try:
|
||||
if counter_key.startswith("spend:key:"):
|
||||
token: Final = counter_key[len("spend:key:") :]
|
||||
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token})
|
||||
elif counter_key.startswith("spend:team_member:"):
|
||||
suffix: Final = counter_key[len("spend:team_member:") :]
|
||||
if ":" not in suffix:
|
||||
async with db_lookup_gate.current():
|
||||
if counter_key.startswith("spend:key:"):
|
||||
token: Final = counter_key[len("spend:key:") :]
|
||||
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token})
|
||||
elif counter_key.startswith("spend:team_member:"):
|
||||
suffix: Final = counter_key[len("spend:team_member:") :]
|
||||
if ":" not in suffix:
|
||||
return None
|
||||
user_id, team_id = suffix.rsplit(":", 1)
|
||||
row = await TeamMembershipRepository(prisma_client).table.find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}
|
||||
)
|
||||
elif counter_key.startswith("spend:team:"):
|
||||
team_id = counter_key[len("spend:team:") :]
|
||||
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
elif counter_key.startswith("spend:user:"):
|
||||
user_id = counter_key[len("spend:user:") :]
|
||||
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
|
||||
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
|
||||
return None
|
||||
elif counter_key.startswith("spend:org:"):
|
||||
org_id: Final = counter_key[len("spend:org:") :]
|
||||
row = await OrganizationRepository(prisma_client).table.find_unique(
|
||||
where={"organization_id": org_id}
|
||||
)
|
||||
else:
|
||||
return None
|
||||
user_id, team_id = suffix.rsplit(":", 1)
|
||||
row = await TeamMembershipRepository(prisma_client).table.find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}
|
||||
)
|
||||
elif counter_key.startswith("spend:team:"):
|
||||
team_id = counter_key[len("spend:team:") :]
|
||||
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
elif counter_key.startswith("spend:user:"):
|
||||
user_id = counter_key[len("spend:user:") :]
|
||||
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
|
||||
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
|
||||
return None
|
||||
elif counter_key.startswith("spend:org:"):
|
||||
org_id: Final = counter_key[len("spend:org:") :]
|
||||
row = await OrganizationRepository(prisma_client).table.find_unique(where={"organization_id": org_id})
|
||||
else:
|
||||
return None
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("SpendCounterReseed.from_db: failed for %s", counter_key)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -266,13 +268,12 @@ from litellm.constants import (
|
|||
WEEKLY_SPEND_REPORT_JOB_ID,
|
||||
)
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail, ModifyResponseException
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
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,
|
||||
|
|
@ -365,8 +366,11 @@ from litellm.proxy.common_utils.load_config_utils import (
|
|||
)
|
||||
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
|
||||
from litellm.proxy.common_utils.model_listing_utils import (
|
||||
ClaudeCodeRoutingNames,
|
||||
TeamModelNameTranslator,
|
||||
claude_code_view_ids,
|
||||
configured_display_names,
|
||||
is_claude_code_client,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_endpoint_utils import (
|
||||
remove_sensitive_info_from_deployment,
|
||||
|
|
@ -6641,6 +6645,14 @@ class ProxyConfig:
|
|||
return parsed
|
||||
return None
|
||||
|
||||
async def get_hierarchical_router_settings(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> dict | None:
|
||||
return await self._get_hierarchical_router_settings(user_api_key_dict, prisma_client, proxy_logging_obj)
|
||||
|
||||
async def _get_hierarchical_router_settings(
|
||||
self,
|
||||
user_api_key_dict: Optional["UserAPIKeyAuth"],
|
||||
|
|
@ -8634,6 +8646,7 @@ _STREAM_KEEPALIVE: Final = object()
|
|||
_KEEPALIVE_MIN_SECONDS: Final = 1.0
|
||||
_KEEPALIVE_MAX_SECONDS: Final = 300.0
|
||||
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_EMPTY_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
async def _iter_with_keepalive(
|
||||
|
|
@ -10471,6 +10484,24 @@ async def model_list(
|
|||
wants_anthropic_format: Final = (
|
||||
http_request is not None and http_request.headers.get("anthropic-version") is not None
|
||||
)
|
||||
client_headers: Final[Mapping[str, str]] = http_request.headers if http_request is not None else _EMPTY_HEADERS
|
||||
view_router_settings: Final = (
|
||||
await proxy_config.get_hierarchical_router_settings(user_api_key_dict, prisma_client, proxy_logging_obj)
|
||||
if wants_anthropic_format and is_claude_code_client(client_headers)
|
||||
else None
|
||||
)
|
||||
view_aliases: Final = (
|
||||
view_router_settings.get("model_group_alias") if isinstance(view_router_settings, Mapping) else None
|
||||
)
|
||||
routing_names: Final = ClaudeCodeRoutingNames(
|
||||
llm_router,
|
||||
team_id or user_api_key_dict.team_id,
|
||||
(
|
||||
user_api_key_dict.aliases,
|
||||
user_api_key_dict.team_model_aliases,
|
||||
view_aliases,
|
||||
),
|
||||
)
|
||||
|
||||
# Validate scope parameter if provided
|
||||
if scope is not None and scope != "expand":
|
||||
|
|
@ -10557,6 +10588,11 @@ async def model_list(
|
|||
return create_anthropic_model_list_response(
|
||||
admin_listing,
|
||||
display_names=configured_display_names(admin_entries, llm_router),
|
||||
listed_ids=claude_code_view_ids(
|
||||
admin_listing,
|
||||
client_headers,
|
||||
routing_names,
|
||||
),
|
||||
)
|
||||
|
||||
return dict(
|
||||
|
|
@ -10605,6 +10641,11 @@ async def model_list(
|
|||
return create_anthropic_model_list_response(
|
||||
listing,
|
||||
display_names=configured_display_names(entries, llm_router),
|
||||
listed_ids=claude_code_view_ids(
|
||||
listing,
|
||||
client_headers,
|
||||
routing_names,
|
||||
),
|
||||
)
|
||||
|
||||
return dict(
|
||||
|
|
@ -12841,7 +12882,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)
|
||||
|
|
@ -12854,7 +12897,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,
|
||||
|
|
@ -17677,6 +17720,66 @@ async def delete_callback(
|
|||
)
|
||||
|
||||
|
||||
def _normalize_callback_alias(callback_name: str) -> str:
|
||||
callback_aliases: Final = (
|
||||
("opentelemetry", "otel"),
|
||||
("s3_v2", "s3"),
|
||||
("aws_sqs", "sqs"),
|
||||
("custom_callback_api", "generic_api"),
|
||||
)
|
||||
return next(
|
||||
(canonical_name for alias, canonical_name in callback_aliases if alias == callback_name),
|
||||
callback_name,
|
||||
)
|
||||
|
||||
|
||||
def _callback_module_name(callback: CustomLogger | Callable[..., object]) -> str:
|
||||
if inspect.ismethod(callback):
|
||||
return callback.__func__.__module__
|
||||
if inspect.isfunction(callback):
|
||||
return callback.__module__
|
||||
return type(callback).__module__
|
||||
|
||||
|
||||
def _is_litellm_internal_callback(callback_name: str, callback: CustomLogger | Callable[..., object]) -> bool:
|
||||
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
|
||||
|
||||
module_owner: Final = _callback_module_name(callback).partition(".")[0]
|
||||
is_registered_integration: Final = callback_name in CustomLoggerRegistry.CALLBACK_CLASS_STR_TO_CLASS_TYPE
|
||||
return not is_registered_integration and module_owner in ("litellm", "litellm_enterprise")
|
||||
|
||||
|
||||
def _is_instance_of_configured_callback(
|
||||
callback_name: str, callback: CustomLogger | Callable[..., object], configured_classes: tuple[type, ...]
|
||||
) -> bool:
|
||||
"""Self-naming OTel-family instances (`arize`, `weave_otel`) match by name, so a configured `logfire` (a bare
|
||||
`OpenTelemetry`) does not hide YAML-configured siblings of the same class."""
|
||||
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
|
||||
|
||||
class_derived_name: Final = CustomLoggerRegistry.get_callback_str_from_class_type(type(callback))
|
||||
return isinstance(callback, configured_classes) and callback_name in (class_derived_name, type(callback).__name__)
|
||||
|
||||
|
||||
def _hidden_runtime_callback_names(configured_callback_names: frozenset[str]) -> frozenset[str]:
|
||||
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
|
||||
|
||||
configured_classes: Final = tuple(
|
||||
CustomLoggerRegistry.CALLBACK_CLASS_STR_TO_CLASS_TYPE[name]
|
||||
for name in configured_callback_names
|
||||
if name in CustomLoggerRegistry.CALLBACK_CLASS_STR_TO_CLASS_TYPE
|
||||
)
|
||||
configured_modules: Final = frozenset(name.rsplit(".", 1)[0] for name in configured_callback_names if "." in name)
|
||||
internal_callback_names: Final = frozenset({"cache", "vector_store_pre_call_hook"})
|
||||
return internal_callback_names | frozenset(
|
||||
callback_name
|
||||
for callback_name, callback in litellm.logging_callback_manager.get_callback_objects()
|
||||
if isinstance(callback, CustomGuardrail)
|
||||
or _is_litellm_internal_callback(callback_name, callback)
|
||||
or _is_instance_of_configured_callback(callback_name, callback, configured_classes)
|
||||
or _callback_module_name(callback) in configured_modules
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/get/config/callbacks",
|
||||
tags=["config.yaml"],
|
||||
|
|
@ -17709,10 +17812,10 @@ async def get_config(
|
|||
# Normalize string callbacks to lists
|
||||
def normalize_callback(callback):
|
||||
if isinstance(callback, str):
|
||||
return [callback]
|
||||
elif callback is None:
|
||||
return []
|
||||
return callback
|
||||
return (callback,)
|
||||
if callback is None:
|
||||
return ()
|
||||
return tuple(callback) if isinstance(callback, (list, dict)) else ()
|
||||
|
||||
_success_callbacks = normalize_callback(_success_callbacks)
|
||||
_failure_callbacks = normalize_callback(_failure_callbacks)
|
||||
|
|
@ -17743,6 +17846,30 @@ async def get_config(
|
|||
for _callback in _success_and_failure_callbacks:
|
||||
_data_to_return.append(process_callback(_callback, "success_and_failure", environment_variables))
|
||||
|
||||
configured_callback_names: Final = frozenset(
|
||||
_normalize_callback_alias(callback)
|
||||
for callback in (_success_callbacks + _failure_callbacks + _success_and_failure_callbacks)
|
||||
)
|
||||
runtime_callbacks_by_type: Final = litellm.logging_callback_manager.get_callbacks_by_type()
|
||||
hidden_callback_names: Final = _hidden_runtime_callback_names(configured_callback_names)
|
||||
runtime_callback_rows: Final = tuple(
|
||||
(_normalize_callback_alias(callback_name), callback_type)
|
||||
for callback_type, callback_names in (
|
||||
("success", runtime_callbacks_by_type["success"]),
|
||||
("failure", runtime_callbacks_by_type["failure"]),
|
||||
("success_and_failure", runtime_callbacks_by_type["success_and_failure"]),
|
||||
)
|
||||
for callback_name in callback_names
|
||||
if callback_name not in hidden_callback_names
|
||||
)
|
||||
runtime_only_rows: Final = sorted(
|
||||
frozenset(row for row in runtime_callback_rows if row[0] not in configured_callback_names)
|
||||
)
|
||||
_data_to_return.extend(
|
||||
dict(process_callback(callback_name, callback_type, environment_variables), read_only=True)
|
||||
for callback_name, callback_type in runtime_only_rows
|
||||
)
|
||||
|
||||
_data_to_return = _apply_callback_role_gate(_data_to_return, is_full_admin)
|
||||
|
||||
# Check if slack alerting is on
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
@ -2900,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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -149,6 +151,7 @@ from litellm.router_utils.common_utils import (
|
|||
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,
|
||||
|
|
@ -5198,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):
|
||||
|
|
@ -5234,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,
|
||||
|
|
@ -5754,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(
|
||||
**{
|
||||
|
|
@ -12113,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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -256,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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
):
|
||||
|
|
|
|||
|
|
@ -154,6 +154,22 @@ LATENCY_BUCKETS: Final = (
|
|||
float("inf"),
|
||||
)
|
||||
|
||||
UNKNOWN_INPUT_SEQUENCE_LENGTH: Final = "unknown"
|
||||
INPUT_SEQUENCE_LENGTH_BUCKETS: Final = (
|
||||
(1_000, "0-1k"),
|
||||
(4_000, "1k-4k"),
|
||||
(16_000, "4k-16k"),
|
||||
(64_000, "16k-64k"),
|
||||
(float("inf"), "64k+"),
|
||||
)
|
||||
|
||||
|
||||
def get_input_sequence_length_bucket(prompt_tokens: object) -> str:
|
||||
if not isinstance(prompt_tokens, int) or isinstance(prompt_tokens, bool) or prompt_tokens < 0:
|
||||
return UNKNOWN_INPUT_SEQUENCE_LENGTH
|
||||
return next(label for upper, label in INPUT_SEQUENCE_LENGTH_BUCKETS if prompt_tokens < upper)
|
||||
|
||||
|
||||
# Batch jobs can run for minutes to hours; buckets span 1 min → 24 h.
|
||||
BATCH_DURATION_BUCKETS: Final = (
|
||||
60.0,
|
||||
|
|
@ -205,6 +221,7 @@ class UserAPIKeyLabelNames(Enum):
|
|||
MCP_TOOL_NAME = "mcp_tool_name"
|
||||
MCP_SERVER_NAME = "mcp_server_name"
|
||||
SERVICE_TIER = "service_tier"
|
||||
INPUT_SEQUENCE_LENGTH = "input_sequence_length"
|
||||
|
||||
|
||||
DEFINED_PROMETHEUS_METRICS = Literal[
|
||||
|
|
@ -857,6 +874,13 @@ class PrometheusMetricLabels:
|
|||
"litellm_images_generated_metric",
|
||||
}
|
||||
)
|
||||
_input_sequence_length_metrics: ClassVar[frozenset[str]] = frozenset(
|
||||
{
|
||||
"litellm_llm_api_latency_metric",
|
||||
"litellm_llm_api_time_to_first_token_metric",
|
||||
"litellm_request_total_latency_metric",
|
||||
}
|
||||
)
|
||||
# Managed batch metrics
|
||||
_batch_user_labels = [
|
||||
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
|
||||
|
|
@ -955,14 +979,23 @@ class PrometheusMetricLabels:
|
|||
custom_labels.append(label)
|
||||
|
||||
if label_name in PrometheusMetricLabels._org_label_metrics:
|
||||
for label in [
|
||||
for label in (
|
||||
UserAPIKeyLabelNames.ORG_ID.value,
|
||||
UserAPIKeyLabelNames.ORG_ALIAS.value,
|
||||
]:
|
||||
):
|
||||
if label not in default_labels and label not in custom_labels:
|
||||
custom_labels.append(label)
|
||||
|
||||
return default_labels + custom_labels
|
||||
input_sequence_length_labels: Final = (
|
||||
(UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value,)
|
||||
if (
|
||||
label_name in PrometheusMetricLabels._input_sequence_length_metrics
|
||||
and litellm.prometheus_emit_input_sequence_length_label is True
|
||||
and UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in custom_labels
|
||||
)
|
||||
else ()
|
||||
)
|
||||
return [*default_labels, *custom_labels, *input_sequence_length_labels]
|
||||
|
||||
|
||||
_USER_API_KEY_LABEL_VALUE_INIT_ALIASES: Final[Mapping[str, str]] = MappingProxyType(
|
||||
|
|
@ -1015,6 +1048,7 @@ class UserAPIKeyLabelValues:
|
|||
mcp_tool_name: str | None = None
|
||||
mcp_server_name: str | None = None
|
||||
service_tier: str | None = None
|
||||
input_sequence_length: str | None = None
|
||||
|
||||
# Added for test compatibility.
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -8858,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,
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ longer signal it.
|
|||
|
||||
### Added
|
||||
|
||||
- **team**: Optional `team_id` argument on `litellm_team`, so teams can be created with a stable, human-readable ID instead of a provider-generated UUID; changing it forces replacement
|
||||
- **jwt_key_mapping**: New `litellm_jwt_key_mapping` resource for the proxy's JWT to virtual key mappings, so JWT clients identified by a claim (`client_id`, `azp`, `sub`) map to virtual keys and inherit their models, budgets and rate limits. Supports `description` and `is_active`, rotating the mapped key in place, and forces replacement when the claim name or value changes
|
||||
- **team**: `soft_budget`, `tags`, and `soft_budget_alerting_emails` attributes on `litellm_team`, matching what `/team/new` and `/team/update` already accept; `soft_budget_alerting_emails` is sent under `metadata`, where the proxy reads it
|
||||
- **user**: New `litellm_user` resource and `litellm_user` / `litellm_users` data sources for managing internal users
|
||||
|
|
@ -38,12 +39,14 @@ longer signal it.
|
|||
|
||||
- **team**: Read now decodes the `team_info` envelope `/team/info` actually returns, so team attributes refresh from the proxy instead of always falling back to the prior state
|
||||
- **key**: Read now unwraps the `info` envelope `/key/info` actually returns; previously reads mapped nothing back into state, so drift on a key was never detected
|
||||
- **key**: Read now picks up `model_rpm_limit`, `model_tpm_limit`, `guardrails`, `tags`, `enforced_params`, `allowed_passthrough_routes`, `rpm_limit_type`, `tpm_limit_type` and `prompts` from `info.metadata`, where the proxy actually stores them; previously they stayed empty in state, so a matching config showed a permanent phantom diff on them and out-of-band changes to them were never detected
|
||||
- **key**: Updates no longer send an empty `budget_duration`, which the proxy rejects with a 400; any update to a key without a configured `budget_duration` previously failed outright
|
||||
- **key**: A config-supplied `key` value (write-only) is now forwarded to `/key/generate`; previously it was silently dropped and the proxy generated a random key instead
|
||||
- **security**: The `litellm_key` data source and `litellm_key_block` resource normalize raw `sk-` keys to their SHA-256 token hash before building request URLs and resource IDs, so plaintext keys no longer land in reverse-proxy access logs, Terraform plan output, or state IDs
|
||||
|
||||
### Changed
|
||||
|
||||
- **key** (breaking): `model_max_budget` on `litellm_key` is now a JSON string of per-model budget objects (`jsonencode({"gpt-4o-mini" = {budget_limit = 50, time_period = "30d"}})`), matching `litellm_user`, `litellm_budget` and `litellm_tag`. The old `map(number)` form sent bare numbers to `/key/generate`, which the proxy rejects with a 500 (`'int' object is not iterable`), so every key with a non-empty `model_max_budget` failed to apply. Existing state upgrades automatically (schema version 1) and the attribute is refilled from the proxy on the next read; configurations still using the map form must be rewritten
|
||||
- **Versioning**: the provider is now published at the LiteLLM version, from the same commit as the proxy, on every LiteLLM release (dev, rc, stable). The `0.x` line ends at `0.4.0`; a `~> 0.4` constraint will not receive further releases, so re-pin to the LiteLLM version your proxy runs (for example `~> 1.99.0`). Existing `0.x` versions remain in the registry and keep verifying
|
||||
|
||||
## [0.4.0] - 2026-08-06
|
||||
|
|
|
|||
|
|
@ -103,9 +103,12 @@ resource "litellm_key" "example_key" {
|
|||
permissions = {
|
||||
can_create_keys = "true"
|
||||
}
|
||||
model_max_budget = {
|
||||
"gpt-4" = 50.0
|
||||
}
|
||||
model_max_budget = jsonencode({
|
||||
"gpt-4" = {
|
||||
budget_limit = 50.0
|
||||
time_period = "30d"
|
||||
}
|
||||
})
|
||||
model_rpm_limit = {
|
||||
"claude-3.5-sonnet" = 30
|
||||
}
|
||||
|
|
|
|||
|
|
@ -30,9 +30,12 @@ resource "litellm_key" "example" {
|
|||
permissions = {
|
||||
"can_create_keys" = "true"
|
||||
}
|
||||
model_max_budget = {
|
||||
"gpt-4" = 50.0
|
||||
}
|
||||
model_max_budget = jsonencode({
|
||||
"gpt-4" = {
|
||||
budget_limit = 50.0
|
||||
time_period = "30d"
|
||||
}
|
||||
})
|
||||
model_rpm_limit = {
|
||||
"gpt-3.5-turbo" = 30
|
||||
}
|
||||
|
|
@ -73,7 +76,7 @@ The following arguments are supported:
|
|||
|
||||
* `key_alias` - (Optional) Alias for this key. This provides a human-readable identifier for the key.
|
||||
|
||||
* `duration` - (Optional) Duration for which this key is valid. This sets an expiration time for the key.
|
||||
* `duration` - (Optional) How long the key stays valid, e.g. "30d" or "12h". The proxy stores this as an absolute `expires` timestamp. Changing the value resets the expiry to the time of the update plus the new duration; removing it from the configuration leaves the current expiry in place.
|
||||
|
||||
* `aliases` - (Optional) Map of model aliases. This allows you to create custom names for models when using this key.
|
||||
|
||||
|
|
@ -81,7 +84,7 @@ The following arguments are supported:
|
|||
|
||||
* `permissions` - (Optional) Permissions associated with this key. This defines what actions are allowed with this key.
|
||||
|
||||
* `model_max_budget` - (Optional) Maximum budget per model. This allows setting different budget limits for each model.
|
||||
* `model_max_budget` - (Optional) JSON string of per-model budget config, e.g. `jsonencode({"gpt-4" = {budget_limit = 50.0, time_period = "30d"}})`. Each model maps to an object with `budget_limit` (or `max_budget`), `time_period` (or `budget_duration`), `tpm_limit` and `rpm_limit`.
|
||||
|
||||
* `model_rpm_limit` - (Optional) Requests per minute limit per model. This allows setting different RPM limits for each model.
|
||||
|
||||
|
|
|
|||
|
|
@ -14,6 +14,16 @@ resource "litellm_team" "engineering" {
|
|||
}
|
||||
```
|
||||
|
||||
### Team with a Custom ID
|
||||
|
||||
```hcl
|
||||
resource "litellm_team" "platform" {
|
||||
team_id = "platform-team"
|
||||
team_alias = "platform"
|
||||
models = ["gpt-4-proxy"]
|
||||
}
|
||||
```
|
||||
|
||||
### Team with Comprehensive Configuration
|
||||
|
||||
```hcl
|
||||
|
|
@ -92,6 +102,8 @@ resource "litellm_team" "model_dependent_team" {
|
|||
|
||||
The following arguments are supported:
|
||||
|
||||
* `team_id` - (Optional) A stable, human-readable ID for the team (for example `platform-team`). If omitted, the provider generates a random UUID. Changing this forces a new team to be created.
|
||||
|
||||
* `team_alias` - (Required) A human-readable identifier for the team.
|
||||
|
||||
* `organization_id` - (Optional) The ID of the organization this team belongs to.
|
||||
|
|
@ -152,7 +164,7 @@ The following arguments are supported:
|
|||
|
||||
In addition to the arguments above, the following attributes are exported:
|
||||
|
||||
* `id` - The unique identifier for the team.
|
||||
* `id` - The unique identifier for the team, equal to `team_id`.
|
||||
|
||||
## Import
|
||||
|
||||
|
|
@ -162,7 +174,7 @@ Teams can be imported using the team ID:
|
|||
terraform import litellm_team.engineering <team-id>
|
||||
```
|
||||
|
||||
Note: The team ID is generated when the team is created and is different from the `team_alias`.
|
||||
Note: Unless `team_id` is set, the team ID is generated when the team is created and is different from the `team_alias`.
|
||||
|
||||
## Note on Team Members
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import (
|
|||
"bytes"
|
||||
"crypto/tls"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
|
|
@ -19,6 +20,20 @@ type Client struct {
|
|||
InsecureSkipVerify bool
|
||||
}
|
||||
|
||||
type apiError struct {
|
||||
StatusCode int
|
||||
Body string
|
||||
}
|
||||
|
||||
func (e *apiError) Error() string {
|
||||
return fmt.Sprintf("API request failed with status code %d: %s", e.StatusCode, e.Body)
|
||||
}
|
||||
|
||||
func isNotFound(err error) bool {
|
||||
var apiErr *apiError
|
||||
return errors.As(err, &apiErr) && apiErr.StatusCode == http.StatusNotFound
|
||||
}
|
||||
|
||||
func NewClient(apiBase, apiKey string, insecureSkipVerify bool) *Client {
|
||||
tr := &http.Transport{
|
||||
TLSClientConfig: &tls.Config{InsecureSkipVerify: insecureSkipVerify},
|
||||
|
|
@ -57,6 +72,9 @@ func (c *Client) CreateKey(key *Key) (*Key, error) {
|
|||
|
||||
func (c *Client) GetKey(keyID string) (*Key, error) {
|
||||
resp, err := c.sendRequest("GET", fmt.Sprintf("/key/info?key=%s", keyID), nil)
|
||||
if isNotFound(err) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -69,32 +87,71 @@ func (c *Client) GetKey(keyID string) (*Key, error) {
|
|||
info["key"] = k
|
||||
}
|
||||
}
|
||||
hoistKeyFieldsStoredInMetadata(info)
|
||||
return c.parseKeyResponse(info)
|
||||
}
|
||||
|
||||
return c.parseKeyResponse(resp)
|
||||
}
|
||||
|
||||
var keyFieldsStoredInMetadata = []string{
|
||||
"model_rpm_limit",
|
||||
"model_tpm_limit",
|
||||
"guardrails",
|
||||
"tags",
|
||||
"enforced_params",
|
||||
"allowed_passthrough_routes",
|
||||
"rpm_limit_type",
|
||||
"tpm_limit_type",
|
||||
"prompts",
|
||||
}
|
||||
|
||||
func hoistKeyFieldsStoredInMetadata(info map[string]interface{}) {
|
||||
metadata, ok := info["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
for _, field := range keyFieldsStoredInMetadata {
|
||||
if existing, present := info[field]; present && existing != nil {
|
||||
continue
|
||||
}
|
||||
if v, present := metadata[field]; present {
|
||||
info[field] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) UpdateKey(key *Key) (*Key, error) {
|
||||
// Create a new map with only the fields that can be updated
|
||||
updateData := map[string]interface{}{
|
||||
"key": key.Key,
|
||||
"team_id": key.TeamID,
|
||||
"metadata": key.Metadata,
|
||||
"key_alias": key.KeyAlias,
|
||||
"aliases": key.Aliases,
|
||||
"permissions": key.Permissions,
|
||||
"model_max_budget": key.ModelMaxBudget,
|
||||
"model_rpm_limit": key.ModelRPMLimit,
|
||||
"model_tpm_limit": key.ModelTPMLimit,
|
||||
"blocked": key.Blocked,
|
||||
}
|
||||
|
||||
// The proxy keeps the stored metadata only when the field is absent, so nil means omit.
|
||||
if key.Metadata != nil {
|
||||
updateData["metadata"] = key.Metadata
|
||||
}
|
||||
if key.ModelRPMLimit != nil {
|
||||
updateData["model_rpm_limit"] = key.ModelRPMLimit
|
||||
}
|
||||
if key.ModelTPMLimit != nil {
|
||||
updateData["model_tpm_limit"] = key.ModelTPMLimit
|
||||
}
|
||||
|
||||
// The proxy rejects an empty-string budget_duration with a 400, so only
|
||||
// send it when set.
|
||||
if key.BudgetDuration != "" {
|
||||
updateData["budget_duration"] = key.BudgetDuration
|
||||
}
|
||||
if key.Duration != "" {
|
||||
updateData["duration"] = key.Duration
|
||||
}
|
||||
|
||||
// Only add pointer fields if they are explicitly set
|
||||
if key.MaxBudget != nil {
|
||||
|
|
@ -366,7 +423,7 @@ func (c *Client) sendRequest(method, path string, body interface{}) (map[string]
|
|||
log.Printf("Response body: %s", c.redactSensitiveData(string(bodyBytes)))
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("API request failed with status code %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
return nil, &apiError{StatusCode: resp.StatusCode, Body: string(bodyBytes)}
|
||||
}
|
||||
|
||||
var result map[string]interface{}
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ package litellm
|
|||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"github.com/hashicorp/go-cty/cty"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/diag"
|
||||
|
|
@ -10,7 +12,7 @@ import (
|
|||
)
|
||||
|
||||
func resourceKey() *schema.Resource {
|
||||
return &schema.Resource{
|
||||
r := &schema.Resource{
|
||||
CreateContext: resourceKeyCreate,
|
||||
ReadContext: resourceKeyRead,
|
||||
UpdateContext: resourceKeyUpdate,
|
||||
|
|
@ -18,6 +20,7 @@ func resourceKey() *schema.Resource {
|
|||
Importer: &schema.ResourceImporter{
|
||||
StateContext: schema.ImportStatePassthroughContext,
|
||||
},
|
||||
SchemaVersion: 1,
|
||||
Schema: map[string]*schema.Schema{
|
||||
"key": {
|
||||
Type: schema.TypeString,
|
||||
|
|
@ -86,8 +89,9 @@ func resourceKey() *schema.Resource {
|
|||
Optional: true,
|
||||
},
|
||||
"duration": {
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
Description: "How long the key stays valid, e.g. \"30d\" or \"12h\". Changing it resets the expiry to the time of the update plus the new duration; removing it leaves the current expiry in place",
|
||||
},
|
||||
"aliases": {
|
||||
Type: schema.TypeMap,
|
||||
|
|
@ -105,9 +109,11 @@ func resourceKey() *schema.Resource {
|
|||
Elem: &schema.Schema{Type: schema.TypeString},
|
||||
},
|
||||
"model_max_budget": {
|
||||
Type: schema.TypeMap,
|
||||
Optional: true,
|
||||
Elem: &schema.Schema{Type: schema.TypeFloat, Computed: true},
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
ValidateFunc: validateKeyModelMaxBudget,
|
||||
DiffSuppressFunc: budgetSuppressEquivalentJSON,
|
||||
Description: "JSON string of per-model budget config (e.g. '{\"gpt-4o-mini\": {\"budget_limit\": 50, \"time_period\": \"30d\"}}')",
|
||||
},
|
||||
"model_rpm_limit": {
|
||||
Type: schema.TypeMap,
|
||||
|
|
@ -182,6 +188,79 @@ func resourceKey() *schema.Resource {
|
|||
},
|
||||
},
|
||||
}
|
||||
r.StateUpgraders = []schema.StateUpgrader{{
|
||||
Version: 0,
|
||||
Type: resourceKeyV0Type(r.Schema),
|
||||
Upgrade: resourceKeyStateUpgradeV0,
|
||||
}}
|
||||
return r
|
||||
}
|
||||
|
||||
// Schema version 0 typed model_max_budget as map(number), which the proxy
|
||||
// rejects; version 1 stores the per-model BudgetConfig objects as a JSON string.
|
||||
func resourceKeyV0Type(current map[string]*schema.Schema) cty.Type {
|
||||
v0 := make(map[string]*schema.Schema, len(current))
|
||||
for k, v := range current {
|
||||
v0[k] = v
|
||||
}
|
||||
v0["model_max_budget"] = &schema.Schema{
|
||||
Type: schema.TypeMap,
|
||||
Optional: true,
|
||||
Elem: &schema.Schema{Type: schema.TypeFloat},
|
||||
}
|
||||
return (&schema.Resource{Schema: v0}).CoreConfigSchema().ImpliedType()
|
||||
}
|
||||
|
||||
func resourceKeyStateUpgradeV0(_ context.Context, rawState map[string]interface{}, _ interface{}) (map[string]interface{}, error) {
|
||||
delete(rawState, "model_max_budget")
|
||||
return rawState, nil
|
||||
}
|
||||
|
||||
var keyModelBudgetFields = map[string]bool{
|
||||
"budget_limit": true,
|
||||
"max_budget": true,
|
||||
"time_period": true,
|
||||
"budget_duration": true,
|
||||
"tpm_limit": true,
|
||||
"rpm_limit": true,
|
||||
}
|
||||
|
||||
func validateKeyModelMaxBudget(v interface{}, k string) ([]string, []error) {
|
||||
var parsed map[string]json.RawMessage
|
||||
if err := json.Unmarshal([]byte(v.(string)), &parsed); err != nil || parsed == nil {
|
||||
return nil, []error{fmt.Errorf("%q must be a JSON object keyed by model name, got %s", k, v)}
|
||||
}
|
||||
for model, cfg := range parsed {
|
||||
var budget map[string]json.RawMessage
|
||||
if err := json.Unmarshal(cfg, &budget); err != nil || len(budget) == 0 {
|
||||
return nil, []error{fmt.Errorf("%q[%q] must be a budget object such as {\"budget_limit\": 50, \"time_period\": \"30d\"}, got %s", k, model, cfg)}
|
||||
}
|
||||
for field := range budget {
|
||||
if !keyModelBudgetFields[field] {
|
||||
return nil, []error{fmt.Errorf("%q[%q] has unknown budget field %q; supported fields are budget_limit, max_budget, time_period, budget_duration, tpm_limit, rpm_limit", k, model, field)}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func parseKeyModelMaxBudget(raw string) map[string]interface{} {
|
||||
var parsed map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(raw), &parsed); err != nil || parsed == nil {
|
||||
return map[string]interface{}{}
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
func keyModelMaxBudgetJSON(modelMaxBudget map[string]interface{}) string {
|
||||
if len(modelMaxBudget) == 0 {
|
||||
return ""
|
||||
}
|
||||
encoded, err := json.Marshal(modelMaxBudget)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(encoded)
|
||||
}
|
||||
|
||||
func resourceKeyCreate(ctx context.Context, d *schema.ResourceData, m interface{}) diag.Diagnostics {
|
||||
|
|
@ -219,10 +298,12 @@ func resourceKeyRead(ctx context.Context, d *schema.ResourceData, m interface{})
|
|||
}
|
||||
|
||||
if key == nil {
|
||||
log.Printf("[WARN] Key %s not found, removing from state", d.Id())
|
||||
d.SetId("")
|
||||
return nil
|
||||
}
|
||||
|
||||
key.Metadata = declaredKeyMetadata(key.Metadata, d.Get("metadata").(map[string]interface{}))
|
||||
mapKeyToResourceData(d, key)
|
||||
return nil
|
||||
}
|
||||
|
|
@ -232,15 +313,75 @@ func resourceKeyUpdate(ctx context.Context, d *schema.ResourceData, m interface{
|
|||
|
||||
key := &Key{Key: d.Id()}
|
||||
mapResourceDataToKey(d, key)
|
||||
if !d.HasChange("duration") {
|
||||
key.Duration = ""
|
||||
}
|
||||
key.ModelRPMLimit = changedMap(d, "model_rpm_limit")
|
||||
key.ModelTPMLimit = changedMap(d, "model_tpm_limit")
|
||||
|
||||
_, err := c.UpdateKey(key)
|
||||
metadata, err := plannedKeyMetadata(c, d)
|
||||
if err != nil {
|
||||
d.Partial(true)
|
||||
return diag.FromErr(fmt.Errorf("error updating key: %s", err))
|
||||
}
|
||||
key.Metadata = metadata
|
||||
|
||||
if _, err := c.UpdateKey(key); err != nil {
|
||||
return diag.FromErr(fmt.Errorf("error updating key: %s", err))
|
||||
}
|
||||
|
||||
return resourceKeyRead(ctx, d, m)
|
||||
}
|
||||
|
||||
func changedMap(d *schema.ResourceData, name string) map[string]interface{} {
|
||||
if !d.HasChange(name) {
|
||||
return nil
|
||||
}
|
||||
return d.Get(name).(map[string]interface{})
|
||||
}
|
||||
|
||||
func plannedKeyMetadata(c *Client, d *schema.ResourceData) (map[string]interface{}, error) {
|
||||
if !d.HasChange("metadata") {
|
||||
return nil, nil
|
||||
}
|
||||
current, err := c.GetKey(d.Id())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if current == nil {
|
||||
return nil, fmt.Errorf("key %s no longer exists", d.Id())
|
||||
}
|
||||
oldDeclared, newDeclared := d.GetChange("metadata")
|
||||
return mergeKeyMetadata(current.Metadata, oldDeclared.(map[string]interface{}), newDeclared.(map[string]interface{})), nil
|
||||
}
|
||||
|
||||
func declaredKeyMetadata(server, declared map[string]interface{}) map[string]interface{} {
|
||||
if server == nil {
|
||||
return nil
|
||||
}
|
||||
result := make(map[string]interface{}, len(declared))
|
||||
for k := range declared {
|
||||
if v, ok := server[k]; ok {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func mergeKeyMetadata(server, oldDeclared, newDeclared map[string]interface{}) map[string]interface{} {
|
||||
result := make(map[string]interface{}, len(server)+len(newDeclared))
|
||||
for k, v := range server {
|
||||
result[k] = v
|
||||
}
|
||||
for k := range oldDeclared {
|
||||
delete(result, k)
|
||||
}
|
||||
for k, v := range newDeclared {
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func resourceKeyDelete(ctx context.Context, d *schema.ResourceData, m interface{}) diag.Diagnostics {
|
||||
c := m.(*Client)
|
||||
|
||||
|
|
@ -285,7 +426,7 @@ func mapResourceDataToKey(d *schema.ResourceData, key *Key) {
|
|||
key.Aliases = d.Get("aliases").(map[string]interface{})
|
||||
key.Config = d.Get("config").(map[string]interface{})
|
||||
key.Permissions = d.Get("permissions").(map[string]interface{})
|
||||
key.ModelMaxBudget = d.Get("model_max_budget").(map[string]interface{})
|
||||
key.ModelMaxBudget = parseKeyModelMaxBudget(d.Get("model_max_budget").(string))
|
||||
key.ModelRPMLimit = d.Get("model_rpm_limit").(map[string]interface{})
|
||||
key.ModelTPMLimit = d.Get("model_tpm_limit").(map[string]interface{})
|
||||
key.Guardrails = expandStringList(d.Get("guardrails").([]interface{}))
|
||||
|
|
@ -358,9 +499,7 @@ func mapKeyToResourceData(d *schema.ResourceData, key *Key) {
|
|||
if key.Permissions != nil {
|
||||
d.Set("permissions", key.Permissions)
|
||||
}
|
||||
if key.ModelMaxBudget != nil {
|
||||
d.Set("model_max_budget", key.ModelMaxBudget)
|
||||
}
|
||||
d.Set("model_max_budget", keyModelMaxBudgetJSON(key.ModelMaxBudget))
|
||||
if key.ModelRPMLimit != nil {
|
||||
d.Set("model_rpm_limit", key.ModelRPMLimit)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,9 +6,11 @@ import (
|
|||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/terraform"
|
||||
)
|
||||
|
||||
func newKeyResourceData(t *testing.T, raw map[string]interface{}) *schema.ResourceData {
|
||||
|
|
@ -193,6 +195,102 @@ func TestCreateKeySendsConfigSuppliedKey(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
// The proxy validates each model_max_budget entry as a BudgetConfig object and
|
||||
// 500s on a bare number, so the JSON string must reach /key/generate as nested
|
||||
// objects and the proxy's response must map back to equivalent JSON in state.
|
||||
func TestCreateKeySendsModelMaxBudgetAsBudgetObjects(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if r.URL.Path == "/key/generate" {
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
json.Unmarshal(body, &captured)
|
||||
w.Write([]byte(`{"key": "sk-test", "token_id": "hash-1"}`))
|
||||
return
|
||||
}
|
||||
w.Write([]byte(`{"key": "hash-1", "info": {"model_max_budget": {"gpt-4o-mini": {"budget_limit": 50, "time_period": "30d", "rpm_limit": 60}}}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := newKeyResourceData(t, map[string]interface{}{
|
||||
"model_max_budget": `{"gpt-4o-mini": {"budget_limit": 50, "time_period": "30d"}}`,
|
||||
})
|
||||
|
||||
if diags := resourceKeyCreate(context.Background(), d, client); diags.HasError() {
|
||||
t.Fatalf("create returned error: %v", diags)
|
||||
}
|
||||
|
||||
budgets, ok := captured["model_max_budget"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("create payload model_max_budget = %v, want object", captured["model_max_budget"])
|
||||
}
|
||||
cfg, ok := budgets["gpt-4o-mini"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("model_max_budget[gpt-4o-mini] = %v, want BudgetConfig object", budgets["gpt-4o-mini"])
|
||||
}
|
||||
if cfg["budget_limit"] != float64(50) || cfg["time_period"] != "30d" {
|
||||
t.Errorf("BudgetConfig = %v, want budget_limit 50 and time_period 30d", cfg)
|
||||
}
|
||||
|
||||
var state map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(d.Get("model_max_budget").(string)), &state); err != nil {
|
||||
t.Fatalf("state model_max_budget %q is not JSON: %v", d.Get("model_max_budget"), err)
|
||||
}
|
||||
if got, _ := state["gpt-4o-mini"].(map[string]interface{}); got["budget_limit"] != float64(50) || got["rpm_limit"] != float64(60) {
|
||||
t.Errorf("state model_max_budget = %v, want the BudgetConfig read back from /key/info", state)
|
||||
}
|
||||
}
|
||||
|
||||
// Schema version 0 stored model_max_budget as map(number); that state cannot
|
||||
// decode into the version 1 string attribute, so the upgrader must drop it.
|
||||
func TestKeyStateUpgradeV0DropsMapModelMaxBudget(t *testing.T) {
|
||||
upgraded, err := resourceKey().StateUpgraders[0].Upgrade(context.Background(), map[string]interface{}{
|
||||
"id": "hash-1",
|
||||
"key_alias": "legacy",
|
||||
"model_max_budget": map[string]interface{}{"gpt-4o-mini": 50.0},
|
||||
}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("upgrade returned error: %v", err)
|
||||
}
|
||||
if _, present := upgraded["model_max_budget"]; present {
|
||||
t.Errorf("upgraded state still carries map model_max_budget: %v", upgraded["model_max_budget"])
|
||||
}
|
||||
if upgraded["key_alias"] != "legacy" {
|
||||
t.Errorf("upgrade dropped unrelated attribute: %v", upgraded)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyModelMaxBudgetValidationRequiresBudgetObjects(t *testing.T) {
|
||||
validate := resourceKey().Schema["model_max_budget"].ValidateFunc
|
||||
for _, valid := range []string{
|
||||
`{}`,
|
||||
`{"gpt-4o-mini": {"budget_limit": 50, "time_period": "30d"}}`,
|
||||
`{"gpt-4o-mini": {"max_budget": 50, "rpm_limit": 60}, "gpt-4o": {"budget_duration": "1d", "tpm_limit": 1000}}`,
|
||||
} {
|
||||
if _, errs := validate(valid, "model_max_budget"); len(errs) != 0 {
|
||||
t.Errorf("validate(%s) = %v, want accepted", valid, errs)
|
||||
}
|
||||
}
|
||||
for _, invalid := range []string{
|
||||
`null`,
|
||||
`[]`,
|
||||
`"gpt-4o-mini"`,
|
||||
`50`,
|
||||
`{"gpt-4o-mini": 50}`,
|
||||
`{"gpt-4o-mini": null}`,
|
||||
`{"gpt-4o-mini": [50]}`,
|
||||
`{"gpt-4o-mini": {}}`,
|
||||
`{"gpt-4o-mini": {"budget_limt": 50}}`,
|
||||
`{"gpt-4o-mini": {"budget_limit": 50, "max_tokens": 100}}`,
|
||||
`not json`,
|
||||
} {
|
||||
if _, errs := validate(invalid, "model_max_budget"); len(errs) == 0 {
|
||||
t.Errorf("validate(%s) accepted a value that would send no per-model budget", invalid)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The proxy 400s on budget_duration: "", so an unset duration must be
|
||||
// omitted from the update payload entirely.
|
||||
func TestUpdateKeyOmitsEmptyBudgetDuration(t *testing.T) {
|
||||
|
|
@ -221,6 +319,47 @@ func TestUpdateKeyOmitsEmptyBudgetDuration(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestResourceKeyUpdateFailureKeepsPriorState(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if r.URL.Path == "/key/update" {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
w.Write([]byte(`{"error":{"message":"Invalid budget_duration 'bad'"}}`))
|
||||
return
|
||||
}
|
||||
w.Write([]byte(`{"key":"hash-1","info":{"key_alias":"demo","models":["fake-model"]}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
res := resourceKey()
|
||||
priorData := newKeyResourceData(t, map[string]interface{}{
|
||||
"key_alias": "demo",
|
||||
"models": []interface{}{"fake-model"},
|
||||
})
|
||||
priorData.SetId("hash-1")
|
||||
prior := priorData.State()
|
||||
config := terraform.NewResourceConfigRaw(map[string]interface{}{
|
||||
"key_alias": "demo",
|
||||
"models": []interface{}{"fake-model"},
|
||||
"budget_duration": "bad",
|
||||
})
|
||||
diff, err := res.Diff(context.Background(), prior, config, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("diff failed: %v", err)
|
||||
}
|
||||
|
||||
newState, diags := res.Apply(context.Background(), prior, diff, NewClient(srv.URL, "test-key", true))
|
||||
if !diags.HasError() {
|
||||
t.Fatal("apply succeeded, want the proxy's 400 surfaced as an error")
|
||||
}
|
||||
if got, ok := newState.Attributes["budget_duration"]; ok {
|
||||
t.Errorf("failed update persisted budget_duration=%q into state, want it absent", got)
|
||||
}
|
||||
if newState.Attributes["key_alias"] != "demo" {
|
||||
t.Errorf("prior key_alias lost from state: %v", newState.Attributes)
|
||||
}
|
||||
}
|
||||
|
||||
// /key/info nests the key's fields under "info"; GetKey must unwrap that
|
||||
// envelope or reads map nothing back into state.
|
||||
func TestGetKeyUnwrapsInfoEnvelope(t *testing.T) {
|
||||
|
|
@ -254,3 +393,296 @@ func TestGetKeyUnwrapsInfoEnvelope(t *testing.T) {
|
|||
t.Errorf("RPMLimit not parsed: %+v", key.RPMLimit)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetKeyReadsFieldsStoredInMetadata(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Write([]byte(`{
|
||||
"key": "hash-1",
|
||||
"info": {
|
||||
"models": ["gpt-4o-mini"],
|
||||
"metadata": {
|
||||
"team": "core-infra",
|
||||
"model_rpm_limit": {"gpt-4o-mini": 7},
|
||||
"model_tpm_limit": {"gpt-4o-mini": 10000},
|
||||
"guardrails": ["pii-guard"],
|
||||
"tags": ["prod"],
|
||||
"enforced_params": ["user"],
|
||||
"allowed_passthrough_routes": ["/v1/foo"],
|
||||
"rpm_limit_type": "guaranteed_throughput",
|
||||
"tpm_limit_type": "dynamic",
|
||||
"prompts": ["p1"]
|
||||
}
|
||||
}
|
||||
}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
key, err := client.GetKey("hash-1")
|
||||
if err != nil {
|
||||
t.Fatalf("GetKey returned error: %v", err)
|
||||
}
|
||||
if got, ok := key.ModelRPMLimit["gpt-4o-mini"].(float64); !ok || got != 7 {
|
||||
t.Errorf("ModelRPMLimit = %v, want gpt-4o-mini=7 read from metadata", key.ModelRPMLimit)
|
||||
}
|
||||
if got, ok := key.ModelTPMLimit["gpt-4o-mini"].(float64); !ok || got != 10000 {
|
||||
t.Errorf("ModelTPMLimit = %v, want gpt-4o-mini=10000 read from metadata", key.ModelTPMLimit)
|
||||
}
|
||||
if len(key.Guardrails) != 1 || key.Guardrails[0] != "pii-guard" {
|
||||
t.Errorf("Guardrails = %v, want [pii-guard]", key.Guardrails)
|
||||
}
|
||||
if len(key.Tags) != 1 || key.Tags[0] != "prod" {
|
||||
t.Errorf("Tags = %v, want [prod]", key.Tags)
|
||||
}
|
||||
if len(key.EnforcedParams) != 1 || key.EnforcedParams[0] != "user" {
|
||||
t.Errorf("EnforcedParams = %v, want [user]", key.EnforcedParams)
|
||||
}
|
||||
if len(key.AllowedPassthroughRoutes) != 1 || key.AllowedPassthroughRoutes[0] != "/v1/foo" {
|
||||
t.Errorf("AllowedPassthroughRoutes = %v, want [/v1/foo]", key.AllowedPassthroughRoutes)
|
||||
}
|
||||
if key.RPMLimitType != "guaranteed_throughput" || key.TPMLimitType != "dynamic" {
|
||||
t.Errorf("limit types = %q/%q, want guaranteed_throughput/dynamic", key.RPMLimitType, key.TPMLimitType)
|
||||
}
|
||||
if len(key.Prompts) != 1 || key.Prompts[0] != "p1" {
|
||||
t.Errorf("Prompts = %v, want [p1]", key.Prompts)
|
||||
}
|
||||
if key.Metadata["team"] != "core-infra" {
|
||||
t.Errorf("Metadata = %v, want team=core-infra preserved", key.Metadata)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetKeyPrefersTopLevelOverMetadataCopy(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Write([]byte(`{
|
||||
"key": "hash-1",
|
||||
"info": {
|
||||
"tags": ["top-level"],
|
||||
"guardrails": null,
|
||||
"metadata": {
|
||||
"tags": ["from-metadata"],
|
||||
"guardrails": ["from-metadata"]
|
||||
}
|
||||
}
|
||||
}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
key, err := client.GetKey("hash-1")
|
||||
if err != nil {
|
||||
t.Fatalf("GetKey returned error: %v", err)
|
||||
}
|
||||
if len(key.Tags) != 1 || key.Tags[0] != "top-level" {
|
||||
t.Errorf("Tags = %v, want [top-level]", key.Tags)
|
||||
}
|
||||
if len(key.Guardrails) != 1 || key.Guardrails[0] != "from-metadata" {
|
||||
t.Errorf("Guardrails = %v, want [from-metadata] (null top-level must not shadow)", key.Guardrails)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceKeyReadDropsMissingKeyFromState(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
w.Write([]byte(`{"error":{"message":"Key not found in database","type":"not_found_error","param":"key","code":"404"}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := newKeyResourceData(t, map[string]interface{}{"key_alias": "stale"})
|
||||
d.SetId("deleted-out-of-band")
|
||||
|
||||
diags := resourceKeyRead(context.Background(), d, NewClient(srv.URL, "test-key", true))
|
||||
if diags.HasError() {
|
||||
t.Fatalf("read of a missing key must not error, got: %v", diags)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Errorf("Id = %q, want empty so Terraform plans a recreate", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceKeyReadStillFailsOnNon404Errors(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
w.Write([]byte(`{"error":{"message":"db down"}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := newKeyResourceData(t, map[string]interface{}{"key_alias": "live"})
|
||||
d.SetId("still-exists")
|
||||
|
||||
diags := resourceKeyRead(context.Background(), d, NewClient(srv.URL, "test-key", true))
|
||||
if !diags.HasError() {
|
||||
t.Fatal("a 500 from /key/info must surface as an error, not be treated as a deleted key")
|
||||
}
|
||||
if d.Id() != "still-exists" {
|
||||
t.Errorf("Id = %q, want unchanged on a transient error", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
// fakeKeyProxy serves /key/info from stored metadata and applies /key/update
|
||||
// the way the proxy does: an absent "metadata" keeps the stored map, a
|
||||
// present one replaces it wholesale.
|
||||
type fakeKeyProxy struct {
|
||||
metadata map[string]interface{}
|
||||
updates []map[string]interface{}
|
||||
}
|
||||
|
||||
func (p *fakeKeyProxy) handler() http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/key/info":
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"key": "hash-1",
|
||||
"info": map[string]interface{}{"key_alias": "alias-1", "models": []string{"gpt-4o-mini"}, "metadata": p.metadata},
|
||||
})
|
||||
case "/key/update":
|
||||
var body map[string]interface{}
|
||||
json.NewDecoder(r.Body).Decode(&body)
|
||||
p.updates = append(p.updates, body)
|
||||
if m, ok := body["metadata"].(map[string]interface{}); ok {
|
||||
p.metadata = m
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{"key": "hash-1", "metadata": p.metadata})
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func applyKeyUpdate(t *testing.T, client *Client, stateAttrs map[string]string, config map[string]interface{}) *terraform.InstanceState {
|
||||
t.Helper()
|
||||
r := resourceKey()
|
||||
state := &terraform.InstanceState{ID: "hash-1", Attributes: stateAttrs}
|
||||
diff, err := r.Diff(context.Background(), state, terraform.NewResourceConfigRaw(config), client)
|
||||
if err != nil {
|
||||
t.Fatalf("Diff returned error: %v", err)
|
||||
}
|
||||
if diff == nil {
|
||||
t.Fatalf("expected a non-empty diff between %v and %v", stateAttrs, config)
|
||||
}
|
||||
newState, diags := r.Apply(context.Background(), state, diff, client)
|
||||
if diags.HasError() {
|
||||
t.Fatalf("Apply returned error: %v", diags)
|
||||
}
|
||||
return newState
|
||||
}
|
||||
|
||||
func TestKeyUpdateWithoutMetadataChangePreservesServerMetadata(t *testing.T) {
|
||||
proxy := &fakeKeyProxy{metadata: map[string]interface{}{"a": "1", "server_side": "x", "model_rpm_limit": map[string]interface{}{"gpt-4o-mini": float64(5)}}}
|
||||
srv := httptest.NewServer(proxy.handler())
|
||||
defer srv.Close()
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
newState := applyKeyUpdate(t, client,
|
||||
map[string]string{"key_alias": "alias-1", "max_budget": "10", "metadata.%": "1", "metadata.a": "1"},
|
||||
map[string]interface{}{"key_alias": "alias-1", "max_budget": 20, "metadata": map[string]interface{}{"a": "1"}},
|
||||
)
|
||||
|
||||
if len(proxy.updates) != 1 {
|
||||
t.Fatalf("expected one /key/update call, got %d", len(proxy.updates))
|
||||
}
|
||||
for _, field := range []string{"metadata", "model_rpm_limit", "model_tpm_limit"} {
|
||||
if _, present := proxy.updates[0][field]; present {
|
||||
t.Errorf("unchanged %q was sent on /key/update: %v", field, proxy.updates[0][field])
|
||||
}
|
||||
}
|
||||
if proxy.metadata["server_side"] != "x" {
|
||||
t.Errorf("server-side metadata lost: %v", proxy.metadata)
|
||||
}
|
||||
if got := newState.Attributes["metadata.%"]; got != "1" {
|
||||
t.Errorf("state metadata should hold only the declared entry, got %v", newState.Attributes)
|
||||
}
|
||||
if got := newState.Attributes["metadata.a"]; got != "1" {
|
||||
t.Errorf("metadata.a = %q, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyUpdateWithMetadataChangeMergesOverServerMetadata(t *testing.T) {
|
||||
proxy := &fakeKeyProxy{metadata: map[string]interface{}{"a": "1", "b": "2", "server_side": "x"}}
|
||||
srv := httptest.NewServer(proxy.handler())
|
||||
defer srv.Close()
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
applyKeyUpdate(t, client,
|
||||
map[string]string{"key_alias": "alias-1", "metadata.%": "2", "metadata.a": "1", "metadata.b": "2"},
|
||||
map[string]interface{}{"key_alias": "alias-1", "metadata": map[string]interface{}{"a": "2", "c": "3"}},
|
||||
)
|
||||
|
||||
want := map[string]interface{}{"a": "2", "c": "3", "server_side": "x"}
|
||||
if !reflect.DeepEqual(proxy.metadata, want) {
|
||||
t.Errorf("metadata after update = %v, want %v", proxy.metadata, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyUpdateSendsChangedModelLimits(t *testing.T) {
|
||||
proxy := &fakeKeyProxy{metadata: map[string]interface{}{}}
|
||||
srv := httptest.NewServer(proxy.handler())
|
||||
defer srv.Close()
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
applyKeyUpdate(t, client,
|
||||
map[string]string{"key_alias": "alias-1", "model_rpm_limit.%": "1", "model_rpm_limit.gpt-4o-mini": "5"},
|
||||
map[string]interface{}{"key_alias": "alias-1", "model_rpm_limit": map[string]interface{}{"gpt-4o-mini": 7}},
|
||||
)
|
||||
|
||||
got, ok := proxy.updates[0]["model_rpm_limit"].(map[string]interface{})
|
||||
if !ok || got["gpt-4o-mini"] != float64(7) {
|
||||
t.Errorf("changed model_rpm_limit not sent: %v", proxy.updates[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyReadKeepsOnlyDeclaredMetadata(t *testing.T) {
|
||||
proxy := &fakeKeyProxy{metadata: map[string]interface{}{"a": "1", "server_side": "x"}}
|
||||
srv := httptest.NewServer(proxy.handler())
|
||||
defer srv.Close()
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
d := newKeyResourceData(t, map[string]interface{}{"metadata": map[string]interface{}{"a": "1"}})
|
||||
d.SetId("hash-1")
|
||||
if diags := resourceKeyRead(context.Background(), d, client); diags.HasError() {
|
||||
t.Fatalf("Read returned error: %v", diags)
|
||||
}
|
||||
|
||||
want := map[string]interface{}{"a": "1"}
|
||||
if got := d.Get("metadata"); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("metadata in state = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyUpdateSendsChangedDuration(t *testing.T) {
|
||||
proxy := &fakeKeyProxy{metadata: map[string]interface{}{}}
|
||||
srv := httptest.NewServer(proxy.handler())
|
||||
defer srv.Close()
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
applyKeyUpdate(t, client,
|
||||
map[string]string{"key_alias": "alias-1", "duration": "30d"},
|
||||
map[string]interface{}{"key_alias": "alias-1", "duration": "90d"},
|
||||
)
|
||||
|
||||
if got := proxy.updates[0]["duration"]; got != "90d" {
|
||||
t.Errorf("update payload duration = %v, want 90d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeyUpdateOmitsUnchangedDuration(t *testing.T) {
|
||||
proxy := &fakeKeyProxy{metadata: map[string]interface{}{}}
|
||||
srv := httptest.NewServer(proxy.handler())
|
||||
defer srv.Close()
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
|
||||
applyKeyUpdate(t, client,
|
||||
map[string]string{"key_alias": "alias-1", "duration": "30d"},
|
||||
map[string]interface{}{"key_alias": "alias-2", "duration": "30d"},
|
||||
)
|
||||
|
||||
if got := proxy.updates[0]["key_alias"]; got != "alias-2" {
|
||||
t.Fatalf("update payload key_alias = %v, want alias-2", got)
|
||||
}
|
||||
if v, present := proxy.updates[0]["duration"]; present {
|
||||
t.Errorf("update payload unexpectedly contains duration = %v", v)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -65,7 +65,7 @@ func buildKeyData(d *schema.ResourceData) map[string]interface{} {
|
|||
keyData["permissions"] = v.(map[string]interface{})
|
||||
}
|
||||
if v, ok := d.GetOkExists("model_max_budget"); ok {
|
||||
keyData["model_max_budget"] = v.(map[string]interface{})
|
||||
keyData["model_max_budget"] = parseKeyModelMaxBudget(v.(string))
|
||||
}
|
||||
if v, ok := d.GetOkExists("model_rpm_limit"); ok {
|
||||
keyData["model_rpm_limit"] = v.(map[string]interface{})
|
||||
|
|
@ -107,7 +107,7 @@ func setKeyResourceData(d *schema.ResourceData, key *Key) error {
|
|||
"aliases": key.Aliases,
|
||||
"config": key.Config,
|
||||
"permissions": key.Permissions,
|
||||
"model_max_budget": key.ModelMaxBudget,
|
||||
"model_max_budget": keyModelMaxBudgetJSON(key.ModelMaxBudget),
|
||||
"model_rpm_limit": key.ModelRPMLimit,
|
||||
"model_tpm_limit": key.ModelTPMLimit,
|
||||
"guardrails": key.Guardrails,
|
||||
|
|
|
|||
|
|
@ -31,6 +31,13 @@ func ResourceLiteLLMTeam() *schema.Resource {
|
|||
},
|
||||
|
||||
Schema: map[string]*schema.Schema{
|
||||
"team_id": {
|
||||
Type: schema.TypeString,
|
||||
Optional: true,
|
||||
Computed: true,
|
||||
ForceNew: true,
|
||||
Description: "Unique ID for the team. Generated by the provider if not provided",
|
||||
},
|
||||
"team_alias": {
|
||||
Type: schema.TypeString,
|
||||
Required: true,
|
||||
|
|
@ -162,7 +169,7 @@ func ResourceLiteLLMTeam() *schema.Resource {
|
|||
func resourceLiteLLMTeamCreate(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
|
||||
teamID := uuid.New().String()
|
||||
teamID := resolveTeamID(d)
|
||||
teamData := buildTeamData(d, teamID)
|
||||
|
||||
// Throughput limit types are only accepted by /team/new, not /team/update.
|
||||
|
|
@ -214,6 +221,7 @@ func resourceLiteLLMTeamRead(d *schema.ResourceData, m interface{}) error {
|
|||
teamResp := infoResp.TeamInfo
|
||||
|
||||
// Update the state with values from the response or fall back to the data passed in during creation
|
||||
d.Set("team_id", d.Id())
|
||||
d.Set("team_alias", GetStringValue(teamResp.TeamAlias, d.Get("team_alias").(string)))
|
||||
d.Set("organization_id", GetStringValue(teamResp.OrganizationID, d.Get("organization_id").(string)))
|
||||
|
||||
|
|
@ -263,11 +271,11 @@ func resourceLiteLLMTeamRead(d *schema.ResourceData, m interface{}) error {
|
|||
d.Set("team_member_tpm_limit", *teamResp.TeamMemberTPMLimit)
|
||||
}
|
||||
d.Set("team_member_key_duration", GetStringValue(teamResp.TeamMemberKeyDuration, d.Get("team_member_key_duration").(string)))
|
||||
if teamResp.ModelRPMLimit != nil {
|
||||
d.Set("model_rpm_limit", teamResp.ModelRPMLimit)
|
||||
if v := teamModelLimit(teamResp.ModelRPMLimit, teamResp.Metadata, "model_rpm_limit"); v != nil {
|
||||
d.Set("model_rpm_limit", v)
|
||||
}
|
||||
if teamResp.ModelTPMLimit != nil {
|
||||
d.Set("model_tpm_limit", teamResp.ModelTPMLimit)
|
||||
if v := teamModelLimit(teamResp.ModelTPMLimit, teamResp.Metadata, "model_tpm_limit"); v != nil {
|
||||
d.Set("model_tpm_limit", v)
|
||||
}
|
||||
if teamResp.AllowedPassthroughRoutes != nil {
|
||||
d.Set("allowed_passthrough_routes", teamResp.AllowedPassthroughRoutes)
|
||||
|
|
@ -354,6 +362,13 @@ func resourceLiteLLMTeamDelete(d *schema.ResourceData, m interface{}) error {
|
|||
return nil
|
||||
}
|
||||
|
||||
func resolveTeamID(d *schema.ResourceData) string {
|
||||
if v, ok := d.GetOk("team_id"); ok {
|
||||
return v.(string)
|
||||
}
|
||||
return uuid.New().String()
|
||||
}
|
||||
|
||||
func buildTeamData(d *schema.ResourceData, teamID string) map[string]interface{} {
|
||||
teamData := map[string]interface{}{
|
||||
"team_id": teamID,
|
||||
|
|
@ -364,14 +379,19 @@ func buildTeamData(d *schema.ResourceData, teamID string) map[string]interface{}
|
|||
"organization_id", "tpm_limit", "rpm_limit", "max_budget", "budget_duration", "models",
|
||||
"blocked", "team_member_permissions", "model_aliases", "guardrails", "prompts",
|
||||
"team_member_budget", "team_member_budget_duration", "team_member_rpm_limit",
|
||||
"team_member_tpm_limit", "team_member_key_duration", "model_rpm_limit",
|
||||
"model_tpm_limit", "allowed_passthrough_routes",
|
||||
"team_member_tpm_limit", "team_member_key_duration", "allowed_passthrough_routes",
|
||||
} {
|
||||
if v, ok := d.GetOk(key); ok {
|
||||
teamData[key] = v
|
||||
}
|
||||
}
|
||||
|
||||
for _, key := range []string{"model_rpm_limit", "model_tpm_limit"} {
|
||||
if v, ok := d.GetOk(key); ok || d.HasChange(key) {
|
||||
teamData[key] = v
|
||||
}
|
||||
}
|
||||
|
||||
if v, ok := d.GetOk("soft_budget"); ok {
|
||||
teamData["soft_budget"] = v
|
||||
} else if d.HasChange("soft_budget") {
|
||||
|
|
@ -404,6 +424,14 @@ func buildTeamMetadata(d *schema.ResourceData) map[string]interface{} {
|
|||
return metadata
|
||||
}
|
||||
|
||||
func teamModelLimit(topLevel, metadata map[string]interface{}, key string) map[string]interface{} {
|
||||
if topLevel != nil {
|
||||
return topLevel
|
||||
}
|
||||
nested, _ := metadata[key].(map[string]interface{})
|
||||
return nested
|
||||
}
|
||||
|
||||
func splitTeamMetadata(raw map[string]interface{}) (map[string]string, []string, []string) {
|
||||
metadata := map[string]string{}
|
||||
var tags, alertEmails []string
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import (
|
|||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/terraform"
|
||||
)
|
||||
|
|
@ -85,6 +86,87 @@ func TestTeamCreateSendsSoftBudgetTagsAndAlertEmails(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestTeamCreateSendsConfiguredTeamID(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := newTeamTestServer(t, &captured, `{"team_id":"platform-team","team_info":{"team_id":"platform-team","team_alias":"platform"},"keys":[],"team_memberships":[]}`)
|
||||
defer srv.Close()
|
||||
|
||||
d := newTeamResourceData(t, map[string]interface{}{
|
||||
"team_id": "platform-team",
|
||||
"team_alias": "platform",
|
||||
})
|
||||
|
||||
if err := resourceLiteLLMTeamCreate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
if got := captured["team_id"]; got != "platform-team" {
|
||||
t.Fatalf("payload team_id = %v, want platform-team", got)
|
||||
}
|
||||
if got := d.Id(); got != "platform-team" {
|
||||
t.Fatalf("resource id = %q, want platform-team", got)
|
||||
}
|
||||
if got := d.Get("team_id"); got != "platform-team" {
|
||||
t.Fatalf("state team_id = %v, want platform-team", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTeamCreateGeneratesTeamIDWhenUnset(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := newTeamTestServer(t, &captured, `{"team_id":"x","team_info":{"team_alias":"eng"},"keys":[],"team_memberships":[]}`)
|
||||
defer srv.Close()
|
||||
|
||||
d := newTeamResourceData(t, map[string]interface{}{"team_alias": "eng"})
|
||||
|
||||
if err := resourceLiteLLMTeamCreate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("create failed: %v", err)
|
||||
}
|
||||
|
||||
sent, _ := captured["team_id"].(string)
|
||||
if _, err := uuid.Parse(sent); err != nil {
|
||||
t.Fatalf("payload team_id = %q, want a generated UUID: %v", sent, err)
|
||||
}
|
||||
if d.Id() != sent || d.Get("team_id") != sent {
|
||||
t.Fatalf("id = %q, state team_id = %v, want both to equal the sent id %q", d.Id(), d.Get("team_id"), sent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTeamReadSetsTeamIDFromResourceID(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := newTeamTestServer(t, &captured, `{"team_id":"imported-team","team_info":{"team_id":"imported-team","team_alias":"imported"},"keys":[],"team_memberships":[]}`)
|
||||
defer srv.Close()
|
||||
|
||||
d := newTeamResourceData(t, map[string]interface{}{})
|
||||
d.SetId("imported-team")
|
||||
|
||||
if err := resourceLiteLLMTeamRead(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("read failed: %v", err)
|
||||
}
|
||||
if got := d.Get("team_id"); got != "imported-team" {
|
||||
t.Fatalf("team_id = %v, want imported-team", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTeamIDChangeForcesReplacement(t *testing.T) {
|
||||
res := ResourceLiteLLMTeam()
|
||||
priorData := schema.TestResourceDataRaw(t, res.Schema, map[string]interface{}{
|
||||
"team_id": "old-team",
|
||||
"team_alias": "eng",
|
||||
})
|
||||
priorData.SetId("old-team")
|
||||
config := terraform.NewResourceConfigRaw(map[string]interface{}{
|
||||
"team_id": "new-team",
|
||||
"team_alias": "eng",
|
||||
})
|
||||
diff, err := res.Diff(context.Background(), priorData.State(), config, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("diff failed: %v", err)
|
||||
}
|
||||
if diff == nil || !diff.RequiresNew() {
|
||||
t.Fatalf("changing team_id must force replacement, diff = %+v", diff)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTeamReadMapsTeamInfoEnvelope(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := newTeamTestServer(t, &captured, teamInfoWithSoftBudget)
|
||||
|
|
@ -250,6 +332,77 @@ func TestTeamReadMapsNewFields(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestTeamReadMapsPerModelLimitsFromMetadata(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := newTeamTestServer(t, &captured, `{
|
||||
"team_id": "team-1",
|
||||
"team_info": {
|
||||
"team_id": "team-1",
|
||||
"team_alias": "eng",
|
||||
"model_rpm_limit": null,
|
||||
"model_tpm_limit": null,
|
||||
"metadata": {
|
||||
"department": "eng",
|
||||
"model_rpm_limit": {"gpt-4o-mini": 250},
|
||||
"model_tpm_limit": {"gpt-4o-mini": 5000}
|
||||
}
|
||||
}
|
||||
}`)
|
||||
defer srv.Close()
|
||||
|
||||
d := newTeamResourceData(t, map[string]interface{}{
|
||||
"team_alias": "eng",
|
||||
"model_rpm_limit": map[string]interface{}{"gpt-4o-mini": 100},
|
||||
})
|
||||
d.SetId("team-1")
|
||||
|
||||
if err := resourceLiteLLMTeamRead(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("read returned error: %v", err)
|
||||
}
|
||||
if got := d.Get("model_rpm_limit"); !reflect.DeepEqual(got, map[string]interface{}{"gpt-4o-mini": 250}) {
|
||||
t.Errorf("model_rpm_limit = %v, want server value 250", got)
|
||||
}
|
||||
if got := d.Get("model_tpm_limit"); !reflect.DeepEqual(got, map[string]interface{}{"gpt-4o-mini": 5000}) {
|
||||
t.Errorf("model_tpm_limit = %v, want server value 5000", got)
|
||||
}
|
||||
if got := d.Get("metadata"); !reflect.DeepEqual(got, map[string]interface{}{"department": "eng"}) {
|
||||
t.Errorf("metadata = %v, want per-model limits kept out of the string map", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTeamUpdateClearsRemovedPerModelLimits(t *testing.T) {
|
||||
var captured map[string]interface{}
|
||||
srv := newTeamTestServer(t, &captured, `{"team_id":"team-1","team_info":{"team_id":"team-1","team_alias":"eng"}}`)
|
||||
defer srv.Close()
|
||||
|
||||
res := ResourceLiteLLMTeam()
|
||||
priorData := schema.TestResourceDataRaw(t, res.Schema, map[string]interface{}{
|
||||
"team_alias": "eng",
|
||||
"model_rpm_limit": map[string]interface{}{"gpt-4o-mini": 100},
|
||||
"model_tpm_limit": map[string]interface{}{"gpt-4o-mini": 5000},
|
||||
})
|
||||
priorData.SetId("team-1")
|
||||
prior := priorData.State()
|
||||
config := terraform.NewResourceConfigRaw(map[string]interface{}{"team_alias": "eng"})
|
||||
diff, err := res.Diff(context.Background(), prior, config, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("diff failed: %v", err)
|
||||
}
|
||||
d, err := schema.InternalMap(res.Schema).Data(prior, diff)
|
||||
if err != nil {
|
||||
t.Fatalf("data failed: %v", err)
|
||||
}
|
||||
|
||||
if err := resourceLiteLLMTeamUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
for _, k := range []string{"model_rpm_limit", "model_tpm_limit"} {
|
||||
if got, ok := captured[k]; !ok || !reflect.DeepEqual(got, map[string]interface{}{}) {
|
||||
t.Errorf("payload %s = %v (present=%v), want explicit empty map", k, got, ok)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// rpm_limit_type / tpm_limit_type are accepted by /team/new but not
|
||||
// /team/update, so create must send them and update must not.
|
||||
func TestTeamLimitTypesSentOnCreateOnly(t *testing.T) {
|
||||
|
|
|
|||
|
|
@ -181,13 +181,15 @@ quota_management.<behavior>.<variant>.<assertion>
|
|||
| team_multi_window | fallback | spend_counter
|
||||
<spend_tracking> chat_completions | stream | messages_bridge | embeddings
|
||||
| cache_hit | key_rollup | concurrent_burst | tags | end_user
|
||||
| per_model | failure | spend_calculate | pagination
|
||||
| per_model | failure | spend_calculate | pagination | key_attribution
|
||||
assertion : blocks_over_limit | resets_after_window | headers_report_remaining | picks_under_tpm
|
||||
| blocks_then_resets | resets_windows_independently | alerts_without_blocking
|
||||
| isolates_per_model | isolates_per_member | isolates_per_group | enforced_across_keys
|
||||
| routes_to_fallback | reseed_matches_db | reports_spend | logs_cost | zero_cost
|
||||
| matches_sum_of_logs | loses_no_spend | attributes_spend | writes_own_rows
|
||||
| writes_failure_row | returns_cost | keeps_total
|
||||
| writes_failure_row | returns_cost | keeps_total | joins_key | reports_alias_and_email
|
||||
| health_rows_keep_service_account | retrieve_batch_cost_joins_retrieving_key
|
||||
| poller_batch_cost_joins_creating_key
|
||||
e.g. quota_management.ratelimit.rpm.blocks_over_limit exercised_on=[chat_completions, messages]
|
||||
quota_management.budget.key.blocks_over_limit exercised_on=[chat_completions]
|
||||
```
|
||||
|
|
|
|||
|
|
@ -58,3 +58,8 @@
|
|||
- {id: quota_management.spend_tracking.service_tier.bills_tier_rates, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier, assertions: [bills_tier_rates], exercised_on: [chat_completions], source: "cost_calculator.py", rationale: "A priority service_tier call bills input, output, and reasoning at the deployment's *_priority rates and records the tier on the row (#35923, #35925)"}
|
||||
- {id: quota_management.spend_tracking.cost_headers.additive_components, module: quota_management, tier: P1, behavior: spend_tracking, variant: cost_headers, assertions: [additive_components], exercised_on: [chat_completions], source: "proxy/common_request_processing.py", rationale: "The x-litellm-response-cost-* component headers sum to the total, input covers only fresh tokens, and reasoning stays a subset of output (#36965)"}
|
||||
- {id: quota_management.spend_tracking.passthrough_stream.injects_usage_cost, module: quota_management, tier: P1, behavior: spend_tracking, variant: passthrough_stream, assertions: [injects_usage_cost], exercised_on: [openai_passthrough], source: "proxy/pass_through_endpoints/streaming_handler.py", rationale: "With include_cost_in_streaming_usage on, the /openai passthrough's final streaming usage frame carries the proxy-computed cost (#36503). Uncovered: the flag is only settable in litellm_settings, and the shared e2e stack does not turn it on yet"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.joins_key, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [joins_key], exercised_on: [chat_completions, messages, responses, embeddings, batches, files, google_native, rust_control_plane], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "Every spend row a virtual key writes across chat, queued chat, messages, responses, embeddings, the Gemini passthrough, file upload, batch create, and a replayed callback log carries api_key equal to the key's token hash and the key alias, the join the usage APIs depend on; a re-hashed token shows up as an unattributed key-hash-* row (#39568, #39572)"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.reports_alias_and_email, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [reports_alias_and_email], exercised_on: [chat_completions, messages, responses, embeddings, batches, files, google_native, rust_control_plane], source: "proxy/management_endpoints/internal_user_endpoints.py", rationale: "/spend/logs?api_key= returns every one of the key's rows with its alias and /user/daily/activity aggregates them under the key's token with key_alias and user_email; /spend/logs carries no email field, so the email is asserted on daily activity only"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.health_rows_keep_service_account, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [health_rows_keep_service_account], exercised_on: [chat_completions], source: "proxy/health_check.py", rationale: "A /health probe's spend row stays keyed by the literal litellm-internal-health-check service account rather than a hash of it, so health spend never appears as an unattributed key"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.retrieve_batch_cost_joins_retrieving_key, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [retrieve_batch_cost_joins_retrieving_key], exercised_on: [batches], source: "proxy/batches_endpoints/endpoints.py", rationale: "The retrieve that first sees a batch in a terminal state prices it inline and writes its {provider_batch_id}_batch_cost row against the retrieving key, so the batch each run creates is one OpenAI fails at validation within seconds and the test retrieves it by its raw provider id with the same key until it is failed; a raw id is never owned by the CheckBatchCost poller, and the row must carry that key's token hash and alias"}
|
||||
- {id: quota_management.spend_tracking.key_attribution.poller_batch_cost_joins_creating_key, module: quota_management, tier: P1, behavior: spend_tracking, variant: key_attribution, assertions: [poller_batch_cost_joins_creating_key], exercised_on: [batches], source: "enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py", rationale: "The CheckBatchCost poller bills a completed, positive-cost batch created through a unified id against the key that created it, a different writer from the inline retrieve. No test claims this cell yet: OpenAI's completion window is 24h and both e2e stacks boot a fresh Postgres per build, so a completed batch is out of one run's reach and the managed list never shows an earlier run's batch; the cell stays visible as a gap until a run can hand a completed batch to the poller"}
|
||||
|
|
|
|||
|
|
@ -84,6 +84,7 @@ class KeyGenerateBody(BaseModel):
|
|||
|
||||
class KeyGenerateResponse(BaseModel):
|
||||
key: str
|
||||
token: str | None = None
|
||||
key_alias: str | None = None
|
||||
models: list[str] = []
|
||||
max_budget: float | None = None
|
||||
|
|
@ -672,6 +673,7 @@ class GuardrailRunRecord(BaseModel):
|
|||
|
||||
|
||||
class SpendLogMetadata(BaseModel):
|
||||
user_api_key_alias: str | None = None
|
||||
applied_guardrails: list[str] | None = None
|
||||
guardrail_information: list[GuardrailRunRecord] | None = None
|
||||
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ DRIVER_MODELS: tuple[tuple[str, str, str], ...] = (
|
|||
("claude-haiku-4-5", "anthropic/claude-haiku-4-5", "ANTHROPIC_API_KEY"),
|
||||
("openai-text-embedding-3-small", "openai/text-embedding-3-small", "OPENAI_API_KEY"),
|
||||
("openai-responses-codex", "openai/gpt-5.3-codex", "OPENAI_API_KEY"),
|
||||
("openai-gpt-4o-mini", "openai/gpt-4o-mini", "OPENAI_API_KEY"),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -15,9 +15,12 @@ import time
|
|||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import (
|
||||
FileUploadForm,
|
||||
Headers,
|
||||
NoBody,
|
||||
ProbeResult,
|
||||
Result,
|
||||
|
|
@ -35,6 +38,8 @@ from models import (
|
|||
DateRangeParams,
|
||||
EmbedBody,
|
||||
EmbedResponse,
|
||||
KeyGenerateBody,
|
||||
KeyGenerateResponse,
|
||||
OpenAPISchema,
|
||||
SpendCalculateBody,
|
||||
SpendCalculateResponse,
|
||||
|
|
@ -43,13 +48,27 @@ from models import (
|
|||
SpendLogsPageParams,
|
||||
SpendTagsResponse,
|
||||
TagSpend,
|
||||
UserDeleteBody,
|
||||
UserDeleteResponse,
|
||||
UserNewBody,
|
||||
UserNewResponse,
|
||||
UserRole,
|
||||
)
|
||||
from proxy_client import ProxyClient
|
||||
from proxy_client import Converged, ProxyClient, await_converged
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
__all__ = [
|
||||
"BatchCreateBody",
|
||||
"CallbackLogMetadata",
|
||||
"CallbackLogPayload",
|
||||
"BatchObject",
|
||||
"DailyActivityKeyBreakdown",
|
||||
"FileObject",
|
||||
"ProbeResult",
|
||||
"ResponseIdentity",
|
||||
"SpendClient",
|
||||
"SpendLogRow",
|
||||
"StreamingResponse",
|
||||
"build_client",
|
||||
"is_ok",
|
||||
"unique_marker",
|
||||
|
|
@ -57,6 +76,139 @@ __all__ = [
|
|||
]
|
||||
|
||||
|
||||
class GeminiApiKeyHeaders(Headers):
|
||||
x_goog_api_key: str = Field(serialization_alias="x-goog-api-key")
|
||||
content_type: str = Field(default="application/json", serialization_alias="Content-Type")
|
||||
|
||||
|
||||
class GeminiPart(BaseModel):
|
||||
text: str
|
||||
|
||||
|
||||
class GeminiContent(BaseModel):
|
||||
parts: list[GeminiPart]
|
||||
|
||||
|
||||
class GeminiGenerationConfig(BaseModel):
|
||||
maxOutputTokens: int
|
||||
|
||||
|
||||
class GeminiGenerateBody(BaseModel):
|
||||
contents: list[GeminiContent]
|
||||
generationConfig: GeminiGenerationConfig
|
||||
|
||||
|
||||
class ResponsesBody(BaseModel):
|
||||
model: str
|
||||
input: str
|
||||
cache: dict[str, bool] | None = {"no-cache": True}
|
||||
|
||||
|
||||
class QueuedChatBody(ChatBody):
|
||||
priority: int = 0
|
||||
|
||||
|
||||
class ResponseIdentity(BaseModel):
|
||||
id: str | None = None
|
||||
|
||||
|
||||
class HealthParams(BaseModel):
|
||||
model: str
|
||||
|
||||
|
||||
class ModelQuery(BaseModel):
|
||||
model: str
|
||||
|
||||
|
||||
class FileObject(BaseModel):
|
||||
id: str
|
||||
|
||||
|
||||
class BatchCreateBody(BaseModel):
|
||||
input_file_id: str
|
||||
endpoint: str = "/v1/chat/completions"
|
||||
completion_window: str = "24h"
|
||||
model: str
|
||||
metadata: dict[str, str]
|
||||
|
||||
|
||||
class BatchObject(BaseModel):
|
||||
id: str
|
||||
status: str
|
||||
|
||||
|
||||
class ProviderQuery(BaseModel):
|
||||
provider: str
|
||||
|
||||
|
||||
class CallbackLogMetadata(BaseModel):
|
||||
user_api_key_hash: str
|
||||
user_api_key_alias: str
|
||||
user_api_key_user_id: str
|
||||
|
||||
|
||||
class CallbackLogPayload(BaseModel):
|
||||
id: str
|
||||
litellm_call_id: str
|
||||
model: str
|
||||
call_type: str = "acompletion"
|
||||
start_time: float = Field(serialization_alias="startTime")
|
||||
end_time: float = Field(serialization_alias="endTime")
|
||||
response_cost: float
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
total_tokens: int
|
||||
metadata: CallbackLogMetadata
|
||||
|
||||
|
||||
class CallbackLogRecord(BaseModel):
|
||||
status: str = "success"
|
||||
standard_logging_payload: CallbackLogPayload
|
||||
|
||||
|
||||
class CallbackLogsRequest(BaseModel):
|
||||
records: list[CallbackLogRecord]
|
||||
|
||||
|
||||
class CallbackLogsResponse(BaseModel):
|
||||
processed: int
|
||||
failed: int
|
||||
|
||||
|
||||
class DailyActivityParams(BaseModel):
|
||||
start_date: str
|
||||
end_date: str
|
||||
api_key: str
|
||||
|
||||
|
||||
class DailyActivityKeyMetadata(BaseModel):
|
||||
key_alias: str | None = None
|
||||
team_id: str | None = None
|
||||
user_email: str | None = None
|
||||
|
||||
|
||||
class DailyActivityKeyMetrics(BaseModel):
|
||||
api_requests: int = 0
|
||||
|
||||
|
||||
class DailyActivityKeyBreakdown(BaseModel):
|
||||
metrics: DailyActivityKeyMetrics
|
||||
metadata: DailyActivityKeyMetadata
|
||||
|
||||
|
||||
class DailyActivityBreakdown(BaseModel):
|
||||
api_keys: dict[str, DailyActivityKeyBreakdown] = {}
|
||||
|
||||
|
||||
class DailyActivityRow(BaseModel):
|
||||
date: str
|
||||
breakdown: DailyActivityBreakdown
|
||||
|
||||
|
||||
class DailyActivityResponse(BaseModel):
|
||||
results: list[DailyActivityRow] = []
|
||||
|
||||
|
||||
def _chat_body(
|
||||
model: str,
|
||||
content: str,
|
||||
|
|
@ -207,6 +359,166 @@ class SpendClient:
|
|||
def probe(self, path: str, *, params: DateRangeParams) -> ProbeResult:
|
||||
return self.proxy.transport.probe(path, params=params)
|
||||
|
||||
def create_user(self, *, email: str, role: UserRole, user_id: str) -> str:
|
||||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/user/new",
|
||||
headers=self.proxy.transport.master,
|
||||
json=UserNewBody(user_email=email, user_role=role, user_id=user_id),
|
||||
response_type=UserNewResponse,
|
||||
)
|
||||
).user_id
|
||||
|
||||
def delete_user(self, user_id: str) -> None:
|
||||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/user/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
json=UserDeleteBody(user_ids=[user_id]),
|
||||
response_type=UserDeleteResponse,
|
||||
)
|
||||
)
|
||||
|
||||
def generate_key_record(self, body: KeyGenerateBody) -> KeyGenerateResponse:
|
||||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/key/generate",
|
||||
headers=self.proxy.transport.master,
|
||||
json=body,
|
||||
response_type=KeyGenerateResponse,
|
||||
)
|
||||
)
|
||||
|
||||
def send_chat(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse:
|
||||
return self.proxy.transport.send(
|
||||
"/chat/completions",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=_chat_body(model, content, max_tokens=max_tokens),
|
||||
)
|
||||
|
||||
def send_queued_chat(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse:
|
||||
return self.proxy.transport.send(
|
||||
"/queue/chat/completions",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=QueuedChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=content)],
|
||||
max_tokens=max_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
def send_messages(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse:
|
||||
return self.proxy.transport.send(
|
||||
"/v1/messages",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=AnthropicMessagesBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=content)],
|
||||
max_tokens=max_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
def send_responses(self, key: str, model: str, content: str) -> StreamingResponse:
|
||||
return self.proxy.transport.send(
|
||||
"/v1/responses",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=ResponsesBody(model=model, input=content),
|
||||
)
|
||||
|
||||
def send_embed(self, key: str, model: str, content: str) -> StreamingResponse:
|
||||
return self.proxy.transport.send(
|
||||
"/embeddings",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=EmbedBody(model=model, input=content),
|
||||
)
|
||||
|
||||
def send_gemini_generate(self, key: str, model: str, content: str, *, max_tokens: int) -> StreamingResponse:
|
||||
return self.proxy.transport.send(
|
||||
f"/gemini/v1beta/models/{model}:generateContent",
|
||||
headers=GeminiApiKeyHeaders(x_goog_api_key=key),
|
||||
json=GeminiGenerateBody(
|
||||
contents=[GeminiContent(parts=[GeminiPart(text=content)])],
|
||||
generationConfig=GeminiGenerationConfig(maxOutputTokens=max_tokens),
|
||||
),
|
||||
)
|
||||
|
||||
def upload_batch_file(self, key: str, model: str, content: bytes) -> FileObject:
|
||||
return unwrap(
|
||||
self.proxy.transport.upload(
|
||||
"/v1/files",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
form=FileUploadForm(purpose="batch"),
|
||||
filename="key_attribution.jsonl",
|
||||
content=content,
|
||||
params=ModelQuery(model=model),
|
||||
response_type=FileObject,
|
||||
)
|
||||
)
|
||||
|
||||
def create_batch(self, key: str, body: BatchCreateBody) -> BatchObject:
|
||||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/v1/batches",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=body,
|
||||
response_type=BatchObject,
|
||||
)
|
||||
)
|
||||
|
||||
def retrieve_batch(self, key: str, batch_id: str, *, provider: str) -> BatchObject:
|
||||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
f"/v1/batches/{batch_id}",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
params=ProviderQuery(provider=provider),
|
||||
response_type=BatchObject,
|
||||
)
|
||||
)
|
||||
|
||||
def replay_callback_log(self, key: str, payload: CallbackLogPayload) -> CallbackLogsResponse:
|
||||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/v1/rust_control_plane/logs",
|
||||
headers=self.proxy.transport.bearer(key),
|
||||
json=CallbackLogsRequest(records=[CallbackLogRecord(standard_logging_payload=payload)]),
|
||||
response_type=CallbackLogsResponse,
|
||||
)
|
||||
)
|
||||
|
||||
def health(self, model: str) -> ProbeResult:
|
||||
return self.proxy.transport.probe("/health", params=HealthParams(model=model))
|
||||
|
||||
def daily_activity_for_key(self, token: str, *, start: datetime, end: datetime) -> DailyActivityKeyBreakdown | None:
|
||||
response: Final = unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/user/daily/activity",
|
||||
headers=self.proxy.transport.master,
|
||||
params=DailyActivityParams(
|
||||
start_date=start.strftime("%Y-%m-%d"),
|
||||
end_date=end.strftime("%Y-%m-%d"),
|
||||
api_key=token,
|
||||
),
|
||||
response_type=DailyActivityResponse,
|
||||
)
|
||||
)
|
||||
return next(
|
||||
(row.breakdown.api_keys[token] for row in response.results if token in row.breakdown.api_keys),
|
||||
None,
|
||||
)
|
||||
|
||||
def poll_daily_activity_for_key(
|
||||
self, token: str, *, start: datetime, end: datetime, min_requests: int
|
||||
) -> DailyActivityKeyBreakdown | None:
|
||||
outcome: Final = await_converged(
|
||||
lambda: self.daily_activity_for_key(token, start=start, end=end),
|
||||
converged=lambda found: found is not None and found.metrics.api_requests >= min_requests,
|
||||
timeout=self.proxy.poll_timeout,
|
||||
interval=self.proxy.poll_interval,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
return outcome.result if isinstance(outcome, Converged) else outcome.last_result
|
||||
|
||||
def openapi(self) -> OpenAPISchema:
|
||||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,405 @@
|
|||
"""Every spend row a live proxy writes joins its virtual key (MAT-180).
|
||||
|
||||
One virtual key with an alias, owned by a user with an email, drives every spend
|
||||
write path a key can reach: /chat/completions, /queue/chat/completions,
|
||||
/v1/messages, /v1/responses, /embeddings, the Gemini native passthrough, a batch
|
||||
input file upload, a batch create, and a replayed callback log (POST
|
||||
/v1/rust_control_plane/logs, the writer an external gateway feeds). Each row those calls write must carry
|
||||
`api_key` equal to the key's LiteLLM_VerificationToken.token (the sha256 hash
|
||||
/key/generate returns as `token`), which is the join /spend/logs?api_key= and
|
||||
/user/daily/activity rely on to report key_alias and user_email. A row keyed by a
|
||||
re-hashed token (v1.99.0's regression, #39568 and #39572) shows up as a
|
||||
key-hash-* row with no alias and no email in the customer's usage exports.
|
||||
|
||||
The health-check service account writes rows too; those must stay keyed by the
|
||||
literal service-account name, never by a hash of it. A batch's cost row is
|
||||
written by the retrieve that first sees the batch in a terminal state, so the
|
||||
batch the run creates is one OpenAI fails at validation within seconds (its one
|
||||
line targets /v1/embeddings under a /v1/chat/completions batch), and the test
|
||||
retrieves it by its raw provider id with the same key until it is failed. A raw
|
||||
id is never owned by the CheckBatchCost poller, so that retrieve prices the batch
|
||||
inline against the retrieving key and its {provider_batch_id}_batch_cost row
|
||||
must join the key's token with its alias. A completed batch with a positive
|
||||
cost is out of a single run's reach (OpenAI's completion window is 24h, and a
|
||||
stack booted fresh per run lists no earlier run's batches), so the poller's own
|
||||
row is not asserted here.
|
||||
|
||||
/spend/logs carries no email field, so the email assertion lives on
|
||||
/user/daily/activity alone; /spend/logs is held to the alias in metadata.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from models import KeyGenerateBody
|
||||
from proxy_client import Converged, await_converged
|
||||
from pydantic import BaseModel
|
||||
from spend_e2e_client import (
|
||||
BatchCreateBody,
|
||||
BatchObject,
|
||||
CallbackLogMetadata,
|
||||
CallbackLogPayload,
|
||||
DailyActivityKeyBreakdown,
|
||||
ResponseIdentity,
|
||||
SpendClient,
|
||||
SpendLogRow,
|
||||
StreamingResponse,
|
||||
unique_marker,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
CHAT_MODEL: Final = "gemini-2.5-flash"
|
||||
MESSAGES_MODEL: Final = "claude-haiku-4-5"
|
||||
RESPONSES_MODEL: Final = "openai-responses-codex"
|
||||
EMBED_MODEL: Final = "openai-text-embedding-3-small"
|
||||
BATCH_MODEL: Final = "openai-gpt-4o-mini"
|
||||
BATCH_BACKEND_MODEL: Final = "gpt-4o-mini"
|
||||
BATCH_PROVIDER: Final = "openai"
|
||||
HEALTH_SERVICE_ACCOUNT: Final = "litellm-internal-health-check"
|
||||
BATCH_TERMINAL_STATUSES: Final = frozenset({"completed", "failed", "cancelled", "expired"})
|
||||
FAILED_BATCH_POLL_SECONDS: Final = 120.0
|
||||
FAILED_BATCH_POLL_INTERVAL_SECONDS: Final = 5.0
|
||||
MAX_TOKENS: Final = 8
|
||||
REPLAY_RESPONSE_COST: Final = 0.0001
|
||||
REPLAY_PROMPT_TOKENS: Final = 5
|
||||
REPLAY_COMPLETION_TOKENS: Final = 1
|
||||
WRITE_PATHS: Final = (
|
||||
"chat_completions",
|
||||
"queue_chat_completions",
|
||||
"messages",
|
||||
"responses",
|
||||
"embeddings",
|
||||
"gemini_passthrough",
|
||||
"batch_file_upload",
|
||||
"batch_create",
|
||||
"callback_replay",
|
||||
)
|
||||
|
||||
|
||||
class EmbeddingLineBody(BaseModel):
|
||||
model: str
|
||||
input: str
|
||||
|
||||
|
||||
class EmbeddingLine(BaseModel):
|
||||
custom_id: str
|
||||
method: str = "POST"
|
||||
url: str = "/v1/embeddings"
|
||||
body: EmbeddingLineBody
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AttributedKey:
|
||||
key: str
|
||||
token: str
|
||||
alias: str
|
||||
email: str
|
||||
user_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WritePath:
|
||||
name: str
|
||||
request_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DrivenKey:
|
||||
identity: AttributedKey
|
||||
paths: tuple[WritePath, ...]
|
||||
started_at: datetime
|
||||
|
||||
|
||||
def _body_id(name: str, sent: StreamingResponse) -> WritePath:
|
||||
assert sent.ok, f"{name} failed with {sent.status_code}: {sent.body[:300]}"
|
||||
response_id: Final = ResponseIdentity.model_validate_json(sent.body).id
|
||||
assert response_id, f"{name} answered without a response id: {sent.body[:300]}"
|
||||
return WritePath(name=name, request_id=response_id)
|
||||
|
||||
|
||||
def _call_id(name: str, sent: StreamingResponse) -> WritePath:
|
||||
assert sent.ok, f"{name} failed with {sent.status_code}: {sent.body[:300]}"
|
||||
assert sent.call_id, f"{name} answered without an x-litellm-call-id header"
|
||||
return WritePath(name=name, request_id=sent.call_id)
|
||||
|
||||
|
||||
def _endpoint_mismatched_jsonl(marker: str) -> bytes:
|
||||
line: Final = EmbeddingLine(custom_id=marker, body=EmbeddingLineBody(model=BATCH_BACKEND_MODEL, input=marker))
|
||||
return f"{line.model_dump_json()}\n".encode()
|
||||
|
||||
|
||||
def _drive_batch(client: SpendClient, identity: AttributedKey, marker: str) -> tuple[WritePath, WritePath]:
|
||||
uploaded: Final = client.upload_batch_file(identity.key, BATCH_MODEL, _endpoint_mismatched_jsonl(marker))
|
||||
created: Final = client.create_batch(
|
||||
identity.key,
|
||||
BatchCreateBody(
|
||||
input_file_id=uploaded.id,
|
||||
model=BATCH_MODEL,
|
||||
metadata={"run": marker},
|
||||
),
|
||||
)
|
||||
return (
|
||||
WritePath(name="batch_file_upload", request_id=uploaded.id),
|
||||
WritePath(name="batch_create", request_id=created.id),
|
||||
)
|
||||
|
||||
|
||||
def _drive_callback_replay(client: SpendClient, identity: AttributedKey, marker: str) -> WritePath:
|
||||
request_id: Final = f"callback-replay-{marker}"
|
||||
finished_at: Final = time.time()
|
||||
replayed: Final = client.replay_callback_log(
|
||||
identity.key,
|
||||
CallbackLogPayload(
|
||||
id=request_id,
|
||||
litellm_call_id=request_id,
|
||||
model=CHAT_MODEL,
|
||||
start_time=finished_at - 1,
|
||||
end_time=finished_at,
|
||||
response_cost=REPLAY_RESPONSE_COST,
|
||||
prompt_tokens=REPLAY_PROMPT_TOKENS,
|
||||
completion_tokens=REPLAY_COMPLETION_TOKENS,
|
||||
total_tokens=REPLAY_PROMPT_TOKENS + REPLAY_COMPLETION_TOKENS,
|
||||
metadata=CallbackLogMetadata(
|
||||
user_api_key_hash=identity.token,
|
||||
user_api_key_alias=identity.alias,
|
||||
user_api_key_user_id=identity.user_id,
|
||||
),
|
||||
),
|
||||
)
|
||||
assert replayed.processed == 1 and replayed.failed == 0, f"callback replay rejected the payload: {replayed}"
|
||||
return WritePath(name="callback_replay", request_id=request_id)
|
||||
|
||||
|
||||
def _drive_every_write_path(client: SpendClient, identity: AttributedKey) -> tuple[WritePath, ...]:
|
||||
marker: Final = unique_marker()
|
||||
prompt: Final = f"Reply with the word ok. {marker}"
|
||||
key: Final = identity.key
|
||||
return (
|
||||
_body_id("chat_completions", client.send_chat(key, CHAT_MODEL, prompt, max_tokens=MAX_TOKENS)),
|
||||
_body_id("queue_chat_completions", client.send_queued_chat(key, CHAT_MODEL, prompt, max_tokens=MAX_TOKENS)),
|
||||
_body_id("messages", client.send_messages(key, MESSAGES_MODEL, prompt, max_tokens=MAX_TOKENS)),
|
||||
_body_id("responses", client.send_responses(key, RESPONSES_MODEL, prompt)),
|
||||
_call_id("embeddings", client.send_embed(key, EMBED_MODEL, prompt)),
|
||||
_call_id("gemini_passthrough", client.send_gemini_generate(key, CHAT_MODEL, prompt, max_tokens=MAX_TOKENS)),
|
||||
*_drive_batch(client, identity, marker),
|
||||
_drive_callback_replay(client, identity, marker),
|
||||
)
|
||||
|
||||
|
||||
def _provider_batch_id(unified_batch_id: str) -> str:
|
||||
encoded: Final = unified_batch_id.removeprefix("batch_")
|
||||
decoded: Final = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4)).decode()
|
||||
return decoded.removeprefix("litellm:").split(";", 1)[0]
|
||||
|
||||
|
||||
def _driven_batch_id(driven: DrivenKey) -> str:
|
||||
return next(path.request_id for path in driven.paths if path.name == "batch_create")
|
||||
|
||||
|
||||
def _await_terminal_batch(client: SpendClient, key: str, provider_batch_id: str) -> BatchObject:
|
||||
outcome: Final = await_converged(
|
||||
lambda: client.retrieve_batch(key, provider_batch_id, provider=BATCH_PROVIDER),
|
||||
converged=lambda batch: batch.status in BATCH_TERMINAL_STATUSES,
|
||||
timeout=FAILED_BATCH_POLL_SECONDS,
|
||||
interval=FAILED_BATCH_POLL_INTERVAL_SECONDS,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
return outcome.result if isinstance(outcome, Converged) else outcome.last_result
|
||||
|
||||
|
||||
def _health_rows_between(client: SpendClient, started_at: datetime) -> list[SpendLogRow]:
|
||||
return [
|
||||
row
|
||||
for row in client.proxy.spend_logs_window(
|
||||
start=started_at - timedelta(minutes=1), end=datetime.now(timezone.utc) + timedelta(minutes=1)
|
||||
)
|
||||
if HEALTH_SERVICE_ACCOUNT in (row.request_tags or [])
|
||||
]
|
||||
|
||||
|
||||
def _health_rows_since(client: SpendClient, started_at: datetime) -> list[SpendLogRow]:
|
||||
outcome: Final = await_converged(
|
||||
lambda: _health_rows_between(client, started_at),
|
||||
converged=lambda rows: bool(rows),
|
||||
timeout=client.proxy.poll_timeout,
|
||||
interval=client.proxy.poll_interval,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
return outcome.result if isinstance(outcome, Converged) else outcome.last_result
|
||||
|
||||
|
||||
class TestKeyAttribution:
|
||||
@pytest.fixture(scope="class")
|
||||
def driven(self, client: SpendClient) -> Iterator[DrivenKey]:
|
||||
marker: Final = unique_marker()
|
||||
user_id: Final = client.create_user(
|
||||
email=f"key-attribution-{marker}@example.com",
|
||||
role="proxy_admin",
|
||||
user_id=f"key-attribution-{marker}",
|
||||
)
|
||||
record: Final = client.generate_key_record(
|
||||
KeyGenerateBody(models=[], user_id=user_id, key_alias=f"key-attribution-{marker}")
|
||||
)
|
||||
assert record.token, "/key/generate answered without the key's token hash"
|
||||
assert record.key_alias, "/key/generate dropped the key alias"
|
||||
identity: Final = AttributedKey(
|
||||
key=record.key,
|
||||
token=record.token,
|
||||
alias=record.key_alias,
|
||||
email=f"key-attribution-{marker}@example.com",
|
||||
user_id=user_id,
|
||||
)
|
||||
started_at: Final = datetime.now(timezone.utc)
|
||||
try:
|
||||
yield DrivenKey(
|
||||
identity=identity,
|
||||
paths=_drive_every_write_path(client, identity),
|
||||
started_at=started_at,
|
||||
)
|
||||
finally:
|
||||
client.proxy.delete_key(identity.key)
|
||||
client.delete_user(identity.user_id)
|
||||
|
||||
@pytest.mark.covers(
|
||||
"quota_management.spend_tracking.key_attribution.joins_key",
|
||||
exercised_on=[
|
||||
"chat_completions",
|
||||
"messages",
|
||||
"responses",
|
||||
"embeddings",
|
||||
"batches",
|
||||
"files",
|
||||
"google_native",
|
||||
"rust_control_plane",
|
||||
],
|
||||
)
|
||||
def test_every_write_path_row_joins_the_key(self, client: SpendClient, driven: DrivenKey) -> None:
|
||||
assert tuple(path.name for path in driven.paths) == WRITE_PATHS
|
||||
found: Final = tuple((path, client.proxy.poll_logs_for_request_id(path.request_id)) for path in driven.paths)
|
||||
unwritten: Final = [path.name for path, rows in found if not rows]
|
||||
assert not unwritten, f"write paths that produced no spend row within the poll window: {unwritten}"
|
||||
unjoined: Final = [
|
||||
(path.name, row.call_type, row.api_key)
|
||||
for path, rows in found
|
||||
for row in rows
|
||||
if row.api_key != driven.identity.token
|
||||
]
|
||||
assert not unjoined, (
|
||||
"spend rows whose api_key does not join LiteLLM_VerificationToken.token "
|
||||
f"{driven.identity.token}: {unjoined}"
|
||||
)
|
||||
unaliased: Final = [
|
||||
(path.name, row.call_type, row.metadata.user_api_key_alias if row.metadata else None)
|
||||
for path, rows in found
|
||||
for row in rows
|
||||
if row.metadata is None or row.metadata.user_api_key_alias != driven.identity.alias
|
||||
]
|
||||
assert not unaliased, f"spend rows written without key alias {driven.identity.alias!r}: {unaliased}"
|
||||
|
||||
@pytest.mark.covers(
|
||||
"quota_management.spend_tracking.key_attribution.reports_alias_and_email",
|
||||
exercised_on=[
|
||||
"chat_completions",
|
||||
"messages",
|
||||
"responses",
|
||||
"embeddings",
|
||||
"batches",
|
||||
"files",
|
||||
"google_native",
|
||||
"rust_control_plane",
|
||||
],
|
||||
)
|
||||
def test_spend_logs_by_key_return_every_row_with_the_alias(self, client: SpendClient, driven: DrivenKey) -> None:
|
||||
expected_ids: Final = frozenset(path.request_id for path in driven.paths)
|
||||
rows: Final = client.poll_logs_for_key(
|
||||
driven.identity.key,
|
||||
min_rows=len(driven.paths),
|
||||
predicate=lambda found: expected_ids <= frozenset(row.request_id or "" for row in found),
|
||||
)
|
||||
missing: Final = expected_ids - frozenset(row.request_id or "" for row in rows)
|
||||
assert not missing, (
|
||||
f"/spend/logs?api_key= does not return {len(missing)} of {len(expected_ids)} rows for the key: "
|
||||
f"{sorted(path.name for path in driven.paths if path.request_id in missing)}"
|
||||
)
|
||||
aliases: Final = frozenset(row.metadata.user_api_key_alias if row.metadata else None for row in rows)
|
||||
assert aliases == {driven.identity.alias}, f"/spend/logs rows carry aliases {sorted(map(str, aliases))}"
|
||||
|
||||
@pytest.mark.covers(
|
||||
"quota_management.spend_tracking.key_attribution.reports_alias_and_email",
|
||||
exercised_on=[
|
||||
"chat_completions",
|
||||
"messages",
|
||||
"responses",
|
||||
"embeddings",
|
||||
"batches",
|
||||
"files",
|
||||
"google_native",
|
||||
"rust_control_plane",
|
||||
],
|
||||
)
|
||||
def test_user_daily_activity_reports_alias_and_email(self, client: SpendClient, driven: DrivenKey) -> None:
|
||||
breakdown: Final[DailyActivityKeyBreakdown | None] = client.poll_daily_activity_for_key(
|
||||
driven.identity.token,
|
||||
start=driven.started_at - timedelta(days=1),
|
||||
end=datetime.now(timezone.utc) + timedelta(days=1),
|
||||
min_requests=len(driven.paths),
|
||||
)
|
||||
assert breakdown is not None, (
|
||||
f"/user/daily/activity?api_key={driven.identity.token} has no api_keys breakdown: "
|
||||
"the key's rows did not aggregate under its token"
|
||||
)
|
||||
assert breakdown.metrics.api_requests >= len(driven.paths), (
|
||||
f"/user/daily/activity counts {breakdown.metrics.api_requests} requests for the key, "
|
||||
f"expected at least {len(driven.paths)}"
|
||||
)
|
||||
assert breakdown.metadata.key_alias == driven.identity.alias, f"key_alias={breakdown.metadata.key_alias!r}"
|
||||
assert breakdown.metadata.user_email == driven.identity.email, f"user_email={breakdown.metadata.user_email!r}"
|
||||
|
||||
@pytest.mark.covers(
|
||||
"quota_management.spend_tracking.key_attribution.health_rows_keep_service_account",
|
||||
exercised_on=["chat_completions"],
|
||||
)
|
||||
def test_health_check_rows_keep_the_service_account_key(self, client: SpendClient) -> None:
|
||||
started_at: Final = datetime.now(timezone.utc)
|
||||
probe: Final = client.health(CHAT_MODEL)
|
||||
assert probe.healthy, f"/health?model={CHAT_MODEL} answered {probe.status_code}: {probe.body[:300]}"
|
||||
rows: Final = _health_rows_since(client, started_at)
|
||||
assert rows, f"/health?model={CHAT_MODEL} wrote no {HEALTH_SERVICE_ACCOUNT}-tagged spend row"
|
||||
rehashed: Final = [(row.request_id, row.api_key) for row in rows if row.api_key != HEALTH_SERVICE_ACCOUNT]
|
||||
assert not rehashed, f"health-check rows keyed by something other than {HEALTH_SERVICE_ACCOUNT!r}: {rehashed}"
|
||||
|
||||
@pytest.mark.covers(
|
||||
"quota_management.spend_tracking.key_attribution.retrieve_batch_cost_joins_retrieving_key",
|
||||
exercised_on=["batches"],
|
||||
)
|
||||
def test_terminal_batch_cost_row_joins_the_retrieving_key(self, client: SpendClient, driven: DrivenKey) -> None:
|
||||
provider_batch_id: Final = _provider_batch_id(_driven_batch_id(driven))
|
||||
fetched: Final = _await_terminal_batch(client, driven.identity.key, provider_batch_id)
|
||||
assert fetched.status == "failed", (
|
||||
f"endpoint-mismatched batch {provider_batch_id} is {fetched.status!r} after "
|
||||
f"{FAILED_BATCH_POLL_SECONDS:.0f}s, so its terminal cost row cannot be asserted"
|
||||
)
|
||||
cost_request_id: Final = f"{provider_batch_id}_batch_cost"
|
||||
rows: Final = client.proxy.poll_logs_for_request_id(cost_request_id)
|
||||
assert rows, f"retrieving failed batch {provider_batch_id} wrote no cost row under {cost_request_id}"
|
||||
call_types: Final = tuple(sorted({row.call_type or "" for row in rows}))
|
||||
assert call_types == ("aretrieve_batch",), f"cost rows under {cost_request_id} carry call types {call_types}"
|
||||
unjoined: Final = [
|
||||
(row.call_type, row.api_key, row.metadata.user_api_key_alias if row.metadata else None)
|
||||
for row in rows
|
||||
if row.api_key != driven.identity.token
|
||||
or row.metadata is None
|
||||
or row.metadata.user_api_key_alias != driven.identity.alias
|
||||
]
|
||||
assert not unjoined, (
|
||||
f"batch cost rows that do not join the retrieving key's token {driven.identity.token} "
|
||||
f"with alias {driven.identity.alias!r}: {unjoined}"
|
||||
)
|
||||
|
|
@ -23,9 +23,9 @@ from litellm.proxy.proxy_server import token_counter
|
|||
|
||||
def _fake_hf_tokenizer(num_tokens: int) -> MagicMock:
|
||||
encoding = MagicMock()
|
||||
encoding.ids = list(range(num_tokens))
|
||||
encoding.__len__.return_value = num_tokens
|
||||
tokenizer = MagicMock()
|
||||
tokenizer.encode.return_value = encoding
|
||||
tokenizer.encode_batch_fast.return_value = [encoding]
|
||||
return tokenizer
|
||||
|
||||
|
||||
|
|
@ -68,13 +68,11 @@ async def test_custom_tokenizer_from_model_info_is_used(monkeypatch):
|
|||
)
|
||||
)
|
||||
|
||||
mock_tokenizer_cls.from_pretrained.assert_called_once_with(
|
||||
"my-org/custom-tokenizer", revision="v2", auth_token=None
|
||||
)
|
||||
mock_tokenizer_cls.from_pretrained.assert_called_once_with("my-org/custom-tokenizer", revision="v2", token=None)
|
||||
assert response.tokenizer_type == "huggingface_tokenizer"
|
||||
assert response.request_model == "my-embedding-model"
|
||||
assert response.model_used == "self-hosted-embedder"
|
||||
assert response.total_tokens > 0
|
||||
assert response.total_tokens >= 7
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue