diff --git a/.github/workflows/create-release.yml b/.github/workflows/create-release.yml index a726a921a2b..4834775e329 100644 --- a/.github/workflows/create-release.yml +++ b/.github/workflows/create-release.yml @@ -52,6 +52,22 @@ jobs: // are stable maintenance releases, not pre-releases. const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag); + // A stable release should only claim the repo "latest" badge when its + // version is >= the current latest. Otherwise a backport (e.g. 1.84.6) + // would steal "latest" from a newer line (e.g. 1.88.1). + const versionKey = (rawTag) => { + const m = String(rawTag).match(/^v?(\d+)\.(\d+)\.(\d+)/); + if (!m) return null; + const maintenance = String(rawTag).match(/(?:\.post|\.patch\.)(\d+)/i); + return [Number(m[1]), Number(m[2]), Number(m[3]), maintenance ? Number(maintenance[1]) : 0]; + }; + const isAtLeast = (a, b) => { + for (let i = 0; i < a.length; i++) { + if (a[i] !== b[i]) return a[i] > b[i]; + } + return true; + }; + const cosignSection = [ `## Verify Docker Image Signature`, ``, @@ -90,6 +106,22 @@ jobs: ].join('\n'); try { + let makeLatest = "false"; + const newVersion = versionKey(tag); + if (!isPrerelease && newVersion) { + let latestVersion = null; + try { + const latest = await github.rest.repos.getLatestRelease({ + owner: context.repo.owner, + repo: context.repo.repo, + }); + latestVersion = versionKey(latest.data.tag_name); + } catch (error) { + if (error.status !== 404) throw error; + } + makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false"; + } + const response = await github.rest.repos.createRelease({ draft: true, generate_release_notes: true, @@ -108,6 +140,7 @@ jobs: release_id: response.data.id, body: updatedBody, draft: false, + make_latest: makeLatest, }); } catch (error) { diff --git a/CLAUDE.md b/CLAUDE.md index 02a9630b486..758eac7e266 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -52,6 +52,19 @@ Do not put names of customers or customer company names in code, PRs, and issues CI supply-chain safety: Never pipe a remote script into a shell (`curl ... | bash`, `wget ... | sh`); download the artifact to a file, verify its SHA-256 checksum, then install. Pin every external tool to a specific version with a full URL (not `latest` or `stable`). Verify checksums for all downloaded binaries, using the provider's official `.sha256` / `.sha256sum` sidecar when available. These rules apply to every download in CI +Follow these coding conventions for new/updated code (a three-line fix in a legacy file shouldn't trigger huge drive-by refactors): + +- Composition over inheritance +- Never-nester: early returns over deep nesting +- Don't throw; model failures as values (One function (e.g., raise_public) maps error union to existing public exception contracts via exhaustive match + assert_never) +- No mutation; instead of mutable lists and dicts, prefer tuples, NamedTuples, frozen dataclasses, etc. +- Use dependency injection +- Fully typed; no `Any` or coarse types like dict[str, Any]. Every function parameter must be strongly typed +- Use tagged unions + match +- No monster files or god objects + +Follow conventional commits for commit names and PR titles + ## Think Before Coding **Don't assume. Don't hide confusion. Surface tradeoffs** diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index d02afe37569..a0d63f5043c 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -129,7 +129,7 @@ "bash_20241022": null, "bash_20250124": null, "code-execution-2025-08-25": null, - "compact-2026-01-12": null, + "compact-2026-01-12": "compact-2026-01-12", "computer-use-2025-01-24": "computer-use-2025-01-24", "computer-use-2025-11-24": "computer-use-2025-11-24", "context-1m-2025-08-07": "context-1m-2025-08-07", diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index c1afde16250..b6cfc8e7907 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -691,6 +691,7 @@ class Cache: self, embedding_response: Any, model: Optional[str], + prompt_tokens: Optional[int] = None, prompt_tokens_details: Optional[dict] = None, ) -> CachedEmbedding: """ @@ -703,6 +704,7 @@ class Cache: "index": embedding_response.get("index"), "object": embedding_response.get("object"), "model": model, + "prompt_tokens": prompt_tokens, "prompt_tokens_details": prompt_tokens_details, } elif hasattr(embedding_response, "model_dump"): @@ -712,6 +714,7 @@ class Cache: "index": data.get("index"), "object": data.get("object"), "model": model, + "prompt_tokens": prompt_tokens, "prompt_tokens_details": prompt_tokens_details, } else: @@ -721,6 +724,7 @@ class Cache: "index": data.get("index"), "object": data.get("object"), "model": model, + "prompt_tokens": prompt_tokens, "prompt_tokens_details": prompt_tokens_details, } except KeyError as e: @@ -769,6 +773,29 @@ class Cache: per_item[key] = value return per_item if per_item else None + def _get_per_item_prompt_tokens( + self, + result: EmbeddingResponse, + idx_in_result_data: int, + ) -> Optional[int]: + """ + Extract the per-item prompt_tokens from a response for caching. + + Single-item responses store the full usage.prompt_tokens. Multi-item + responses distribute it evenly (with remainder) so that summing all + per-item values on retrieval reconstructs the original total. + """ + if result.usage is None or result.usage.prompt_tokens is None: + return None + + total = result.usage.prompt_tokens + num_items = len(result.data) + if num_items <= 1: + return total + + quotient, remainder = divmod(total, num_items) + return quotient + (1 if idx_in_result_data < remainder else 0) + def add_embedding_response_to_cache( self, result: EmbeddingResponse, @@ -780,7 +807,11 @@ class Cache: kwargs["cache_key"] = preset_cache_key embedding_response = result.data[idx_in_result_data] - # Extract per-item prompt_tokens_details from response usage + # Extract per-item prompt_tokens + details from response usage + prompt_tokens = self._get_per_item_prompt_tokens( + result=result, + idx_in_result_data=idx_in_result_data, + ) prompt_tokens_details = self._get_per_item_prompt_tokens_details( result=result, idx_in_result_data=idx_in_result_data, @@ -791,6 +822,7 @@ class Cache: embedding_dict: CachedEmbedding = self._convert_to_cached_embedding( embedding_response, model_name, + prompt_tokens=prompt_tokens, prompt_tokens_details=prompt_tokens_details, ) diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 3f4e54382c9..48691335b40 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -394,7 +394,7 @@ class LLMCachingHandler: return cr["model"] return None - def _process_async_embedding_cached_response( + def _process_async_embedding_cached_response( # noqa: PLR0915 self, final_embedding_cached_response: Optional[EmbeddingResponse], cached_result: List[Optional[CachedEmbedding]], @@ -456,7 +456,10 @@ class LLMCachingHandler: index=idx, object="embedding", ) - if isinstance(kwargs_input_as_list[idx], str): + cached_prompt_tokens = cr.get("prompt_tokens") + if cached_prompt_tokens is not None: + prompt_tokens += cached_prompt_tokens + elif isinstance(kwargs_input_as_list[idx], str): from litellm.utils import token_counter prompt_tokens += token_counter( diff --git a/litellm/constants.py b/litellm/constants.py index f10cec034f0..57f55e6c177 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1158,6 +1158,7 @@ BEDROCK_CONVERSE_MODELS = [ "openai.gpt-oss-120b-1:0", "anthropic.claude-haiku-4-5-20251001-v1:0", "anthropic.claude-sonnet-4-5-20250929-v1:0", + "anthropic.claude-fable-5", "anthropic.claude-opus-4-8", "anthropic.claude-opus-4-7", "anthropic.claude-opus-4-6-v1:0", diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index c6ae7d9e82b..8899089500d 100644 --- a/litellm/integrations/compression_interception/handler.py +++ b/litellm/integrations/compression_interception/handler.py @@ -72,8 +72,13 @@ class CompressionInterceptionLogger(CustomLogger): compression_params: CompressionInterceptionConfig = {} if "compression_interception_params" in litellm_settings: compression_params = litellm_settings["compression_interception_params"] - elif "compression_interception" in callback_specific_params: - compression_params = callback_specific_params["compression_interception"] + elif "compression_interception" in callback_specific_params and isinstance( + callback_specific_params["compression_interception"], dict + ): + compression_params = cast( + CompressionInterceptionConfig, + callback_specific_params["compression_interception"], + ) return CompressionInterceptionLogger.from_config_yaml(compression_params) async def async_pre_call_deployment_hook( diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 37528e7dcd5..79f9b16bba0 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1339,8 +1339,13 @@ class WebSearchInterceptionLogger(CustomLogger): websearch_params: WebSearchInterceptionConfig = {} if "websearch_interception_params" in litellm_settings: websearch_params = litellm_settings["websearch_interception_params"] - elif "websearch_interception" in callback_specific_params: - websearch_params = callback_specific_params["websearch_interception"] + elif "websearch_interception" in callback_specific_params and isinstance( + callback_specific_params["websearch_interception"], dict + ): + websearch_params = cast( + WebSearchInterceptionConfig, + callback_specific_params["websearch_interception"], + ) # Use classmethod to initialize from config return WebSearchInterceptionLogger.from_config_yaml(websearch_params) diff --git a/litellm/litellm_core_utils/cli_token_utils.py b/litellm/litellm_core_utils/cli_token_utils.py index 3776d276912..eb01359cdc0 100644 --- a/litellm/litellm_core_utils/cli_token_utils.py +++ b/litellm/litellm_core_utils/cli_token_utils.py @@ -37,7 +37,7 @@ def get_litellm_gateway_api_key( """ Get the stored CLI API key for use with LiteLLM SDK. - This function reads the token file created by `litellm-proxy login` + This function reads the token file created by `lite login` and returns the API key for use in Python scripts. Args: diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index b32803b5dfc..f80cb41dc3f 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -34,6 +34,7 @@ _OPTIONAL_KWARGS_KEYS = frozenset( "aws_bedrock_runtime_endpoint", "tpm", "rpm", + "use_xai_oauth", } ) diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 2547fd4d8c6..4e5b53a13d7 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -633,11 +633,6 @@ def convert_to_model_response_object( # noqa: PLR0915 thinking_blocks = choice["message"]["thinking_blocks"] provider_specific_fields["thinking_blocks"] = thinking_blocks - if reasoning_content: - provider_specific_fields["reasoning_content"] = ( - reasoning_content - ) - message = Message( content=content, role=choice["message"]["role"] or "assistant", diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 81a4c8b14b6..b09f2bb130e 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4290,6 +4290,49 @@ def _deduplicate_bedrock_tool_content( return _deduplicate_bedrock_content_blocks(tool_content, "toolResult") +def _rename_duplicate_bedrock_document_names( + contents: List[BedrockMessageBlock], +) -> List[BedrockMessageBlock]: + """ + Rename duplicate document names across all messages in a Bedrock request. + + Document names are derived from a content hash, so the same file appearing + in multiple conversation turns produces identical names and Bedrock rejects + the request with "Messages can not contain duplicate document names". The + first occurrence keeps its original name so prompt-cache prefixes stay + stable; later occurrences get a deterministic positional suffix + (``_2``, ``_3``, ...), bumped further if the suffixed name already + belongs to another document (e.g. an organic name ending in ``_2``). + """ + used_names: Set[str] = set() + for message in contents: + for block in message.get("content") or []: + document = block.get("document") + if isinstance(document, dict) and document.get("name"): + used_names.add(document["name"]) + + name_counts: Dict[str, int] = {} + for message in contents: + for block in message.get("content") or []: + document = block.get("document") + if not isinstance(document, dict): + continue + name = document.get("name") + if not name: + continue + count = name_counts.get(name, 0) + 1 + name_counts[name] = count + if count > 1: + suffix = count + new_name = f"{name}_{suffix}" + while new_name in used_names: + suffix += 1 + new_name = f"{name}_{suffix}" + used_names.add(new_name) + document["name"] = new_name + return contents + + def _sort_bedrock_assistant_content_blocks( blocks: List[BedrockContentBlock], ) -> List[BedrockContentBlock]: @@ -4938,7 +4981,7 @@ class BedrockConverseMessagesProcessor: llm_provider=llm_provider, ) - return contents + return _rename_duplicate_bedrock_document_names(contents) @staticmethod def translate_thinking_blocks_to_reasoning_content_blocks( @@ -5360,7 +5403,7 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 llm_provider=llm_provider, ) - return contents + return _rename_duplicate_bedrock_document_names(contents) def make_valid_bedrock_tool_name(input_tool_name: str) -> str: diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 3f30d5d6807..9ecd0df0cb8 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1455,10 +1455,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _value = self._map_stop_sequences(value) if _value is not None: optional_params["stop_sequences"] = _value - elif param == "temperature": - optional_params["temperature"] = value - elif param == "top_p": - optional_params["top_p"] = value + elif param == "temperature" or param == "top_p": + AnthropicConfig._apply_sampling_param( + optional_params=optional_params, + model=model, + param=param, + value=value, + drop_params=drop_params, + output_key=param, + ) elif param == "response_format" and isinstance(value, dict): if any( substring in model @@ -1975,6 +1980,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): optional_params.pop("is_vertex_request", None) optional_params.pop("client_metadata", None) + # ``top_k`` is a provider-specific kwarg that bypasses + # ``map_openai_params``; gate it here, the single boundary shared by + # the direct Anthropic, Bedrock invoke, Vertex, and Azure paths. + top_k = optional_params.pop("top_k", None) + if top_k is not None: + AnthropicConfig._apply_sampling_param( + optional_params=optional_params, + model=model, + param="top_k", + value=top_k, + drop_params=litellm_params.get("drop_params") is True, + output_key="top_k", + ) + data = { "model": model, "messages": anthropic_messages, diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 3f002d73cbc..5741513903c 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -272,23 +272,68 @@ class AnthropicModelInfo(BaseLLMModelInfo): ) @staticmethod - def _supports_model_capability(model: str, key: str) -> bool: - """Check a boolean capability ``key`` in the model map. + def _supports_sampling_params(model: str) -> bool: + """Claude 4.7+ (Opus 4.7/4.8, Fable 5) removed sampling params: the API + rejects ``top_p``, ``top_k``, and any ``temperature`` other than 1 with + a 400 ("`temperature` is deprecated for this model"). - Strips bedrock/vertex prefixes so a provider-routed Claude still - resolves to the Anthropic model-map entry. - """ - from litellm.utils import _supports_factory + Driven by the ``supports_sampling_params`` flag in the model map; the + name check remains only as a fallback for provider-routed ids whose + map entries predate the flag.""" + flag = AnthropicModelInfo._get_model_capability( + model, "supports_sampling_params" + ) + if flag is not None: + return flag + model_lower = model.lower() + return not any( + v in model_lower + for v in ( + "fable", + "opus-4-7", + "opus_4_7", + "opus-4.7", + "opus_4.7", + "opus-4-8", + "opus_4_8", + "opus-4.8", + "opus_4.8", + ) + ) - try: - if _supports_factory( - model=model, - custom_llm_provider="anthropic", - key=key, - ): - return True - except Exception: - pass + @staticmethod + def _apply_sampling_param( + optional_params: dict, + model: str, + param: str, + value: Any, + drop_params: bool, + output_key: str, + ) -> None: + """Forward ``temperature``/``top_p``/``top_k`` to + ``optional_params[output_key]`` unless the model removed sampling + params, in which case drop the param (with drop_params) or raise a + clean client-side 400.""" + if AnthropicModelInfo._supports_sampling_params(model) or ( + param == "temperature" and value == 1 + ): + optional_params[output_key] = value + elif not (litellm.drop_params or drop_params): + supported_hint = ( + "Only temperature=1 is supported. " if param == "temperature" else "" + ) + raise litellm.utils.UnsupportedParamsError( + message=( + f"{model} does not support {param}={value}. {supported_hint}" + "To drop unsupported params, set `litellm.drop_params = True`." + ), + status_code=400, + ) + + @staticmethod + def _model_map_lookup_candidates(model: str) -> List[str]: + """Model-map keys to try for ``model``, stripping bedrock/vertex + prefixes so a provider-routed Claude still resolves to its entry.""" candidates = [model] for prefix in ( "bedrock/converse/", @@ -307,15 +352,40 @@ class AnthropicModelInfo(BaseLLMModelInfo): candidates.append(f"bedrock/{base}") except Exception: pass + return candidates + + @staticmethod + def _get_model_capability(model: str, key: str) -> Optional[bool]: + """Read boolean capability ``key`` from the model map, or None when + no entry declares it.""" try: - for cand in candidates: - if cand in litellm.model_cost and ( - litellm.model_cost[cand].get(key) is True - ): - return True + for cand in AnthropicModelInfo._model_map_lookup_candidates(model): + value = litellm.model_cost.get(cand, {}).get(key) + if isinstance(value, bool): + return value except Exception: pass - return False + return None + + @staticmethod + def _supports_model_capability(model: str, key: str) -> bool: + """Check a boolean capability ``key`` in the model map. + + Strips bedrock/vertex prefixes so a provider-routed Claude still + resolves to the Anthropic model-map entry. + """ + from litellm.utils import _supports_factory + + try: + if _supports_factory( + model=model, + custom_llm_provider="anthropic", + key=key, + ): + return True + except Exception: + pass + return AnthropicModelInfo._get_model_capability(model, key) is True @staticmethod def _is_adaptive_thinking_model(model: str) -> bool: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index bacb9f8ddf6..8c20f4c430e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -1,5 +1,6 @@ # What is this? ## Translates OpenAI call to Anthropic `/v1/messages` format +import copy import json import traceback from collections import deque @@ -29,6 +30,98 @@ if TYPE_CHECKING: from litellm.types.utils import ModelResponseStream +class _CombinedChunkSplitter: + """ + Splits a streaming chunk that carries BOTH response content and a + ``finish_reason`` into two chunks: a content-only chunk followed by a + finish-only chunk. + + ``AnthropicStreamWrapper`` (via ``translate_streaming_openai_response_to_anthropic``) + assumes content and ``finish_reason`` never arrive in the same chunk — true for + real provider streams, but false for fake-streamed providers (e.g. Vertex AI + Gemma ``:predict``) where ``MockResponseIterator`` collapses the entire response + into a single chunk. Without this split the assumption causes all content to be + silently dropped (only the ``message_delta`` stop event is emitted). + + Supports both sync and async iteration, since ``AnthropicStreamWrapper`` exposes + both ``__next__`` and ``__anext__``. An instance is single-mode: callers must + iterate it either synchronously or asynchronously, never both — the two modes + hold independent iterator references on the upstream stream and mixing them + would advance them out of sync. + """ + + def __init__(self, completion_stream: Any): + self._stream = completion_stream + self._sync_iter: Optional[Iterator[Any]] = None + self._async_iter: Optional[AsyncIterator[Any]] = None + self._buffer: deque = deque() + + @staticmethod + def _is_combined(chunk: Any) -> bool: + """True if ``chunk`` carries response content AND a finish_reason.""" + choices = getattr(chunk, "choices", None) + if not choices: + return False + choice = choices[0] + if getattr(choice, "finish_reason", None) is None: + return False + delta = getattr(choice, "delta", None) + if delta is None: + return False + return bool( + getattr(delta, "content", None) + or getattr(delta, "tool_calls", None) + or getattr(delta, "reasoning_content", None) + or getattr(delta, "thinking_blocks", None) + ) + + @staticmethod + def _split(chunk: Any) -> List[Any]: + """Return ``[chunk]``, or ``[content_chunk, finish_chunk]`` if combined.""" + if not _CombinedChunkSplitter._is_combined(chunk): + return [chunk] + + # Content chunk: keep the delta payload, clear the finish_reason. + content_chunk = copy.deepcopy(chunk) + content_chunk.choices[0].finish_reason = None + + # Finish chunk: keep finish_reason (and usage), clear the delta payload. + finish_chunk = copy.deepcopy(chunk) + finish_delta = finish_chunk.choices[0].delta + finish_delta.content = None + if hasattr(finish_delta, "tool_calls"): + finish_delta.tool_calls = None + if hasattr(finish_delta, "reasoning_content"): + finish_delta.reasoning_content = None + if hasattr(finish_delta, "thinking_blocks"): + finish_delta.thinking_blocks = None + return [content_chunk, finish_chunk] + + def __iter__(self) -> "Iterator[Any]": + return self + + def __next__(self) -> Any: + if self._buffer: + return self._buffer.popleft() + if self._sync_iter is None: + self._sync_iter = iter(self._stream) + chunk = next(self._sync_iter) # propagates StopIteration when exhausted + self._buffer.extend(self._split(chunk)) + return self._buffer.popleft() + + def __aiter__(self) -> "AsyncIterator[Any]": + return self + + async def __anext__(self) -> Any: + if self._buffer: + return self._buffer.popleft() + if self._async_iter is None: + self._async_iter = self._stream.__aiter__() + chunk = await self._async_iter.__anext__() # propagates StopAsyncIteration + self._buffer.extend(self._split(chunk)) + return self._buffer.popleft() + + class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): """ - first chunk return 'message_start' @@ -62,7 +155,10 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): compaction_block: Optional[CompactionBlock] = None, iterations_usage: Optional[List[UsageIteration]] = None, ): - super().__init__(completion_stream) + # Wrap the upstream stream so chunks that carry both content and a + # finish_reason (fake-streamed providers) are split into two — see + # _CombinedChunkSplitter. + super().__init__(_CombinedChunkSplitter(completion_stream)) self.model = model # Mapping of truncated tool names to original names (for OpenAI's 64-char limit) self.tool_name_mapping = tool_name_mapping or {} diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py index 94c5200be64..5f1362e259f 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py @@ -155,10 +155,24 @@ class AnthropicResponsesStreamWrapper: event.get("delta", "") if isinstance(event, dict) else "" ) block_idx = ( - self._item_id_to_block_index.get(item_id, self._current_block_index) + self._item_id_to_block_index.get(item_id, -1) if item_id else self._current_block_index ) + if block_idx < 0: + # Some providers (e.g. LMStudio) skip response.output_item.added, + # so no text block is open yet; synthesize content_block_start + # instead of emitting a delta with index -1 + block_idx = self._next_block_index() + if item_id: + self._item_id_to_block_index[item_id] = block_idx + self._chunk_queue.append( + { + "type": "content_block_start", + "index": block_idx, + "content_block": {"type": "text", "text": ""}, + } + ) self._chunk_queue.append( { "type": "content_block_delta", diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index cf1fd6f786e..bf1bfd06537 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -50,6 +50,11 @@ def convert_model_response_to_streaming( model=model_response.model, choices=streaming_choices, ) + # Carry usage onto the streaming chunk so fake-streamed responses + # (e.g. Vertex AI Gemma :predict) still report token counts. + usage = getattr(model_response, "usage", None) + if usage is not None: + setattr(processed_chunk, "usage", usage) return processed_chunk except Exception as e: raise ValueError( diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index ea0326dffd1..b5e5e4de6fc 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -920,10 +920,15 @@ class AmazonConverseConfig(BaseConfig): continue value = [value] optional_params["stopSequences"] = value - if param == "temperature": - optional_params["temperature"] = value - if param == "top_p": - optional_params["topP"] = value + if param == "temperature" or param == "top_p": + AnthropicConfig._apply_sampling_param( + optional_params=optional_params, + model=model, + param=param, + value=value, + drop_params=drop_params, + output_key="topP" if param == "top_p" else param, + ) if param == "tools" and isinstance(value, list): self._apply_tool_call_transformation( tools=cast(List[OpenAIChatCompletionToolParam], value), @@ -1221,7 +1226,9 @@ class AmazonConverseConfig(BaseConfig): inference_params["topK"] = inference_params.pop("top_k") return InferenceConfig(**inference_params) - def _handle_top_k_value(self, model: str, inference_params: dict) -> dict: + def _handle_top_k_value( + self, model: str, inference_params: dict, drop_params: bool = False + ) -> dict: base_model = BedrockModelInfo.get_base_model(model) val_top_k = None @@ -1230,16 +1237,25 @@ class AmazonConverseConfig(BaseConfig): elif "top_k" in inference_params: val_top_k = inference_params.pop("top_k") - if val_top_k: + if val_top_k is not None: if base_model.startswith("anthropic"): - return {"top_k": val_top_k} + top_k_params: dict = {} + AnthropicConfig._apply_sampling_param( + optional_params=top_k_params, + model=model, + param="top_k", + value=val_top_k, + drop_params=drop_params, + output_key="top_k", + ) + return top_k_params if base_model.startswith("amazon.nova"): return {"inferenceConfig": {"topK": val_top_k}} return {} def _prepare_request_params( - self, optional_params: dict, model: str + self, optional_params: dict, model: str, drop_params: bool = False ) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]: """Prepare and separate request parameters.""" # Consume the internal ``_output_config_normalized`` marker set by @@ -1338,7 +1354,7 @@ class AmazonConverseConfig(BaseConfig): # Only set the topK value in for models that support it additional_request_params.update( - self._handle_top_k_value(model, inference_params) + self._handle_top_k_value(model, inference_params, drop_params) ) # Filter out internal/MCP-related parameters that shouldn't be sent to the API @@ -1572,6 +1588,7 @@ class AmazonConverseConfig(BaseConfig): optional_params: dict, messages: Optional[List[AllMessageValues]] = None, headers: Optional[dict] = None, + drop_params: bool = False, ) -> CommonRequestObject: ## VALIDATE REQUEST """ @@ -1618,7 +1635,7 @@ class AmazonConverseConfig(BaseConfig): additional_request_params, request_metadata, output_config, - ) = self._prepare_request_params(optional_params, model) + ) = self._prepare_request_params(optional_params, model, drop_params) original_tools = inference_params.pop("tools", []) @@ -1701,6 +1718,7 @@ class AmazonConverseConfig(BaseConfig): optional_params=optional_params, messages=messages, headers=headers, + drop_params=litellm_params.get("drop_params") is True, ) bedrock_messages = ( @@ -1758,6 +1776,7 @@ class AmazonConverseConfig(BaseConfig): optional_params=optional_params, messages=messages, headers=headers, + drop_params=litellm_params.get("drop_params") is True, ) ## TRANSFORMATION ## diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py index c967fd334bc..bdef3349e00 100644 --- a/litellm/llms/bedrock/count_tokens/transformation.py +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -11,6 +11,11 @@ from typing import Any, Dict, List, Optional from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.bedrock.common_utils import get_bedrock_base_model +# Placeholder satisfying the Anthropic InvokeModel schema's required +# max_tokens field; CountTokens only counts input, so it has no effect +# on any generation. +DEFAULT_ANTHROPIC_INVOKE_MODEL_MAX_TOKENS = 1024 + class BedrockCountTokensConfig(BaseAWSLLM): """ @@ -32,8 +37,20 @@ class BedrockCountTokensConfig(BaseAWSLLM): Returns: 'converse' or 'invokeModel' """ - # If the request has messages in the expected Anthropic format, use converse - if "messages" in request_data and isinstance(request_data["messages"], list): + messages = request_data.get("messages") + if isinstance(messages, list): + # Anthropic content blocks carry a "type" key ({"type": "text", ...}); + # Converse blocks don't ({"text": ...}, {"toolUse": ...}). Converse + # rejects Anthropic-shape blocks, so route those to invokeModel, + # which forwards the body verbatim. + for message in messages: + if not isinstance(message, dict): + continue + content = message.get("content") + if isinstance(content, list) and any( + isinstance(block, dict) and "type" in block for block in content + ): + return "invokeModel" return "converse" # For raw text or other formats, use invokeModel @@ -68,7 +85,7 @@ class BedrockCountTokensConfig(BaseAWSLLM): { "input": { "invokeModel": { - "body": "{...raw model input...}" + "body": "" } } } @@ -168,13 +185,24 @@ class BedrockCountTokensConfig(BaseAWSLLM): self, request_data: Dict[str, Any] ) -> Dict[str, Any]: """Transform to InvokeModel input format.""" + import base64 import json # For InvokeModel, we need to provide the raw body that would be sent to the model # Remove the 'model' field from the body as it's not part of the model input body_data = {k: v for k, v in request_data.items() if k != "model"} - return {"input": {"invokeModel": {"body": json.dumps(body_data)}}} + if "messages" in body_data: + # Bedrock validates the body against the model's InvokeModel schema; + # Anthropic Messages bodies require these fields. + body_data.setdefault("anthropic_version", "bedrock-2023-05-31") + body_data.setdefault( + "max_tokens", DEFAULT_ANTHROPIC_INVOKE_MODEL_MAX_TOKENS + ) + + # The CountTokens API expects invokeModel.body as a base64-encoded blob + encoded_body = base64.b64encode(json.dumps(body_data).encode()).decode() + return {"input": {"invokeModel": {"body": encoded_body}}} def get_bedrock_count_tokens_endpoint( self, diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index df219091074..dfa108833ac 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -1,8 +1,10 @@ """ Amazon Bedrock Mantle - Responses API backend. -gpt-5.5 / gpt-5.4 on Mantle are exposed ONLY on the `/openai/v1/responses` -path (not the standard `/v1/responses`). Payloads and SSE follow the OpenAI +Mantle serves Responses on two upstream paths: gpt frontier models (gpt-5.5 / +gpt-5.4) on `/openai/v1/responses`, and everything else that supports Responses +(e.g. gpt-oss) on the standard `/v1/responses`. The gate picks the path per +model and injects it via `use_openai_path`. Payloads and SSE follow the OpenAI Responses spec, so this config inherits OpenAIResponsesAPIConfig and overrides only the endpoint URL and authentication. @@ -48,9 +50,14 @@ _MANTLE_HOST_RE = re.compile( class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): - def __init__(self, aws_signer: Optional[BaseAWSLLM] = None): + def __init__( + self, + aws_signer: Optional[BaseAWSLLM] = None, + use_openai_path: bool = True, + ): super().__init__() self._aws_signer = aws_signer or BaseAWSLLM() + self.use_openai_path = use_openai_path @property def custom_llm_provider(self) -> LlmProviders: @@ -94,7 +101,8 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): # single resolved region so aws_region_name wins; preserve custom proxy hosts. if _MANTLE_HOST_RE.match(base): base = f"https://bedrock-mantle.{region}.api.aws" - return f"{base}/openai/v1/responses" + path = "/openai/v1/responses" if self.use_openai_path else "/v1/responses" + return f"{base}{path}" def validate_environment( self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] diff --git a/litellm/llms/databricks/streaming_utils.py b/litellm/llms/databricks/streaming_utils.py index eebe3182881..7a7330227d6 100644 --- a/litellm/llms/databricks/streaming_utils.py +++ b/litellm/llms/databricks/streaming_utils.py @@ -25,6 +25,28 @@ class ModelResponseIterator: finish_reason = "" usage: Optional[ChatCompletionUsageBlock] = None + # Usage-only final chunk (OpenAI ``stream_options.include_usage``) + # arrives with an empty ``choices`` list — return usage without + # indexing ``choices[0]``. + if len(processed_chunk.choices) == 0: + final_usage = getattr(processed_chunk, "usage", None) + return GenericStreamingChunk( + text="", + tool_use=None, + is_finished=False, + finish_reason="", + usage=( + ChatCompletionUsageBlock( + prompt_tokens=final_usage.prompt_tokens or 0, + completion_tokens=final_usage.completion_tokens or 0, + total_tokens=final_usage.total_tokens or 0, + ) + if final_usage is not None + else None + ), + index=0, + ) + if processed_chunk.choices[0].delta.content is not None: # type: ignore text = processed_chunk.choices[0].delta.content # type: ignore diff --git a/litellm/llms/litellm_proxy/skills/README.md b/litellm/llms/litellm_proxy/skills/README.md index 1dfeff1a42c..a896aa1166e 100644 --- a/litellm/llms/litellm_proxy/skills/README.md +++ b/litellm/llms/litellm_proxy/skills/README.md @@ -18,7 +18,7 @@ flowchart TB F[Request with container.skills] --> G[SkillsInjectionHook] G --> H{skill_id prefix?} - H -->|"litellm:skill_abc"| I[Fetch from LiteLLM DB] + H -->|"litellm_skill_abc"| I[Fetch from LiteLLM DB] H -->|"skill_xyz" no prefix| J[Pass to Anthropic as native skill] I --> K{Model provider?} @@ -57,7 +57,7 @@ sequenceDiagram Note over LiteLLM,PreHook: PRE-CALL HOOK LiteLLM->>PreHook: Intercept request - PreHook->>PreHook: Fetch skill from DB (litellm:skill_id) + PreHook->>PreHook: Fetch skill from DB (litellm_skill_id) PreHook->>PreHook: Extract SKILL.md from ZIP PreHook->>PreHook: Inject SKILL.md into system prompt PreHook->>PreHook: Add litellm_code_execution tool @@ -105,7 +105,7 @@ response = await litellm.acompletion( model="gpt-4o-mini", messages=[{"role": "user", "content": "Create a bouncing ball GIF"}], container={ - "skills": [{"type": "custom", "skill_id": "litellm:skill_abc123"}] + "skills": [{"type": "custom", "skill_id": "litellm_skill_abc123"}] }, ) @@ -261,7 +261,7 @@ response = litellm.completion( messages=[{"role": "user", "content": "Analyze this data..."}], container={ "skills": [ - {"type": "custom", "skill_id": "litellm:skill_abc123"} # litellm: prefix + {"type": "custom", "skill_id": "litellm_skill_abc123"} # litellm_skill_ prefix ] } ) @@ -277,7 +277,7 @@ response = litellm.completion( "messages": [{"role": "user", "content": "Help me analyze data"}], "container": { "skills": [ - {"type": "custom", "skill_id": "litellm:skill_abc123"} + {"type": "custom", "skill_id": "litellm_skill_abc123"} ] } } @@ -287,7 +287,7 @@ response = litellm.completion( The hook (`litellm/proxy/hooks/litellm_skills/main.py`) intercepts the request: -1. **Detects `litellm:` prefix** → Fetches skill from database +1. **Detects `litellm_skill_` prefix** → Fetches skill from database 2. **Checks model provider** → Bedrock is not Anthropic 3. **Extracts SKILL.md** from stored ZIP file 4. **Converts skill to tool** + **Injects content into system prompt** @@ -361,8 +361,8 @@ model LiteLLM_SkillsTable { | Create skill on Anthropic | `anthropic` | N/A | Forward to Anthropic API | | Create skill in LiteLLM DB | `litellm_proxy` | N/A | Store in database | | Use Anthropic native skill | N/A | `skill_xyz` | Pass to Anthropic container.skills | -| Use LiteLLM skill on Anthropic | N/A | `litellm:skill_abc` | Convert to tools | -| Use LiteLLM skill on Bedrock/OpenAI | N/A | `litellm:skill_abc` | Convert to tools + inject SKILL.md | +| Use LiteLLM skill on Anthropic | N/A | `litellm_skill_abc` | Convert to tools | +| Use LiteLLM skill on Bedrock/OpenAI | N/A | `litellm_skill_abc` | Convert to tools + inject SKILL.md | ## Testing diff --git a/litellm/llms/litellm_proxy/skills/constants.py b/litellm/llms/litellm_proxy/skills/constants.py index a8c2697fcee..0c60a60842a 100644 --- a/litellm/llms/litellm_proxy/skills/constants.py +++ b/litellm/llms/litellm_proxy/skills/constants.py @@ -4,6 +4,10 @@ Constants for LiteLLM Skills Centralized constants for skills processing, code execution, and sandbox configuration. """ +LITELLM_SKILL_ID_PREFIX: str = "litellm_skill_" +"""Prefix for DB-backed skill IDs. The model-facing tool name is the skill ID +with hyphens/spaces replaced by underscores, which leaves this prefix intact.""" + # Code execution loop settings DEFAULT_MAX_ITERATIONS: int = 10 """Maximum number of iterations for the automatic code execution loop.""" diff --git a/litellm/llms/litellm_proxy/skills/handler.py b/litellm/llms/litellm_proxy/skills/handler.py index 7b259c1ed66..9138b9a712f 100644 --- a/litellm/llms/litellm_proxy/skills/handler.py +++ b/litellm/llms/litellm_proxy/skills/handler.py @@ -10,6 +10,7 @@ from typing import Any, Dict, List, Optional from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache +from litellm.llms.litellm_proxy.skills.constants import LITELLM_SKILL_ID_PREFIX from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest, UserAPIKeyAuth from litellm.proxy.common_utils.resource_ownership import ( get_primary_resource_owner_scope, @@ -68,7 +69,7 @@ class LiteLLMSkillsHandler: ) -> LiteLLM_SkillsTable: prisma_client = await LiteLLMSkillsHandler._get_prisma_client() - skill_id = f"litellm_skill_{uuid.uuid4()}" + skill_id = f"{LITELLM_SKILL_ID_PREFIX}{uuid.uuid4()}" owner = get_primary_resource_owner_scope(user_api_key_dict) or user_id if owner is None: # Identity-less callers (no user_id / team_id / org_id / diff --git a/litellm/llms/openai/completion/handler.py b/litellm/llms/openai/completion/handler.py index 1641615126e..63d39151254 100644 --- a/litellm/llms/openai/completion/handler.py +++ b/litellm/llms/openai/completion/handler.py @@ -49,6 +49,8 @@ class OpenAITextCompletion(BaseLLM): headers: Optional[dict] = None, ): try: + if headers: + optional_params = {**optional_params, "extra_headers": headers} if headers is None: headers = self.validate_environment(api_key=api_key) if model is None or messages is None: diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index e9f08f403f9..103801a1e8d 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -19,7 +19,7 @@ from litellm.types.llms.vertex_ai import ( VertexAICachedContentResponseObject, ) -from ..common_utils import VertexAIError +from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase from .transformation import ( separate_cached_messages, @@ -69,17 +69,13 @@ class ContextCachingEndpoints(VertexBase): elif custom_llm_provider == "vertex_ai": auth_header = vertex_auth_header endpoint = "cachedContents" - if vertex_location == "global": - url = f"https://aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" - else: - url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" + base_url = get_vertex_base_url(vertex_location) + url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" else: auth_header = vertex_auth_header endpoint = "cachedContents" - if vertex_location == "global": - url = f"https://aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" - else: - url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" + base_url = get_vertex_base_url(vertex_location) + url = f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" return self._check_custom_proxy( api_base=api_base, diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index c06928516ef..8019bb67991 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -5,6 +5,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.constants import XAI_API_BASE +from litellm.exceptions import AuthenticationError from litellm.litellm_core_utils.prompt_templates.common_utils import ( filter_value_from_dict, strip_name_from_messages, @@ -39,6 +40,72 @@ class XAIChatConfig(OpenAIGPTConfig): dynamic_api_key = XAIModelInfo.get_api_key(api_key) return api_base, dynamic_api_key + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + dynamic_api_key = XAIModelInfo.get_api_key(api_key) + if should_use_xai_oauth(litellm_params) and not dynamic_api_key: + try: + headers["Authorization"] = ( + f"Bearer {XAIOAuthAuthenticator().get_access_token()}" + ) + except XAIOAuthError as exc: + raise AuthenticationError( + model=model, + llm_provider=self.custom_llm_provider or "xai", + message=str(exc), + ) from exc + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + return headers + + return super().validate_environment( + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=dynamic_api_key, + api_base=api_base, + ) + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + dynamic_api_key = XAIModelInfo.get_api_key(api_key) + if should_use_xai_oauth(litellm_params) and not dynamic_api_key: + api_base = XAIOAuthAuthenticator().get_api_base() + + return super().get_complete_url( + api_base=api_base, + api_key=dynamic_api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + stream=stream, + ) + def get_supported_openai_params(self, model: str) -> list: base_openai_params = [ "logit_bias", diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py new file mode 100644 index 00000000000..30c717b7ca0 --- /dev/null +++ b/litellm/llms/xai/oauth.py @@ -0,0 +1,421 @@ +import base64 +import hashlib +import json +import os +import secrets +import sys +import threading +import time +import uuid +import webbrowser +from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any, Dict, Optional, Tuple, Union +from urllib.parse import parse_qs, urlencode, urlparse + +import httpx + +from litellm._logging import verbose_logger +from litellm.constants import XAI_API_BASE +from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client +from litellm.secret_managers.main import get_secret_str + +XAI_OAUTH_ISSUER = "https://auth.x.ai" +XAI_OAUTH_DISCOVERY_URL = f"{XAI_OAUTH_ISSUER}/.well-known/openid-configuration" +XAI_OAUTH_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828" +XAI_OAUTH_SCOPE = "openid profile email offline_access grok-cli:access api:access" +XAI_OAUTH_REDIRECT_HOST = "127.0.0.1" +XAI_OAUTH_REDIRECT_PORT = 56121 +XAI_OAUTH_REDIRECT_PATH = "/callback" +XAI_OAUTH_EXPIRY_SKEW_SECONDS = 120 +XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS = 180 +_XAI_OAUTH_REFRESH_LOCK = threading.Lock() + + +class XAIOAuthError(Exception): + pass + + +class XAIOAuthLoginRequiredError(XAIOAuthError): + pass + + +class _CallbackHandler(BaseHTTPRequestHandler): + server: "_CallbackServer" + + def do_GET(self) -> None: + parsed = urlparse(self.path) + if parsed.path != XAI_OAUTH_REDIRECT_PATH: + self.send_response(404) + self.end_headers() + return + + params = parse_qs(parsed.query) + result = { + "code": params.get("code", [None])[0], + "state": params.get("state", [None])[0], + "error": params.get("error", [None])[0], + "error_description": params.get("error_description", [None])[0], + } + self.server.callback_result = result + + if result["state"] != self.server.expected_state: + self.send_response(400) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.end_headers() + self.wfile.write( + b"

xAI authorization state mismatch.

" + ) + return + + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.end_headers() + body = ( + b"

xAI authorization failed.

You can close this tab." + if result["error"] + else b"

xAI authorization received.

You can close this tab." + ) + self.wfile.write(body) + + def log_message(self, format: str, *args: Any) -> None: + return + + +class _CallbackServer(HTTPServer): + expected_state: str + callback_result: Optional[Dict[str, Optional[str]]] + + +class XAIOAuthAuthenticator: + def __init__( + self, http_client: Optional[Union[httpx.Client, HTTPHandler]] = None + ) -> None: + self.token_dir = get_secret_str("XAI_OAUTH_TOKEN_DIR") or os.path.expanduser( + "~/.config/litellm/xai_oauth" + ) + self.auth_file = os.path.join( + self.token_dir, get_secret_str("XAI_OAUTH_AUTH_FILE") or "auth.json" + ) + self.http_client = http_client + + def get_api_base(self) -> str: + return ( + get_secret_str("XAI_OAUTH_API_BASE") + or get_secret_str("XAI_API_BASE") + or XAI_API_BASE + ) + + def get_access_token(self) -> str: + auth_data = self._read_auth_file() + if not auth_data: + raise XAIOAuthLoginRequiredError( + "xAI OAuth login required. Run `litellm xai-oauth login`." + ) + + access_token = auth_data.get("access_token") + if access_token and not self._is_expired(auth_data): + return access_token + + refresh_token = auth_data.get("refresh_token") + if not refresh_token: + raise XAIOAuthLoginRequiredError( + "xAI OAuth refresh token missing. Run `litellm xai-oauth login`." + ) + + with _XAI_OAUTH_REFRESH_LOCK: + locked_auth_data = self._read_auth_file() or auth_data + access_token = locked_auth_data.get("access_token") + if access_token and not self._is_expired(locked_auth_data): + return access_token + + refreshed = self._refresh_tokens(locked_auth_data) + return refreshed["access_token"] + + def login(self, force: bool = False, no_browser: bool = False) -> Dict[str, Any]: + existing = self._read_auth_file() + if existing and not force and existing.get("access_token"): + if not self._is_expired(existing): + return existing + if existing.get("refresh_token"): + try: + return self._refresh_tokens(existing) + except XAIOAuthError: + pass + + discovery = self._discover() + verifier, challenge = self._pkce_pair() + state = uuid.uuid4().hex + nonce = uuid.uuid4().hex + server, redirect_uri = self._start_callback_server(state) + authorize_url = self._build_authorize_url( + authorization_endpoint=discovery["authorization_endpoint"], + redirect_uri=redirect_uri, + challenge=challenge, + state=state, + nonce=nonce, + ) + + if no_browser or not webbrowser.open(authorize_url): + sys.stdout.write( + f"Open this URL to authenticate with xAI:\n{authorize_url}\n" + ) + sys.stdout.flush() + + result = self._wait_for_callback(server) + if result.get("state") != state: + raise XAIOAuthError("xAI OAuth state mismatch") + if result.get("error"): + description = result.get("error_description") or result["error"] + raise XAIOAuthError(f"xAI authorization failed: {description}") + code = result.get("code") + if not code: + raise XAIOAuthError("xAI authorization failed: no code returned") + + token_payload = self._exchange_token( + discovery["token_endpoint"], + { + "grant_type": "authorization_code", + "code": code, + "redirect_uri": redirect_uri, + "client_id": XAI_OAUTH_CLIENT_ID, + "code_verifier": verifier, + }, + ) + auth_data = self._build_auth_record(token_payload, discovery["token_endpoint"]) + self._write_auth_file(auth_data) + return auth_data + + def _client(self) -> Union[httpx.Client, HTTPHandler]: + return self.http_client or _get_httpx_client() + + def _ensure_token_dir(self) -> None: + os.makedirs(self.token_dir, mode=0o700, exist_ok=True) + try: + os.chmod(self.token_dir, 0o700) + except OSError: + verbose_logger.debug("Could not chmod xAI OAuth token directory") + + def _read_auth_file(self) -> Optional[Dict[str, Any]]: + try: + with open(self.auth_file, "r") as f: + data = json.load(f) + return data if isinstance(data, dict) else None + except (IOError, json.JSONDecodeError): + return None + + def _write_auth_file(self, data: Dict[str, Any]) -> None: + self._ensure_token_dir() + tmp_file = os.path.join( + self.token_dir, + f".{os.path.basename(self.auth_file)}.{uuid.uuid4().hex}.tmp", + ) + flags = os.O_WRONLY | os.O_CREAT | os.O_EXCL + if hasattr(os, "O_NOFOLLOW"): + flags |= os.O_NOFOLLOW + fd = os.open(tmp_file, flags, 0o600) + try: + with os.fdopen(fd, "w") as f: + json.dump(data, f) + f.flush() + os.fsync(f.fileno()) + os.replace(tmp_file, self.auth_file) + try: + os.chmod(self.auth_file, 0o600) + except OSError: + verbose_logger.debug("Could not chmod xAI OAuth auth file") + except Exception: + try: + os.close(fd) + except OSError: + pass + try: + os.unlink(tmp_file) + except OSError: + pass + raise + + def _is_expired(self, auth_data: Dict[str, Any]) -> bool: + expires_at = auth_data.get("expires_at") + if expires_at is None: + return True + try: + return time.time() >= float(expires_at) - XAI_OAUTH_EXPIRY_SKEW_SECONDS + except (TypeError, ValueError): + return True + + def _discover(self) -> Dict[str, str]: + try: + response = self._client().get( + XAI_OAUTH_DISCOVERY_URL, headers={"Accept": "application/json"} + ) + response.raise_for_status() + except httpx.HTTPStatusError as exc: + raise XAIOAuthError( + f"xAI OAuth discovery request failed: {exc.response.status_code} {exc.response.text}" + ) from exc + try: + data = response.json() + except ValueError as exc: + raise XAIOAuthError( + "xAI OAuth discovery response was not valid JSON" + ) from exc + authorization_endpoint = data.get("authorization_endpoint") + token_endpoint = data.get("token_endpoint") + if not authorization_endpoint or not token_endpoint: + raise XAIOAuthError("xAI OAuth discovery missing endpoints") + return { + "authorization_endpoint": self._validate_xai_endpoint( + authorization_endpoint + ), + "token_endpoint": self._validate_xai_endpoint(token_endpoint), + } + + def _validate_xai_endpoint(self, url: str) -> str: + parsed = urlparse(url) + host = (parsed.hostname or "").lower() + if parsed.scheme != "https" or (host != "x.ai" and not host.endswith(".x.ai")): + raise XAIOAuthError( + f"xAI OAuth discovery returned unexpected endpoint: {url}" + ) + return url + + def _pkce_pair(self) -> Tuple[str, str]: + verifier = ( + base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode() + ) + challenge = ( + base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()) + .rstrip(b"=") + .decode() + ) + return verifier, challenge + + def _start_callback_server(self, state: str) -> Tuple[_CallbackServer, str]: + last_error: Optional[OSError] = None + for port in (XAI_OAUTH_REDIRECT_PORT, 0): + try: + server = _CallbackServer( + (XAI_OAUTH_REDIRECT_HOST, port), _CallbackHandler + ) + server.expected_state = state + server.callback_result = None + actual_port = server.server_address[1] + redirect_uri = f"http://{XAI_OAUTH_REDIRECT_HOST}:{actual_port}{XAI_OAUTH_REDIRECT_PATH}" + return server, redirect_uri + except OSError as exc: + last_error = exc + raise XAIOAuthError(f"Could not start xAI OAuth callback server: {last_error}") + + def _build_authorize_url( + self, + authorization_endpoint: str, + redirect_uri: str, + challenge: str, + state: str, + nonce: str, + ) -> str: + params = { + "response_type": "code", + "client_id": XAI_OAUTH_CLIENT_ID, + "redirect_uri": redirect_uri, + "scope": XAI_OAUTH_SCOPE, + "code_challenge": challenge, + "code_challenge_method": "S256", + "state": state, + "nonce": nonce, + } + return f"{authorization_endpoint}?{urlencode(params)}" + + def _wait_for_callback(self, server: _CallbackServer) -> Dict[str, Optional[str]]: + server.timeout = 1 + deadline = time.time() + XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS + try: + while time.time() < deadline: + server.handle_request() + if server.callback_result is not None: + return server.callback_result + finally: + server.server_close() + raise XAIOAuthError("Timed out waiting for xAI OAuth callback") + + def _exchange_token( + self, token_endpoint: str, data: Dict[str, str] + ) -> Dict[str, Any]: + try: + response = self._client().post( + token_endpoint, + headers={ + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + }, + data=data, + ) + response.raise_for_status() + except httpx.HTTPStatusError as exc: + raise XAIOAuthError( + f"xAI OAuth token request failed: {exc.response.status_code} {exc.response.text}" + ) from exc + try: + body = response.json() + except ValueError as exc: + raise XAIOAuthError("xAI OAuth token response was not valid JSON") from exc + if not isinstance(body, dict): + raise XAIOAuthError("xAI OAuth token response was not an object") + return body + + def _build_auth_record( + self, + token_payload: Dict[str, Any], + token_endpoint: str, + fallback_refresh_token: Optional[str] = None, + ) -> Dict[str, Any]: + access_token = token_payload.get("access_token") + refresh_token = token_payload.get("refresh_token") or fallback_refresh_token + if not access_token: + raise XAIOAuthError("xAI OAuth token response missing access_token") + if not refresh_token: + raise XAIOAuthError("xAI OAuth token response missing refresh_token") + expires_in = token_payload.get("expires_in") or 3600 + try: + expires_at = int(time.time() + int(expires_in)) + except (TypeError, ValueError): + expires_at = int(time.time() + 3600) + return { + "access_token": access_token, + "refresh_token": refresh_token, + "id_token": token_payload.get("id_token"), + "token_type": token_payload.get("token_type") or "Bearer", + "token_endpoint": token_endpoint, + "expires_at": expires_at, + } + + def _refresh_tokens(self, auth_data: Dict[str, Any]) -> Dict[str, Any]: + token_endpoint = auth_data.get("token_endpoint") + if not token_endpoint: + token_endpoint = self._discover()["token_endpoint"] + token_endpoint = self._validate_xai_endpoint(token_endpoint) + refresh_token = auth_data.get("refresh_token") + if not refresh_token: + raise XAIOAuthLoginRequiredError( + "xAI OAuth refresh token missing. Run `litellm xai-oauth login`." + ) + + token_payload = self._exchange_token( + token_endpoint, + { + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "client_id": XAI_OAUTH_CLIENT_ID, + }, + ) + refreshed = self._build_auth_record( + token_payload, + token_endpoint, + fallback_refresh_token=refresh_token, + ) + self._write_auth_file(refreshed) + return refreshed + + +def should_use_xai_oauth(litellm_params: Optional[Dict[str, Any]]) -> bool: + return bool((litellm_params or {}).get("use_xai_oauth")) diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py index 55805ddaede..f81e860a8ce 100644 --- a/litellm/llms/xai/responses/transformation.py +++ b/litellm/llms/xai/responses/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import litellm from litellm._logging import verbose_logger from litellm.constants import XAI_API_BASE +from litellm.exceptions import AuthenticationError from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.llms.xai.common_utils import XAIModelInfo from litellm.secret_managers.main import get_secret_str @@ -220,10 +221,27 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): litellm_params.api_key, legacy_generic_before_env=True ) + if not api_key: + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + if should_use_xai_oauth(litellm_params.model_dump()): + try: + api_key = XAIOAuthAuthenticator().get_access_token() + except XAIOAuthError as exc: + raise AuthenticationError( + model=model, + llm_provider=self.custom_llm_provider.value, + message=str(exc), + ) from exc + if not api_key: raise ValueError( "XAI API key is required. Set api_key, litellm.xai_key, " - "litellm.api_key, or XAI_API_KEY." + "litellm.api_key, XAI_API_KEY, or use_xai_oauth=True." ) headers.update( @@ -244,12 +262,20 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): Returns: str: The full URL for the XAI /responses endpoint """ - api_base = ( - api_base - or litellm.api_base - or get_secret_str("XAI_API_BASE") - or XAI_API_BASE + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + api_key = XAIModelInfo.get_api_key( + litellm_params.get("api_key"), legacy_generic_before_env=True ) + if should_use_xai_oauth(litellm_params) and not api_key: + api_base = XAIOAuthAuthenticator().get_api_base() + else: + api_base = ( + api_base + or litellm.api_base + or get_secret_str("XAI_API_BASE") + or XAI_API_BASE + ) # Remove trailing slashes api_base = api_base.rstrip("/") diff --git a/litellm/main.py b/litellm/main.py index 1a0d0312d73..2c416a595c4 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1638,6 +1638,7 @@ def completion( # type: ignore # noqa: PLR0915 litellm_request_debug=kwargs.get("litellm_request_debug", False), tpm=kwargs.get("tpm"), rpm=kwargs.get("rpm"), + use_xai_oauth=kwargs.get("use_xai_oauth", False), ) cast(LiteLLMLoggingObj, logging).update_environment_variables( model=model, @@ -2134,9 +2135,6 @@ def completion( # type: ignore # noqa: PLR0915 headers = headers or litellm.headers - if extra_headers is not None: - optional_params["extra_headers"] = extra_headers - ## LOAD CONFIG - if set config = litellm.OpenAITextCompletionConfig.get_config() for k, v in config.items(): @@ -2162,6 +2160,7 @@ def completion( # type: ignore # noqa: PLR0915 _response = openai_text_completions.completion( model=model, messages=messages, + headers=headers, model_response=model_response, print_verbose=print_verbose, api_key=api_key, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 282a292ab17..aab0e4264d0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1156,6 +1156,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1202,6 +1203,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1233,6 +1235,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1264,6 +1267,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1295,6 +1299,139 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "global.anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "us.anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.375e-05, + "cache_creation_input_token_cost_above_1hr": 2.2e-05, + "cache_read_input_token_cost": 1.1e-06, + "input_cost_per_token": 1.1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "eu.anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.375e-05, + "cache_creation_input_token_cost_above_1hr": 2.2e-05, + "cache_read_input_token_cost": 1.1e-06, + "input_cost_per_token": 1.1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1327,6 +1464,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1359,6 +1497,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1391,6 +1530,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1423,6 +1563,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1455,6 +1596,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1485,6 +1627,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -2208,6 +2351,37 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true + }, + "azure_ai/claude-fable-5": { + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -2237,6 +2411,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -10170,6 +10345,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -10204,6 +10380,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -10214,6 +10391,40 @@ }, "supports_output_config": true }, + "claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "provider_specific_entry": { + "us": 1.1 + }, + "supports_output_config": true + }, "claude-opus-4-8": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -10238,6 +10449,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -24180,9 +24392,12 @@ "max_output_tokens": 8192 }, "minimax/MiniMax-M3": { - "input_cost_per_token": 6e-07, - "output_cost_per_token": 2.4e-06, - "cache_read_input_token_cost": 1.2e-07, + "input_cost_per_token": 3e-07, + "input_cost_per_token_above_512k_tokens": 6e-07, + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_above_512k_tokens": 2.4e-06, + "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_above_512k_tokens": 1.2e-07, "litellm_provider": "minimax", "mode": "chat", "supports_function_calling": true, @@ -24191,7 +24406,7 @@ "supports_reasoning": true, "supports_system_messages": true, "supports_vision": true, - "max_input_tokens": 512000, + "max_input_tokens": 1000000, "max_output_tokens": 128000 }, "mistral.devstral-2-123b": { @@ -34004,6 +34219,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -34032,6 +34248,67 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true + }, + "vertex_ai/claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true + }, + "vertex_ai/claude-fable-5@default": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -34061,6 +34338,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -34090,6 +34368,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -41370,6 +41649,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", + "use_openai_responses_path": true, "supported_endpoints": ["/v1/responses"], "supported_modalities": ["text", "image"], "supported_output_modalities": ["text"], @@ -41389,6 +41669,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", + "use_openai_responses_path": true, "supported_endpoints": ["/v1/responses"], "supported_modalities": ["text", "image"], "supported_output_modalities": ["text"], @@ -41690,5 +41971,164 @@ "/v1/audio/transcriptions" ], "supports_audio_input": true + }, + "tensormesh/Qwen/Qwen3.5-397B-A17B-FP8": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 6e-07, + "output_cost_per_token": 3.6e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 4.5e-07, + "output_cost_per_token": 1.8e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/Qwen/Qwen3.6-27B-FP8": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 3.2e-07, + "output_cost_per_token": 3.2e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/lukealonso/GLM-5.1-NVFP4-MTP": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.4e-06, + "output_cost_per_token": 4.4e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 202752, + "max_output_tokens": 202752, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/deepseek-ai/DeepSeek-V4-Flash": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 2.8e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/moonshotai/Kimi-K2.6": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 9.6e-07, + "output_cost_per_token": 4e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/MiniMaxAI/MiniMax-M2.5": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 196608, + "max_output_tokens": 196608, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/google/gemma-4-31B-it": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 5.6e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/openai/gpt-oss-120b": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/openai/gpt-oss-20b": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 7e-08, + "output_cost_per_token": 2.8e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" } } \ No newline at end of file diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index c52752940c3..8edb831a9df 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -568,10 +568,12 @@ async def delete_mcp_server( """ Delete the mcp server from the db by server_id - The server-row delete is the commit point. Per-user env var rows have no FK - cascade, so they are cleaned up afterwards on a best-effort basis: a transient - failure there leaves only orphaned rows pointing at a now-missing server and - must not turn a successful delete into a caller-visible error. + The server-row delete is the commit point. Per-user credential and env var + rows have no FK cascade, so they are cleaned up afterwards on a best-effort + basis: a transient failure there leaves only orphaned rows pointing at a + now-missing server and must not turn a successful delete into a + caller-visible error. Each table is cleaned independently so a failure on one + still attempts the other. Returns the deleted mcp server record if it exists, otherwise None """ @@ -581,17 +583,20 @@ async def delete_mcp_server( }, ) if deleted_server is not None: - try: - await prisma_client.db.litellm_mcpuserenvvars.delete_many( - where={"server_id": server_id} - ) - except Exception as e: - verbose_proxy_logger.warning( - "MCP server %s deleted but per-user env var cleanup failed; " - "orphaned rows can be removed on a later delete: %s", - server_id, - e, - ) + for model, label in ( + (prisma_client.db.litellm_mcpusercredentials, "credential"), + (prisma_client.db.litellm_mcpuserenvvars, "env var"), + ): + try: + await model.delete_many(where={"server_id": server_id}) + except Exception as e: + verbose_proxy_logger.warning( + "MCP server %s deleted but per-user %s cleanup failed; " + "orphaned rows can be removed on a later delete: %s", + server_id, + label, + e, + ) return deleted_server diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index a60138dd340..51918509441 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -19,3 +19,9 @@ _mcp_active_toolset_id: ContextVar[Optional[str]] = ContextVar( _mcp_gateway_initialize_instructions: ContextVar[Optional[str]] = ContextVar( "_mcp_gateway_initialize_instructions", default=None ) + +# Per-request scoped server name; set in MCP HTTP/SSE handlers when the path +# identifies exactly one upstream server. Never populated from client-supplied headers. +_mcp_gateway_server_name: ContextVar[Optional[str]] = ContextVar( + "_mcp_gateway_server_name", default=None +) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 85ac6b399f4..73935beeb3a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -354,6 +354,52 @@ def _deserialize_json_list(data: Any) -> Optional[List[Dict[str, Any]]]: ] +def _normalize_mcp_server_cost_info(mcp_info: MCPInfo) -> None: + """Coerce ``mcp_server_cost_info`` numeric fields to ``float`` at ingest. + + YAML 1.1 parses scientific notation without a decimal point (e.g. + ``7e-05``) as a string, and ``MCPServerCostInfo`` is a TypedDict with no + runtime validation, so string-typed costs flow through to the UI and + crash its ``.toFixed`` formatting. Values that cannot be coerced are + dropped with a warning instead of failing the server load. + """ + cost_info = mcp_info.get("mcp_server_cost_info") + if not isinstance(cost_info, dict): + return + + server_name = mcp_info.get("server_name") + normalized = dict(cost_info) + + default_cost = normalized.get("default_cost_per_query") + if default_cost is not None: + try: + normalized["default_cost_per_query"] = float(default_cost) + except (TypeError, ValueError): + verbose_logger.warning( + "MCP server '%s' has non-numeric default_cost_per_query %r; ignoring it", + server_name, + default_cost, + ) + del normalized["default_cost_per_query"] + + tool_costs = normalized.get("tool_name_to_cost_per_query") + if isinstance(tool_costs, dict): + normalized_tool_costs = {} + for tool_name, cost in tool_costs.items(): + try: + normalized_tool_costs[tool_name] = float(cost) + except (TypeError, ValueError): + verbose_logger.warning( + "MCP server '%s' has non-numeric cost %r for tool '%s'; ignoring it", + server_name, + cost, + tool_name, + ) + normalized["tool_name_to_cost_per_query"] = normalized_tool_costs + + mcp_info["mcp_server_cost_info"] = normalized + + def _create_sampling_callback(user_api_key_auth: Optional[Any] = None): """ Create a sampling callback for MCP ClientSession. @@ -621,6 +667,7 @@ class MCPServerManager: mcp_info["server_name"] = server_name if "description" not in mcp_info and server_config.get("description"): mcp_info["description"] = server_config.get("description") + _normalize_mcp_server_cost_info(mcp_info) # Use alias for name if present, else server_name alias = server_config.get("alias", None) @@ -1091,6 +1138,7 @@ class MCPServerManager: mcp_info["server_name"] = mcp_server.server_name or mcp_server.server_id if "description" not in mcp_info and mcp_server.description: mcp_info["description"] = mcp_server.description + _normalize_mcp_server_cost_info(mcp_info) auth_type = cast(MCPAuthType, mcp_server.auth_type) server_url = mcp_server.url diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 0477a5d3244..731493b1337 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -47,6 +47,7 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( from litellm.proxy._experimental.mcp_server.mcp_context import ( _mcp_active_toolset_id, _mcp_gateway_initialize_instructions, + _mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug from litellm.proxy._experimental.mcp_server.utils import ( @@ -323,10 +324,14 @@ if MCP_AVAILABLE: notification_options=notification_options, experimental_capabilities=experimental_capabilities or {}, ) + updates: Dict[str, Any] = {} merged = _mcp_gateway_initialize_instructions.get() if merged is not None: - return opts.model_copy(update={"instructions": merged}) - return opts + updates["instructions"] = merged + scoped_server_name = _mcp_gateway_server_name.get() + if scoped_server_name is not None: + updates["server_name"] = scoped_server_name + return opts.model_copy(update=updates) if updates else opts ######################################################## ############ Initialize the MCP Server ################# @@ -1544,6 +1549,7 @@ if MCP_AVAILABLE: user_api_key_auth: Optional[UserAPIKeyAuth], mcp_servers: Optional[List[str]], client_ip: Optional[str], + scoped_server_endpoint: bool = False, ) -> AsyncIterator[None]: allowed = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, @@ -1565,11 +1571,22 @@ if MCP_AVAILABLE: return_exceptions=True, ) merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed) - tok = _mcp_gateway_initialize_instructions.set(merged) + scoped_server_name = None + if scoped_server_endpoint and len(allowed) == 1: + scoped_server = allowed[0] + scoped_server_name = ( + scoped_server.alias + or scoped_server.server_name + or scoped_server.name + or scoped_server.server_id + ) + instructions_token = _mcp_gateway_initialize_instructions.set(merged) + server_name_token = _mcp_gateway_server_name.set(scoped_server_name) try: yield finally: - _mcp_gateway_initialize_instructions.reset(tok) + _mcp_gateway_initialize_instructions.reset(instructions_token) + _mcp_gateway_server_name.reset(server_name_token) async def _get_tools_from_mcp_servers( # noqa: PLR0915 user_api_key_auth: Optional[UserAPIKeyAuth], @@ -3620,6 +3637,7 @@ if MCP_AVAILABLE: oauth2_headers, raw_headers, ) = await extract_mcp_auth_context(scope, path) + scoped_server_endpoint = len(_get_mcp_servers_in_path(path) or []) == 1 # Extract client IP for MCP access control _client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) @@ -3896,6 +3914,7 @@ if MCP_AVAILABLE: user_api_key_auth, mcp_servers, _client_ip, + scoped_server_endpoint=scoped_server_endpoint, ): await target_manager.handle_request(scope, receive, local_send) if use_stateful and session_id and scope.get("method") == "DELETE": @@ -3980,6 +3999,7 @@ if MCP_AVAILABLE: oauth2_headers, raw_headers, ) = await extract_mcp_auth_context(scope, path) + scoped_server_endpoint = len(_get_mcp_servers_in_path(path) or []) == 1 # Extract client IP for MCP access control _sse_client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) @@ -4052,6 +4072,7 @@ if MCP_AVAILABLE: user_api_key_auth, mcp_servers, _sse_client_ip, + scoped_server_endpoint=scoped_server_endpoint, ): await sse_session_manager.handle_request(scope, receive, send) except MCPUpstreamAuthError as e: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f1443edf455..1b594e20d32 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -23,6 +23,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( from litellm.types.integrations.slack_alerting import AlertType from litellm.types.llms.openai import ( AllMessageValues, + ResponsesAPIResponse, ) from litellm.types.mcp import ( MCPAuthType, @@ -2176,6 +2177,17 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "`statement_cache_size`). Keys here override any default LiteLLM sets." ), ) + database_disable_prepared_statements: Optional[bool] = Field( + None, + description=( + "Disable server-side prepared statements by setting Prisma's " + "`pgbouncer=true` URL param. Use this for pgbouncer transaction-pooling " + "deployments, or to prevent the 'cached plan must not change result " + "type' error that pooled connections hit during rolling schema " + "migrations. An explicit `pgbouncer` in `database_extra_connection_params` " + "takes precedence." + ), + ) database_type: Optional[Literal["dynamo_db"]] = Field( None, description="to use dynamodb instead of postgres db" ) @@ -3834,6 +3846,7 @@ PassThroughEndpointLoggingResultValues = Union[ EmbeddingResponse, VideoObject, StandardPassThroughResponseObject, + ResponsesAPIResponse, ] diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 45007861d55..6eae9d0d475 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3516,10 +3516,13 @@ async def _virtual_key_max_budget_check( if valid_token.max_budget is not None: from litellm.proxy.proxy_server import get_current_spend + fallback_spend = valid_token.spend or 0.0 + counter_key = f"spend:key:{valid_token.token}" + # Read spend from cross-pod counter (Redis-first) or cached object (fallback) spend = await get_current_spend( - counter_key=f"spend:key:{valid_token.token}", - fallback_spend=valid_token.spend or 0.0, + counter_key=counter_key, + fallback_spend=fallback_spend, ) #################################### diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index f76949f4d11..83f18173182 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -168,6 +168,16 @@ class UserAPIKeyAuthExceptionHandler: ) elif isinstance(e, ProxyException): raise e + if PrismaDBExceptionHandler.is_database_service_unavailable_error(e): + raise ProxyException( + message=( + "Service Unavailable, the authentication database is " + "temporarily unreachable. Please retry shortly." + ), + type=ProxyErrorTypes.no_db_connection, + param="None", + code=status.HTTP_503_SERVICE_UNAVAILABLE, + ) raise ProxyException( message="Authentication Error, " + str(e), type=ProxyErrorTypes.auth_error, diff --git a/litellm/proxy/client/README.md b/litellm/proxy/client/README.md index 9fbc6f2197d..c2ce28884c7 100644 --- a/litellm/proxy/client/README.md +++ b/litellm/proxy/client/README.md @@ -338,9 +338,9 @@ sequenceDiagram The CLI provides three authentication commands: -- **`litellm-proxy login`** - Start SSO authentication flow -- **`litellm-proxy logout`** - Clear stored authentication token -- **`litellm-proxy whoami`** - Show current authentication status +- **`lite login`** - Start SSO authentication flow +- **`lite logout`** - Clear stored authentication token +- **`lite whoami`** - Show current authentication status ### Authentication Flow Steps @@ -382,14 +382,14 @@ Once authenticated, the CLI will automatically use the stored token for all requ ```bash # Login -litellm-proxy login +lite login # Use CLI without specifying API key -litellm-proxy models list +lite models list # Check authentication status -litellm-proxy whoami +lite whoami # Logout -litellm-proxy logout +lite logout ``` diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index 6ef837cb521..333e2029e46 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -22,11 +22,11 @@ The CLI can be configured using environment variables or command-line options: Example: ```bash -litellm-proxy version +lite version # or -litellm-proxy --version +lite --version # or -litellm-proxy -v +lite -v ``` ## Commands @@ -40,7 +40,7 @@ The CLI provides several commands for managing models on your LiteLLM proxy serv View all available models: ```bash -litellm-proxy models list [--format table|json] +lite models list [--format table|json] ``` Options: @@ -52,7 +52,7 @@ Options: Get detailed information about all models: ```bash -litellm-proxy models info [options] +lite models info [options] ``` Options: @@ -75,7 +75,7 @@ Default columns: `public_model`, `upstream_model`, `updated_at` Add a new model to the proxy: ```bash -litellm-proxy models add [options] +lite models add [options] ``` Options: @@ -86,7 +86,7 @@ Options: Example: ```bash -litellm-proxy models add gpt-4 -p api_key=sk-123 -p api_base=https://api.openai.com -i description="GPT-4 model" +lite models add gpt-4 -p api_key=sk-123 -p api_base=https://api.openai.com -i description="GPT-4 model" ``` #### Get Model Info @@ -94,7 +94,7 @@ litellm-proxy models add gpt-4 -p api_key=sk-123 -p api_base=https://api.openai. Get information about a specific model: ```bash -litellm-proxy models get [--id MODEL_ID] [--name MODEL_NAME] +lite models get [--id MODEL_ID] [--name MODEL_NAME] ``` Options: @@ -107,7 +107,7 @@ Options: Delete a model from the proxy: ```bash -litellm-proxy models delete +lite models delete ``` #### Update Model @@ -115,7 +115,7 @@ litellm-proxy models delete Update an existing model's configuration: ```bash -litellm-proxy models update [options] +lite models update [options] ``` Options: @@ -128,7 +128,7 @@ Options: Import models from a YAML file: ```bash -litellm-proxy models import models.yaml +lite models import models.yaml ``` Options: @@ -142,31 +142,31 @@ Examples: 1. Import all models from a YAML file: ```bash -litellm-proxy models import models.yaml +lite models import models.yaml ``` 2. Dry run (show what would be imported): ```bash -litellm-proxy models import models.yaml --dry-run +lite models import models.yaml --dry-run ``` 3. Only import models where the model name contains 'gpt': ```bash -litellm-proxy models import models.yaml --only-models-matching-regex gpt +lite models import models.yaml --only-models-matching-regex gpt ``` 4. Only import models with access group containing 'beta': ```bash -litellm-proxy models import models.yaml --only-access-groups-matching-regex beta +lite models import models.yaml --only-access-groups-matching-regex beta ``` 5. Combine both filters: ```bash -litellm-proxy models import models.yaml --only-models-matching-regex gpt --only-access-groups-matching-regex beta +lite models import models.yaml --only-models-matching-regex gpt --only-access-groups-matching-regex beta ``` ### Credentials Management @@ -178,7 +178,7 @@ The CLI provides commands for managing credentials on your LiteLLM proxy server: View all available credentials: ```bash -litellm-proxy credentials list [--format table|json] +lite credentials list [--format table|json] ``` Options: @@ -194,7 +194,7 @@ The table format displays: Create a new credential: ```bash -litellm-proxy credentials create --info --values +lite credentials create --info --values ``` Options: @@ -205,7 +205,7 @@ Options: Example: ```bash -litellm-proxy credentials create azure-cred \ +lite credentials create azure-cred \ --info '{"custom_llm_provider": "azure"}' \ --values '{"api_key": "sk-123", "api_base": "https://example.azure.openai.com"}' ``` @@ -215,7 +215,7 @@ litellm-proxy credentials create azure-cred \ Get information about a specific credential: ```bash -litellm-proxy credentials get +lite credentials get ``` #### Delete Credential @@ -223,7 +223,7 @@ litellm-proxy credentials get Delete a credential: ```bash -litellm-proxy credentials delete +lite credentials delete ``` ### Keys Management @@ -235,7 +235,7 @@ The CLI provides commands for managing API keys on your LiteLLM proxy server: View all API keys: ```bash -litellm-proxy keys list [--format table|json] [options] +lite keys list [--format table|json] [options] ``` Options: @@ -256,7 +256,7 @@ Options: Generate a new API key: ```bash -litellm-proxy keys generate [options] +lite keys generate [options] ``` Options: @@ -274,7 +274,7 @@ Options: Example: ```bash -litellm-proxy keys generate --models gpt-4,gpt-3.5-turbo --spend 100 --duration 24h --key-alias my-key --team-id team123 +lite keys generate --models gpt-4,gpt-3.5-turbo --spend 100 --duration 24h --key-alias my-key --team-id team123 ``` #### Delete Keys @@ -282,7 +282,7 @@ litellm-proxy keys generate --models gpt-4,gpt-3.5-turbo --spend 100 --duration Delete API keys by key or alias: ```bash -litellm-proxy keys delete [--keys ] [--key-aliases ] +lite keys delete [--keys ] [--key-aliases ] ``` Options: @@ -293,7 +293,7 @@ Options: Example: ```bash -litellm-proxy keys delete --keys sk-key1,sk-key2 --key-aliases alias1,alias2 +lite keys delete --keys sk-key1,sk-key2 --key-aliases alias1,alias2 ``` #### Get Key Info @@ -301,7 +301,7 @@ litellm-proxy keys delete --keys sk-key1,sk-key2 --key-aliases alias1,alias2 Get information about a specific API key: ```bash -litellm-proxy keys info --key +lite keys info --key ``` Options: @@ -311,7 +311,7 @@ Options: Example: ```bash -litellm-proxy keys info --key sk-key1 +lite keys info --key sk-key1 ``` ### User Management @@ -323,7 +323,7 @@ The CLI provides commands for managing users on your LiteLLM proxy server: View all users: ```bash -litellm-proxy users list +lite users list ``` #### Get User Info @@ -331,7 +331,7 @@ litellm-proxy users list Get information about a specific user: ```bash -litellm-proxy users get --id +lite users get --id ``` #### Create User @@ -339,7 +339,7 @@ litellm-proxy users get --id Create a new user: ```bash -litellm-proxy users create --email user@example.com --role internal_user --alias "Alice" --team team1 --max-budget 100.0 +lite users create --email user@example.com --role internal_user --alias "Alice" --team team1 --max-budget 100.0 ``` #### Delete User @@ -347,7 +347,7 @@ litellm-proxy users create --email user@example.com --role internal_user --alias Delete one or more users by user_id: ```bash -litellm-proxy users delete +lite users delete ``` ### Chat Commands @@ -359,7 +359,7 @@ The CLI provides commands for interacting with chat models through your LiteLLM Create a chat completion: ```bash -litellm-proxy chat completions [options] +lite chat completions [options] ``` Arguments: @@ -379,12 +379,12 @@ Examples: 1. Simple completion: ```bash -litellm-proxy chat completions gpt-4 -m "user:Hello, how are you?" +lite chat completions gpt-4 -m "user:Hello, how are you?" ``` 2. Multi-message conversation: ```bash -litellm-proxy chat completions gpt-4 \ +lite chat completions gpt-4 \ -m "system:You are a helpful assistant" \ -m "user:What's the capital of France?" \ -m "assistant:The capital of France is Paris." \ @@ -393,7 +393,7 @@ litellm-proxy chat completions gpt-4 \ 3. With generation parameters: ```bash -litellm-proxy chat completions gpt-4 \ +lite chat completions gpt-4 \ -m "user:Write a story" \ --temperature 0.7 \ --max-tokens 500 \ @@ -409,7 +409,7 @@ The CLI provides commands for making direct HTTP requests to your LiteLLM proxy Make an HTTP request to any endpoint: ```bash -litellm-proxy http request [options] +lite http request [options] ``` Arguments: @@ -425,19 +425,46 @@ Examples: 1. List models: ```bash -litellm-proxy http request GET /models +lite http request GET /models ``` 2. Create a chat completion: ```bash -litellm-proxy http request POST /chat/completions -j '{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]}' +lite http request POST /chat/completions -j '{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]}' ``` 3. Test connection with custom headers: ```bash -litellm-proxy http request GET /health/test_connection -H "X-Custom-Header:value" +lite http request GET /health/test_connection -H "X-Custom-Header:value" ``` +### Run a Coding Agent + +Launch a coding agent with all of its LLM traffic routed through your LiteLLM proxy. Each supported agent is its own command, so there is nothing to remember beyond the agent's name: + +```bash +lite claude +lite codex +lite opencode +``` + +Anything you type after the agent name is forwarded to it untouched, so the usual flags keep working: + +```bash +lite claude --resume +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. 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). + +Options (these belong to the wrapper, so put them before the agent's own flags): + +- `--skip-verify`: Skip the pre-launch key check (useful offline or with non-standard auth). + +To pin the model, pass the agent's own model flag (for example `lite claude --model my-proxy-model` or `lite codex -m my-proxy-model`), or export the variable the agent reads (`ANTHROPIC_MODEL` / `ANTHROPIC_SMALL_FAST_MODEL` for Claude Code); the wrapper preserves anything you already have set. Whatever model the agent ends up requesting must exist on the proxy, since requests land on the proxy's `/v1/messages` (Anthropic) or `/v1/chat/completions` and `/v1/responses` (OpenAI) endpoints. + ## Environment Variables The CLI respects the following environment variables: @@ -450,37 +477,37 @@ The CLI respects the following environment variables: 1. List all models in table format: ```bash -litellm-proxy models list +lite models list ``` 2. Add a new model with parameters: ```bash -litellm-proxy models add gpt-4 -p api_key=sk-123 -p max_tokens=2048 +lite models add gpt-4 -p api_key=sk-123 -p max_tokens=2048 ``` 3. Get model information in JSON format: ```bash -litellm-proxy models info --format json +lite models info --format json ``` 4. Update model parameters: ```bash -litellm-proxy models update model-123 -p temperature=0.7 -i description="Updated model" +lite models update model-123 -p temperature=0.7 -i description="Updated model" ``` 5. List all credentials in table format: ```bash -litellm-proxy credentials list +lite credentials list ``` 6. Create a new credential for Azure: ```bash -litellm-proxy credentials create azure-prod \ +lite credentials create azure-prod \ --info '{"custom_llm_provider": "azure"}' \ --values '{"api_key": "sk-123", "api_base": "https://prod.azure.openai.com"}' ``` @@ -488,7 +515,7 @@ litellm-proxy credentials create azure-prod \ 7. Make a custom HTTP request: ```bash -litellm-proxy http request POST /chat/completions \ +lite http request POST /chat/completions \ -j '{"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]}' \ -H "X-Custom-Header:value" ``` @@ -497,29 +524,29 @@ litellm-proxy http request POST /chat/completions \ ```bash # List users -litellm-proxy users list +lite users list # Get user info -litellm-proxy users get --id u1 +lite users get --id u1 # Create a user -litellm-proxy users create --email a@b.com --role internal_user --alias "Alice" --team team1 --max-budget 100.0 +lite users create --email a@b.com --role internal_user --alias "Alice" --team team1 --max-budget 100.0 # Delete users -litellm-proxy users delete u1 u2 +lite users delete u1 u2 ``` 9. Import models from a YAML file (with filters): ```bash # Only import models where the model name contains 'gpt' -litellm-proxy models import models.yaml --only-models-matching-regex gpt +lite models import models.yaml --only-models-matching-regex gpt # Only import models with access group containing 'beta' -litellm-proxy models import models.yaml --only-access-groups-matching-regex beta +lite models import models.yaml --only-access-groups-matching-regex beta # Combine both filters -litellm-proxy models import models.yaml --only-models-matching-regex gpt --only-access-groups-matching-regex beta +lite models import models.yaml --only-models-matching-regex gpt --only-access-groups-matching-regex beta ``` ## Error Handling diff --git a/litellm/proxy/client/cli/commands/agents.py b/litellm/proxy/client/cli/commands/agents.py new file mode 100644 index 00000000000..f39ffb3e864 --- /dev/null +++ b/litellm/proxy/client/cli/commands/agents.py @@ -0,0 +1,303 @@ +import os +import shutil +import sys +from typing import Callable, Dict, FrozenSet, List, Mapping, Optional, Sequence, Tuple + +import click +import requests + +from .auth import get_stored_api_key, login + +ANTHROPIC_BASE_URL_ENV = "ANTHROPIC_BASE_URL" +ANTHROPIC_AUTH_TOKEN_ENV = "ANTHROPIC_AUTH_TOKEN" +ANTHROPIC_API_KEY_ENV = "ANTHROPIC_API_KEY" +OPENAI_BASE_URL_ENV = "OPENAI_BASE_URL" +OPENAI_API_KEY_ENV = "OPENAI_API_KEY" + +PROFILE_ANTHROPIC = "anthropic" +PROFILE_OPENAI = "openai" + +_KNOWN_AGENTS: Dict[str, Tuple[str, FrozenSet[str]]] = { + "claude": ("Claude Code", frozenset({PROFILE_ANTHROPIC})), + "codex": ("Codex", frozenset({PROFILE_OPENAI})), + "opencode": ("OpenCode", frozenset({PROFILE_OPENAI})), +} + +_INSTALL_DOCS: Dict[str, str] = { + "claude": "https://docs.claude.com/en/docs/claude-code/setup", + "codex": "https://developers.openai.com/codex/cli", + "opencode": "https://opencode.ai/docs", +} + +CODEX_PROXY_PROVIDER = "litellm" + + +class AgentRunError(Exception): + """Raised for any user-actionable failure while preparing to run an agent.""" + + +def agent_profile(command: str) -> Tuple[str, FrozenSet[str]]: + """Return the (display name, env profiles) for a wrapped command. + + Known agents map to the API family they speak. Anything else gets both + families so it works regardless of which env vars the tool reads. + """ + base = os.path.basename(command) + if base in _KNOWN_AGENTS: + return _KNOWN_AGENTS[base] + return base, frozenset({PROFILE_ANTHROPIC, PROFILE_OPENAI}) + + +def build_agent_env( + base_env: Mapping[str, str], + base_url: str, + api_key: str, + profiles: FrozenSet[str], +) -> Dict[str, str]: + """Return a copy of base_env wired to route the agent through the proxy. + + Anthropic clients (Claude Code) append /v1/messages to ANTHROPIC_BASE_URL, + so it stays the bare proxy root; OpenAI clients (Codex, OpenCode) expect the + /v1 suffix on OPENAI_BASE_URL. ANTHROPIC_API_KEY is dropped so a stray + Anthropic key cannot win over the bearer token we set. + """ + env = dict(base_env) + root = base_url.rstrip("/") + if PROFILE_ANTHROPIC in profiles: + env[ANTHROPIC_BASE_URL_ENV] = root + env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key + env.pop(ANTHROPIC_API_KEY_ENV, None) + if PROFILE_OPENAI in profiles: + env[OPENAI_BASE_URL_ENV] = root + "/v1" + env[OPENAI_API_KEY_ENV] = api_key + return env + + +def _codex_proxy_args(base_url: str) -> List[str]: + """Codex `-c` overrides that point it at the proxy. + + Codex ignores OPENAI_BASE_URL (it always dials api.openai.com), so the env + profile alone cannot route it. It does honor a custom provider, so define one + inline; supports_websockets=false forces the HTTP/SSE Responses transport + because the proxy does not speak the Responses WebSocket protocol. The key is + read from OPENAI_API_KEY, which build_agent_env already exports. + """ + root = base_url.rstrip("/") + "/v1" + provider = f"model_providers.{CODEX_PROXY_PROVIDER}" + return [ + "-c", + f'model_provider="{CODEX_PROXY_PROVIDER}"', + "-c", + f'{provider}.name="LiteLLM proxy"', + "-c", + f'{provider}.base_url="{root}"', + "-c", + f'{provider}.env_key="{OPENAI_API_KEY_ENV}"', + "-c", + f'{provider}.wire_api="responses"', + "-c", + f"{provider}.supports_websockets=false", + ] + + +_PROXY_ARGS: Dict[str, Callable[[str], List[str]]] = { + "codex": _codex_proxy_args, +} + + +def agent_launch_args(command: str, base_url: str) -> List[str]: + """Extra CLI args an agent needs to actually honor the proxy. + + Claude Code and OpenCode respect the exported env vars, so they get nothing + here; Codex needs its provider pointed via config overrides. + """ + builder = _PROXY_ARGS.get(os.path.basename(command)) + return builder(base_url) if builder else [] + + +def verify_proxy_key( + base_url: str, + api_key: str, + *, + get: Callable[..., requests.Response] = requests.get, +) -> None: + """Probe the proxy with the key so bad creds fail here, not inside the agent. + + Raises AgentRunError when the proxy is unreachable or rejects the key. Other + non-2xx responses are tolerated; the agent's own call is the real test. + """ + url = base_url.rstrip("/") + "/v1/models" + try: + resp = get(url, headers={"Authorization": f"Bearer {api_key}"}, timeout=10) + except requests.RequestException as e: + raise AgentRunError( + f"Could not reach the LiteLLM proxy at {base_url.rstrip('/')}: {e}. " + "Is it running, and is --base-url (or LITELLM_PROXY_URL) correct?" + ) + if resp.status_code in (401, 403): + raise AgentRunError( + f"LiteLLM rejected your key (HTTP {resp.status_code}). " + "Run `lite login` to refresh it, or pass a valid --api-key." + ) + + +def _exec(path: str, args: Sequence[str], env: Mapping[str, str]) -> None: + os.execvpe(path, list(args), dict(env)) + + +def _restore_controlling_terminal() -> None: + """Reattach the controlling terminal to stdin before handing off to the agent. + + Completing the browser SSO login can leave stdin detached from the terminal, + which makes a TUI agent like Claude Code start in non-interactive mode and + exit immediately. Reopening /dev/tty onto fd 0 gives the agent a live + terminal; when stdin is still a tty (no login happened) this is a no-op. + """ + if sys.stdin.isatty(): + return + try: + fd = os.open("/dev/tty", os.O_RDONLY) + except OSError: + return + try: + os.dup2(fd, 0) + finally: + os.close(fd) + + +def run_agent( + base_url: str, + api_key: str, + command: Sequence[str], + *, + skip_verify: bool = False, + base_env: Optional[Mapping[str, str]] = None, + which: Callable[[str], Optional[str]] = shutil.which, + verify: Callable[[str, str], None] = verify_proxy_key, + launcher: Callable[[str, Sequence[str], Mapping[str, str]], None] = _exec, + reattach_terminal: Optional[Callable[[], None]] = None, +) -> None: + """Validate, wire the environment, and hand off to the agent. + + On success this replaces the current process and never returns. Raises + AgentRunError for missing binaries, an unreachable proxy, or a rejected key. + reattach_terminal, when given, runs just before handoff to restore stdin. + """ + if not command: + raise AgentRunError("Nothing to run.") + + _, profiles = agent_profile(command[0]) + binary = which(command[0]) + if binary is None: + docs = _INSTALL_DOCS.get(os.path.basename(command[0])) + hint = f" Install it first: {docs}" if docs else "" + raise AgentRunError(f"Could not find `{command[0]}` on your PATH.{hint}") + + if not skip_verify: + verify(base_url, api_key) + + env = build_agent_env( + base_env if base_env is not None else os.environ, + base_url, + api_key, + profiles, + ) + extra_args = agent_launch_args(command[0], base_url) + if reattach_terminal is not None: + reattach_terminal() + launcher(binary, [command[0], *extra_args, *command[1:]], env) + + +def _is_interactive() -> bool: + return sys.stdin.isatty() + + +def _resolve_api_key(ctx: click.Context) -> str: + base_url = ctx.obj["base_url"] + api_key = ctx.obj.get("api_key") + if api_key: + return api_key + + if not _is_interactive(): + raise click.ClickException( + "No LiteLLM key found. Set LITELLM_PROXY_API_KEY (or pass --api-key) for " + "non-interactive use, or run `lite login` from a terminal." + ) + + click.echo("No LiteLLM credentials found; starting login...") + ctx.invoke(login) + api_key = get_stored_api_key(expected_base_url=base_url) + if not api_key: + raise click.ClickException( + "Login did not produce an API key; cannot start the agent." + ) + return api_key + + +_SKIP_VERIFY_HELP = "Skip the pre-launch key check against the proxy." + + +def _launch( + ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify: bool +) -> None: + base_url = ctx.obj["base_url"] + started_interactive = _is_interactive() + api_key = _resolve_api_key(ctx) + + display_name, _ = agent_profile(binary) + click.echo( + f"litellm: routing {display_name} through proxy at {base_url.rstrip('/')}" + ) + + try: + run_agent( + base_url, + api_key, + [binary, *args], + skip_verify=skip_verify, + reattach_terminal=( + _restore_controlling_terminal if started_interactive else None + ), + ) + except AgentRunError as e: + raise click.ClickException(str(e)) + + +def _make_agent_command(binary: str, display_name: str) -> click.Command: + @click.command( + name=binary, + context_settings={"ignore_unknown_options": True}, + short_help=f"Run {display_name} through your LiteLLM proxy", + ) + @click.option("--skip-verify", is_flag=True, default=False, help=_SKIP_VERIFY_HELP) + @click.argument("args", nargs=-1, type=click.UNPROCESSED) + @click.pass_context + def _command(ctx: click.Context, skip_verify: bool, args: Sequence[str]) -> None: + _launch(ctx, binary, list(args), skip_verify=skip_verify) + + _command.help = ( + f"Run {display_name} routed through your LiteLLM proxy.\n\n" + f"Logs in with LiteLLM if needed, verifies your key against the proxy, " + f"exports the env vars {binary} reads, then hands off. Any arguments are " + f"forwarded to `{binary}`." + ) + return _command + + +def agent_commands() -> List[click.Command]: + """Build one top-level command per known agent, e.g. `lite claude`.""" + return [ + _make_agent_command(binary, name) + for binary, (name, _profiles) in _KNOWN_AGENTS.items() + ] + + +__all__ = [ + "agent_commands", + "run_agent", + "build_agent_env", + "agent_launch_args", + "verify_proxy_key", + "agent_profile", + "AgentRunError", +] diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 447837c35e7..b06d86d5965 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -624,7 +624,7 @@ def whoami(): token_data = load_token() if not token_data: - click.echo("❌ Not authenticated. Run 'litellm-proxy login' to authenticate.") + click.echo("❌ Not authenticated. Run 'lite login' to authenticate.") return click.echo("✅ Authenticated") diff --git a/litellm/proxy/client/cli/commands/chat.py b/litellm/proxy/client/cli/commands/chat.py index a078b766107..696e34c3ecd 100644 --- a/litellm/proxy/client/cli/commands/chat.py +++ b/litellm/proxy/client/cli/commands/chat.py @@ -122,13 +122,13 @@ def chat( Examples: # Chat with a specific model - litellm-proxy chat gpt-4 + lite chat gpt-4 # Chat without specifying model (will show model selection) - litellm-proxy chat + lite chat # Chat with custom settings - litellm-proxy chat gpt-4 --temperature 0.9 --system "You are a helpful coding assistant" + lite chat gpt-4 --temperature 0.9 --system "You are a helpful coding assistant" """ console = Console() diff --git a/litellm/proxy/client/cli/interface.py b/litellm/proxy/client/cli/interface.py index eba693dc18e..a32d60aadd9 100644 --- a/litellm/proxy/client/cli/interface.py +++ b/litellm/proxy/client/cli/interface.py @@ -80,6 +80,8 @@ def styled_prompt(): def show_commands(): """Display available commands.""" + from .commands.agents import agent_commands + commands = [ ("login", "Authenticate with the LiteLLM proxy server"), ("logout", "Clear stored authentication"), @@ -91,6 +93,9 @@ def show_commands(): ("keys", "Manage API keys"), ("teams", "Manage teams and team assignments"), ("users", "Manage users"), + ] + commands += [(c.name, c.get_short_help_str()) for c in agent_commands()] + commands += [ ("version", "Show version information"), ("help", "Show this help message"), ("quit", "Exit the interactive session"), @@ -156,7 +161,7 @@ def execute_command(user_input: str, ctx: click.Context): # Execute the command try: # Create a new argument list for click to parse - sys.argv = ["litellm-proxy"] + [command] + args + sys.argv = ["lite"] + [command] + args # Get the command object and invoke it cmd = cli.commands[command] diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index be55f79c066..b8c483f4b08 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -7,6 +7,7 @@ import click from litellm._version import version as litellm_version from litellm.proxy.client.health import HealthManagementClient +from .commands.agents import agent_commands from .commands.auth import get_stored_api_key, login, logout, whoami from .commands.chat import chat from .commands.credentials import credentials @@ -112,6 +113,9 @@ cli.add_command(keys) cli.add_command(teams) # Add the users command group cli.add_command(users) +# Add a top-level command per coding agent (claude, codex, opencode, ...) +for agent_command in agent_commands(): + cli.add_command(agent_command) if __name__ == "__main__": diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 6558543370d..b9a9f3cebb7 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1949,12 +1949,26 @@ class ProxyBaseLLMRequestProcessing: code=status.HTTP_400_BAD_REQUEST, headers=headers, ) + # Extract status_code from the exception if it carries one. + # Provider exceptions (NotFoundError, BadRequestError, GeminiError, + # VertexAIError, etc.) all have a status_code attribute reflecting + # the upstream API response. Use it to return the correct HTTP code + # instead of defaulting to 500. + _exc_status_code = getattr(e, "status_code", None) + if ( + _exc_status_code is not None + and isinstance(_exc_status_code, int) + and 400 <= _exc_status_code <= 599 + ): + _code = _exc_status_code + else: + _code = status.HTTP_500_INTERNAL_SERVER_ERROR raise ProxyException( message=getattr(e, "message", error_msg), type=getattr(e, "type", "None"), param=getattr(e, "param", "None"), openai_code=getattr(e, "code", None), - code=getattr(e, "status_code", 500), + code=_code, provider_specific_fields=getattr(e, "provider_specific_fields", None), headers=headers, ) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index a65e737f248..c630294c1ec 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -40,8 +40,10 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 premium_user: bool, config_file_path: str, litellm_settings: dict, - callback_specific_params: dict = {}, + callback_specific_params: Optional[dict] = None, ): + if not isinstance(callback_specific_params, dict): + callback_specific_params = {} from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.logging_callback_manager import ( LoggingCallbackManager, @@ -166,7 +168,12 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 ) init_params = {} - if "lakera_prompt_injection" in callback_specific_params: + if ( + "lakera_prompt_injection" in callback_specific_params + and isinstance( + callback_specific_params["lakera_prompt_injection"], dict + ) + ): init_params = callback_specific_params["lakera_prompt_injection"] lakera_moderations_object = lakeraAI_Moderation(**init_params) imported_list.append(lakera_moderations_object) diff --git a/litellm/proxy/common_utils/html_forms/cli_sso_success.py b/litellm/proxy/common_utils/html_forms/cli_sso_success.py index 51f0775d90b..345f3ca5b42 100644 --- a/litellm/proxy/common_utils/html_forms/cli_sso_success.py +++ b/litellm/proxy/common_utils/html_forms/cli_sso_success.py @@ -135,7 +135,7 @@ def render_cli_sso_success_page() -> str: font-size: 14px; }} - .countdown {{ + .status {{ color: #64748b; font-size: 14px; font-weight: 500; @@ -183,23 +183,11 @@ def render_cli_sso_success_page() -> str:

You can now use LiteLLM CLI commands with your authenticated session.

-
This window will close in 3 seconds...
+
You can now close this window and return to your terminal.
- + diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index ab9d341aa51..c500e727595 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -109,6 +109,92 @@ class PrismaDBExceptionHandler: return True return False + @staticmethod + def is_prisma_engine_internal_error(e: Exception) -> bool: + """True iff ``e`` is a non-``PrismaError`` exception raised from inside + prisma-client-py's query-engine layer. + + During the instant a DB connection is torn down, the query engine can + return a malformed error payload (``user_facing_error.meta`` is + ``null``). prisma-client-py's ``handle_response_errors`` then crashes + with ``AttributeError: 'NoneType' object has no attribute 'get'`` + before it can raise the proper P1001 "can't reach database server" + error. That AttributeError carries no connection keyword, so it can't + be matched by message; identify it by its ``prisma.engine`` origin + instead. + + Recognized ``PrismaError`` subclasses are excluded: connectivity ones + are already classified by type/keyword above, and data-layer ones + (the DB IS reachable) must stay 401. + """ + import prisma + + if isinstance(e, prisma.errors.PrismaError): + return False + tb = getattr(e, "__traceback__", None) + while tb is not None: + if tb.tb_frame.f_globals.get("__name__", "").startswith("prisma.engine"): + return True + tb = tb.tb_next + return False + + @staticmethod + def is_database_service_unavailable_error(e: Exception) -> bool: + """True iff the exception means the database could not answer at the + infrastructure level (connection refused, socket/interface failure, + timeout) rather than a genuine auth failure (key not found) or a + data-layer error (the DB IS reachable and rejected the data). + + Auth must answer 401 only for a key the DB confirms is invalid. When + the DB itself is unreachable, the request has to surface as 503 so + callers retry instead of treating valid keys as invalid during an + outage. + + Note: prisma-client-py mislabels the P1001 "can't reach database + server" connectivity failure as a ``DataError`` (a data-layer type), + so a type-only check misses real outages. ``is_database_transport_error`` + keyword-matches the connection message and catches that masquerade, + while genuine data errors (no connection keyword) correctly stay 401. + + The Postgres "cached plan must not change result type" error is matched + here, not in ``is_database_transport_error``: it is a transient stale-DB- + state condition (not an invalid key), but the connection is healthy so it + must not trigger a reconnect. + + A non-``PrismaError`` raised from inside the prisma query engine (e.g. + the ``AttributeError`` from ``handle_response_errors`` when the engine + returns a malformed error payload mid-tear-down) is also treated as + unavailable; see ``is_prisma_engine_internal_error``. + """ + import asyncio + + if PrismaDBExceptionHandler.is_database_connection_error(e): + return True + if PrismaDBExceptionHandler.is_database_transport_error(e): + return True + if PrismaDBExceptionHandler.is_prisma_engine_internal_error(e): + return True + if "cached plan must not change result type" in str(e).lower(): + return True + + # OSError already covers ConnectionError and (Py3.3+) TimeoutError. + # asyncio.TimeoutError is a distinct class before Py3.11. + if isinstance(e, (OSError, asyncio.TimeoutError)): + return True + + try: + import asyncpg + except ImportError: + return False + + return isinstance( + e, + ( + asyncpg.exceptions.PostgresConnectionError, + asyncpg.exceptions.InterfaceError, + ), + ) + @staticmethod def handle_db_exception(e: Exception): """ diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index bbcc7396d67..08bc8944b92 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ToolDiscoveryQueueItem +from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.table_repositories import ToolRepository from litellm.types.tool_management import ( @@ -309,7 +310,11 @@ class ToolPolicyRegistry: async def sync_tool_policy_from_db(self, prisma_client: "PrismaClient") -> None: """Load all tool policies and object-permission blocked_tools from DB.""" try: - tools = await ToolRepository(prisma_client).table.find_many() + tools = await call_with_db_reconnect_retry( + prisma_client, + lambda: ToolRepository(prisma_client).table.find_many(), + reason="sync_tool_policy_from_db_tools_lookup_failure", + ) self._tool_input_policies = { row.tool_name: getattr(row, "input_policy", "untrusted") or "untrusted" for row in tools @@ -319,7 +324,11 @@ class ToolPolicyRegistry: for row in tools } - perms = await ObjectPermissionRepository(prisma_client).table.find_many() + perms = await call_with_db_reconnect_retry( + prisma_client, + lambda: ObjectPermissionRepository(prisma_client).table.find_many(), + reason="sync_tool_policy_from_db_perms_lookup_failure", + ) self._blocked_tools_by_op_id = {} for row in perms: op_id = getattr(row, "object_permission_id", None) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index fc414ab7b54..033e1d0b8e7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -1225,6 +1225,39 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): for chunk in all_chunks: yield chunk + @staticmethod + def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: Dict[str, str]) -> bytes: + try: + text = chunk.decode("utf-8") + except UnicodeDecodeError: + return chunk + + result_lines: List[str] = [] + for line in text.split("\n"): + line = line.rstrip("\r") + if line.startswith("data: ") and line != "data: [DONE]": + raw_json = line[6:] + try: + event = json.loads(raw_json) + delta = event.get("delta") if isinstance(event, dict) else None + if ( + isinstance(delta, dict) + and event.get("type") == "content_block_delta" + and delta.get("type") == "text_delta" + and isinstance(delta.get("text"), str) + ): + unmasked = _OPTIONAL_PresidioPIIMasking._unmask_pii_text( + delta["text"], pii_tokens + ) + if unmasked != delta["text"]: + event["delta"]["text"] = unmasked + line = "data: " + json.dumps(event, ensure_ascii=False) + except (json.JSONDecodeError, KeyError, TypeError): + pass + result_lines.append(line) + + return "\n".join(result_lines).encode("utf-8") + async def _stream_pii_unmasking( self, response: Any, @@ -1237,13 +1270,19 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): from litellm.main import stream_chunk_builder from litellm.types.utils import ModelResponse + metadata = (request_data.get("metadata") or {}) if request_data else {} + pii_tokens: Dict[str, str] = metadata.get("pii_tokens", {}) + remaining_chunks: List[ModelResponseStream] = [] try: async for chunk in response: if isinstance(chunk, ModelResponseStream): remaining_chunks.append(chunk) elif isinstance(chunk, bytes): - yield chunk # type: ignore[misc] + if pii_tokens: + yield self._unmask_sse_bytes_chunk(chunk, pii_tokens) # type: ignore[misc] + else: + yield chunk # type: ignore[misc] continue if not remaining_chunks: diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 3957e3a7fbb..5b691beccbf 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -531,11 +531,17 @@ class _PROXY_BatchRateLimiter(CustomLogger): # Check if this is a managed file (base64 encoded unified file ID) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + get_models_from_unified_file_id, ) # Managed files require bypassing the HTTP endpoint (which runs access-check hooks) # and calling the managed files hook directly with the user's credentials. is_managed_file = _is_base64_encoded_unified_file_id(file_id) + target_model_names = ( + get_models_from_unified_file_id(is_managed_file) + if is_managed_file + else [] + ) if is_managed_file and user_api_key_dict is not None: file_content = await self._fetch_managed_file_content( file_id=file_id, @@ -573,6 +579,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): await self._enforce_batch_file_model_access( user_api_key_dict=user_api_key_dict, file_content_as_dict=file_content_as_dict, + target_model_names=target_model_names or None, ) input_file_usage = _get_batch_job_input_file_usage( @@ -608,9 +615,13 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, file_content_as_dict: List[dict], + target_model_names: Optional[List[str]] = None, ) -> None: - """Reject the batch if the caller is not authorized for every - ``body.model`` named inside the JSONL. + """Reject the batch if the caller is not authorized for the upload target. + + For managed files, ``target_model_names`` (from the unified file id) is + the proxy alias the file was uploaded for and is used directly for auth. + For legacy/non-managed files, falls back to ``body.model`` values in the JSONL. Reuses standard auth helpers so the same model access rules the proxy enforces on `/chat/completions` apply here. @@ -627,9 +638,12 @@ class _PROXY_BatchRateLimiter(CustomLogger): from litellm.proxy.proxy_server import proxy_logging_obj from litellm.proxy.proxy_server import user_api_key_cache - models = _get_models_from_batch_input_file_content(file_content_as_dict) - if not models: - return + if target_model_names: + models = target_model_names + else: + models = _get_models_from_batch_input_file_content(file_content_as_dict) + if not models: + return team_object = None if ( @@ -660,12 +674,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): llm_model_list = llm_router.model_list if llm_router is not None else None for model in models: - # body.model may be the provider id after replace_model_in_jsonl; map to proxy model_name for auth. model_to_check = model - if llm_router is not None: - proxy_model_name = llm_router.resolve_model_name_from_model_id(model) - if proxy_model_name is not None: - model_to_check = proxy_model_name try: if team_object is not None: try: diff --git a/litellm/proxy/hooks/litellm_skills/main.py b/litellm/proxy/hooks/litellm_skills/main.py index 21e8bbbd308..77ed3493a0c 100644 --- a/litellm/proxy/hooks/litellm_skills/main.py +++ b/litellm/proxy/hooks/litellm_skills/main.py @@ -19,7 +19,7 @@ Usage: response = await litellm.acompletion( model="gpt-4o-mini", messages=[{"role": "user", "content": "Create a bouncing ball GIF"}], - container={"skills": [{"skill_id": "litellm:skill_abc123"}]}, + container={"skills": [{"skill_id": "litellm_skill_abc123"}]}, ) # Response includes file_ids for generated files """ @@ -31,6 +31,7 @@ from typing import Any, Dict, List, Optional, Union from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.litellm_proxy.skills.constants import LITELLM_SKILL_ID_PREFIX from litellm.llms.litellm_proxy.skills.prompt_injection import ( SkillPromptInjectionHandler, ) @@ -43,7 +44,7 @@ class SkillsInjectionHook(CustomLogger): Pre/Post-call hook that processes skills from container.skills parameter. Pre-call (async_pre_call_hook): - - Skills with 'litellm:' prefix are fetched from LiteLLM DB + - Skills with 'litellm_skill_' prefix are fetched from LiteLLM DB - For Anthropic models: native skills pass through, LiteLLM skills converted to tools - For non-Anthropic models: LiteLLM skills are converted to tools + execute_code tool @@ -78,7 +79,7 @@ class SkillsInjectionHook(CustomLogger): Process skills from container.skills before the LLM call. 1. Check if container.skills exists in request - 2. Separate skills by prefix (litellm: vs native) + 2. Separate skills by prefix (litellm_skill_ vs native) 3. Fetch LiteLLM skills from database 4. For Anthropic: keep native skills in container 5. For non-Anthropic: convert LiteLLM skills to tools, inject content, add execute_code @@ -108,7 +109,7 @@ class SkillsInjectionHook(CustomLogger): continue skill_id = skill.get("skill_id", "") - if skill_id.startswith("litellm_"): + if skill_id.startswith(LITELLM_SKILL_ID_PREFIX): # Fetch from LiteLLM DB db_skill = await self._fetch_skill_from_db( skill_id, @@ -287,7 +288,7 @@ class SkillsInjectionHook(CustomLogger): Fetch a skill from the LiteLLM database. Args: - skill_id: The skill ID (without 'litellm:' prefix) + skill_id: The skill ID (including the 'litellm_skill_' prefix) Returns: LiteLLM_SkillsTable or None if not found @@ -382,10 +383,10 @@ class SkillsInjectionHook(CustomLogger): has_executable_tool = False for tc in tool_calls: tool_name = tc.get("name", "") - # Execute if it's litellm_code_execution OR a skill tool (skill_xxx) + # Execute if it's litellm_code_execution OR a skill tool (litellm_skill_xxx) if ( tool_name == LiteLLMInternalTools.CODE_EXECUTION.value - or tool_name.startswith("skill_") + or tool_name.startswith(LITELLM_SKILL_ID_PREFIX) ): has_executable_tool = True break @@ -543,7 +544,7 @@ class SkillsInjectionHook(CustomLogger): result = await self._execute_code( code, skill_files, executor, generated_files ) - elif tool_name.startswith("skill_"): + elif tool_name.startswith(LITELLM_SKILL_ID_PREFIX): # Skill tool - execute the skill's code result = await self._execute_skill_tool( tool_name, tool_input, skill_files, executor, generated_files diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b622241dfa5..874e5aa1939 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -7,7 +7,7 @@ from pydantic import BaseModel from typing_extensions import TypedDict import litellm -from litellm import DualCache, ModelResponse +from litellm import DualCache, EmbeddingResponse, ModelResponse, TextCompletionResponse from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs @@ -570,7 +570,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): total_tokens = 0 - if isinstance(response_obj, ModelResponse): + if isinstance( + response_obj, (ModelResponse, EmbeddingResponse, TextCompletionResponse) + ): total_tokens = response_obj.usage.total_tokens # type: ignore # ------------ @@ -659,7 +661,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): if user_api_key_user_id is not None: total_tokens = 0 - if isinstance(response_obj, ModelResponse): + if isinstance( + response_obj, + (ModelResponse, EmbeddingResponse, TextCompletionResponse), + ): total_tokens = response_obj.usage.total_tokens # type: ignore request_count_api_key = ( @@ -692,7 +697,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): if user_api_key_team_id is not None: total_tokens = 0 - if isinstance(response_obj, ModelResponse): + if isinstance( + response_obj, + (ModelResponse, EmbeddingResponse, TextCompletionResponse), + ): total_tokens = response_obj.usage.total_tokens # type: ignore request_count_api_key = ( @@ -725,7 +733,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): if user_api_key_end_user_id is not None: total_tokens = 0 - if isinstance(response_obj, ModelResponse): + if isinstance( + response_obj, + (ModelResponse, EmbeddingResponse, TextCompletionResponse), + ): total_tokens = response_obj.usage.total_tokens # type: ignore request_count_api_key = ( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 62751fb68a4..6b70cea65a3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -39,7 +39,13 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject -from litellm.types.utils import CallTypes, ModelResponse, Usage +from litellm.types.utils import ( + CallTypes, + EmbeddingResponse, + ModelResponse, + TextCompletionResponse, + Usage, +) if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -2736,9 +2742,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Get total tokens from response total_tokens = 0 - # spot fix for /responses api - if isinstance(response_obj, ModelResponse) or isinstance( - response_obj, BaseLiteLLMOpenAIResponseObject + if isinstance( + response_obj, + ( + ModelResponse, + EmbeddingResponse, + TextCompletionResponse, + BaseLiteLLMOpenAIResponseObject, + ), ): _usage = getattr(response_obj, "usage", None) total_tokens = self._get_total_tokens_from_usage( diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index d8eb5dfee92..b6ddf2d8e07 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -27,6 +27,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.repositories.table_repositories import CacheConfigRepository from litellm.types.management_endpoints import ( CACHE_SETTINGS_FIELDS, @@ -160,8 +161,12 @@ class CacheSettingsManager: import json try: - cache_config = await CacheConfigRepository(prisma_client).table.find_unique( - where={"id": "cache_config"} + cache_config = await call_with_db_reconnect_retry( + prisma_client, + lambda: CacheConfigRepository(prisma_client).table.find_unique( + where={"id": "cache_config"} + ), + reason="init_cache_settings_in_db_lookup_failure", ) if cache_config is not None and cache_config.cache_settings: # Parse cache settings JSON diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 8f606fdf90d..eba16c077b0 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1850,11 +1850,10 @@ async def prepare_key_update_data( if "budget_duration" in non_default_values: budget_duration = non_default_values.pop("budget_duration") - if ( - budget_duration - and (isinstance(budget_duration, str)) - and len(budget_duration) > 0 - ): + if budget_duration is None: + non_default_values["budget_duration"] = None + non_default_values["budget_reset_at"] = None + elif isinstance(budget_duration, str) and len(budget_duration) > 0: from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time key_reset_at = get_budget_reset_time(budget_duration=budget_duration) @@ -2518,7 +2517,7 @@ async def update_key_fn( # noqa: PLR0915 }, ) - data_json: dict = data.model_dump(exclude_unset=True, exclude_none=True) + data_json: dict = data.model_dump(exclude_unset=True) key = data_json.pop("key") # get the row from db @@ -2588,6 +2587,17 @@ async def update_key_fn( # noqa: PLR0915 proxy_logging_obj=proxy_logging_obj, ) + if data.spend is not None: + try: + from litellm.proxy.proxy_server import _invalidate_spend_counter + + token_to_invalidate = _hash_token_if_needed(key) + await _invalidate_spend_counter( + counter_key=f"spend:key:{token_to_invalidate}" + ) + except Exception: + pass + asyncio.create_task( KeyManagementEventHooks.async_key_updated_hook( data=data, @@ -4775,6 +4785,13 @@ async def reset_key_spend_fn( proxy_logging_obj=proxy_logging_obj, ) + try: + from litellm.proxy.proxy_server import _invalidate_spend_counter + + await _invalidate_spend_counter(counter_key=f"spend:key:{hashed_api_key}") + except Exception: + pass + max_budget = updated_key.max_budget budget_reset_at = updated_key.budget_reset_at diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 435a8cae379..c894813ada4 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -357,6 +357,42 @@ class TeamMemberBudgetHandler: data_dict.pop("team_member_rpm_limit", None) data_dict.pop("team_member_tpm_limit", None) + @staticmethod + async def clear_team_member_budget_fields( + team_table: LiteLLM_TeamTable, + user_api_key_dict: "UserAPIKeyAuth", + updated_kv: dict, + explicitly_set_fields: set, + ) -> dict: + """Clear explicitly-nulled fields on the team member budget row.""" + from litellm.proxy._types import BudgetNewRequest + from litellm.proxy.management_endpoints.budget_management_endpoints import ( + update_budget, + ) + + if team_table.metadata is None: + team_table.metadata = {} + + team_member_budget_id = team_table.metadata.get("team_member_budget_id") + if team_member_budget_id is not None and isinstance(team_member_budget_id, str): + budget_request = BudgetNewRequest(budget_id=team_member_budget_id) + if "team_member_budget" in explicitly_set_fields: + budget_request.max_budget = None + if "team_member_budget_duration" in explicitly_set_fields: + budget_request.budget_duration = None + budget_request.budget_reset_at = None + if "team_member_rpm_limit" in explicitly_set_fields: + budget_request.rpm_limit = None + if "team_member_tpm_limit" in explicitly_set_fields: + budget_request.tpm_limit = None + await update_budget( + budget_obj=budget_request, + user_api_key_dict=user_api_key_dict, + ) + + TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + return updated_kv + @staticmethod async def backfill_team_member_budget_entries( team_id: str, @@ -1872,11 +1908,25 @@ async def update_team( # noqa: PLR0915 # Check budget_duration and budget_reset_at _set_budget_reset_at(data, updated_kv) - if TeamMemberBudgetHandler.should_create_budget( - team_member_budget=data.team_member_budget, - team_member_rpm_limit=data.team_member_rpm_limit, - team_member_tpm_limit=data.team_member_tpm_limit, - team_member_budget_duration=data.team_member_budget_duration, + _team_member_fields_in_request = { + field + for field in [ + "team_member_budget", + "team_member_rpm_limit", + "team_member_tpm_limit", + "team_member_budget_duration", + ] + if field in updated_kv + } + + if ( + _team_member_fields_in_request + and TeamMemberBudgetHandler.should_create_budget( + team_member_budget=data.team_member_budget, + team_member_rpm_limit=data.team_member_rpm_limit, + team_member_tpm_limit=data.team_member_tpm_limit, + team_member_budget_duration=data.team_member_budget_duration, + ) ): updated_kv = await TeamMemberBudgetHandler.upsert_team_member_budget_table( team_table=existing_team_row, @@ -1899,6 +1949,13 @@ async def update_team( # noqa: PLR0915 team_member_budget_id=_backfill_budget_id, prisma_client=prisma_client, ) + elif _team_member_fields_in_request: + updated_kv = await TeamMemberBudgetHandler.clear_team_member_budget_fields( + team_table=existing_team_row, + user_api_key_dict=user_api_key_dict, + updated_kv=updated_kv, + explicitly_set_fields=_team_member_fields_in_request, + ) else: TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) @@ -1987,6 +2044,8 @@ def _set_budget_reset_at(data: UpdateTeamRequest, updated_kv: dict) -> None: reset_at = get_budget_reset_time(budget_duration=data.budget_duration) updated_kv["budget_reset_at"] = reset_at + elif "budget_duration" in updated_kv and updated_kv["budget_duration"] is None: + updated_kv["budget_reset_at"] = None if data.budget_limits is not None and len(data.budget_limits) > 0: from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 7c3a6f19013..c7db818a07e 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -426,12 +426,7 @@ async def mistral_proxy_route( ) ## check for streaming - is_streaming_request = False - # anthropic is streaming when 'stream' = True is in the body - if request.method == "POST": - _request_body = await request.json() - if _request_body.get("stream"): - is_streaming_request = True + is_streaming_request = await is_streaming_request_fn(request) ## CREATE PASS-THROUGH endpoint_func = create_pass_through_route( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 6dd1f8548eb..9f353226dd0 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -5,7 +5,7 @@ Handles cost tracking and logging for OpenAI passthrough endpoints, specifically """ from datetime import datetime -from typing import List, Optional, Union +from typing import List, Optional, Tuple, Union from urllib.parse import urlparse import httpx @@ -18,6 +18,7 @@ from litellm.litellm_core_utils.litellm_logging import ( ) from litellm.llms.openai.openai import OpenAIConfig from litellm.llms.openai.openai import OpenAIConfig as OpenAIConfigType +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.proxy._types import PassThroughEndpointLoggingTypedDict from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import ( BasePassthroughLoggingHandler, @@ -29,6 +30,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( EndpointType, PassthroughStandardLoggingPayload, ) +from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import ImageResponse, LlmProviders, PassthroughCallTypes from litellm.utils import ModelResponse, TextCompletionResponse @@ -236,6 +238,42 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): ) return 0.0 + @staticmethod + def _build_responses_api_response_and_cost( + model: str, + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + custom_llm_provider: str, + ) -> Tuple[ResponsesAPIResponse, float]: + """Transform a Responses API raw response into a ResponsesAPIResponse + and compute its cost. + + The Responses API has a different on-the-wire shape from chat + completions (`output: [...]` instead of `choices: [...]`), so the + chat-completions `transform_response` raises KeyError 'choices' on + a Responses payload. Use the dedicated Responses-API transformer + (`OpenAIResponsesAPIConfig.transform_response_api_response`) here. + + Returns (litellm_model_response, response_cost) — symmetric with the + chat-completions branch which produces the same two values inline, + and analogous to the image branches' `_calculate_image_*_cost` helpers + (which return cost only because the image-response object is trivial + to build inline; the Responses payload needs a real transformer). + """ + responses_config = OpenAIResponsesAPIConfig() + litellm_model_response = responses_config.transform_response_api_response( + model=model, + raw_response=httpx_response, + logging_obj=logging_obj, + ) + response_cost = litellm.completion_cost( + completion_response=litellm_model_response, + model=model, + custom_llm_provider=custom_llm_provider, + call_type="responses", + ) + return litellm_model_response, response_cost + @staticmethod def openai_passthrough_handler( # noqa: PLR0915 httpx_response: httpx.Response, @@ -301,7 +339,12 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): try: response_cost = 0.0 litellm_model_response: Optional[ - Union[ModelResponse, TextCompletionResponse, ImageResponse] + Union[ + ModelResponse, + TextCompletionResponse, + ImageResponse, + ResponsesAPIResponse, + ] ] = None handler_instance = OpenAIPassthroughLoggingHandler() @@ -384,29 +427,18 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): litellm_model_response._hidden_params = {} litellm_model_response._hidden_params["response_cost"] = response_cost elif is_responses: - # Handle responses API cost calculation - provider_config = handler_instance.get_provider_config(model=model) - existing_litellm_params = kwargs.get("litellm_params", {}) or {} - litellm_model_response = provider_config.transform_response( - raw_response=httpx_response, - model_response=litellm.ModelResponse(), + # Responses-API cost tracking — see + # `_build_responses_api_response_and_cost` for why this needs + # a dedicated transformer (the chat-completions transform + # crashes on the Responses payload shape). + ( + litellm_model_response, + response_cost, + ) = OpenAIPassthroughLoggingHandler._build_responses_api_response_and_cost( model=model, - messages=request_body.get("messages", []), + httpx_response=httpx_response, logging_obj=logging_obj, - optional_params=request_body.get("optional_params", {}), - api_key="", - request_data=request_body, - encoding=litellm.encoding, - json_mode=False, - litellm_params=existing_litellm_params, - ) - - # Calculate cost using LiteLLM's cost calculator with responses call type - response_cost = litellm.completion_cost( - completion_response=litellm_model_response, - model=model, custom_llm_provider=custom_llm_provider, - call_type="responses", ) # Update kwargs with cost information diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index af1d39da020..46043d10a06 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -458,7 +458,16 @@ class PassThroughEndpointLogging: return False def _is_supported_openai_endpoint(self, url_route: str) -> bool: - """Check if the OpenAI endpoint is supported by the passthrough logging handler.""" + """Check if the OpenAI endpoint is supported by the passthrough logging handler. + + The Responses API route is included because + `openai_passthrough_handler` has a dedicated `elif is_responses:` + branch that knows how to extract usage + cost from the + Responses-API on-the-wire shape. Without including it here, the + outer dispatch filters Responses calls out before reaching the + handler — the inner branch is then unreachable and Responses + calls land in `LiteLLM_SpendLogs` with zero tokens / zero spend. + """ from .llm_provider_handlers.openai_passthrough_logging_handler import ( OpenAIPassthroughLoggingHandler, ) @@ -469,6 +478,7 @@ class PassThroughEndpointLogging: url_route ) or OpenAIPassthroughLoggingHandler.is_openai_image_editing_route(url_route) + or OpenAIPassthroughLoggingHandler.is_openai_responses_route(url_route) ) def _set_cost_per_request( diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index e4567b9f494..8c3fa952903 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -44,15 +44,19 @@ def _build_db_connection_url_params( pool_timeout: Optional[Union[int, float]], connect_timeout: Optional[Union[int, float]] = None, socket_timeout: Optional[Union[int, float]] = None, + disable_prepared_statements: bool = False, extra_params: Optional[dict] = None, ) -> dict: """Build the Prisma DATABASE_URL query params controlling connection pool behavior. `connect_timeout` / `socket_timeout` map to the Prisma URL params of the same name (https://www.prisma.io/docs/orm/overview/databases/postgresql) and are - omitted when None so Prisma's defaults apply. `extra_params` is an - untyped passthrough — keys it provides win over the named arguments above, - so it can be used to override any default we set here. + omitted when None so Prisma's defaults apply. `disable_prepared_statements` + sets `pgbouncer=true`, which makes Prisma stop using server-side prepared + statements (pgbouncer transaction-pool compatible; also sidesteps the + "cached plan must not change result type" error during rolling migrations). + `extra_params` is an untyped passthrough — keys it provides win over the + named arguments above, so it can be used to override any default we set here. """ params: dict = { "connection_limit": connection_limit, @@ -63,6 +67,8 @@ def _build_db_connection_url_params( params["connect_timeout"] = connect_timeout if socket_timeout is not None: params["socket_timeout"] = socket_timeout + if disable_prepared_statements: + params["pgbouncer"] = "true" if extra_params: params.update(extra_params) return params @@ -555,6 +561,7 @@ class ProxyInitializationHelpers: @click.command() +@click.argument("cli_args", nargs=-1) @click.option( "--host", default="0.0.0.0", help="Host for the server to listen on.", envvar="HOST" ) @@ -808,6 +815,7 @@ class ProxyInitializationHelpers: help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", ) def run_server( # noqa: PLR0915 + cli_args, host, port, api_base, @@ -854,6 +862,20 @@ def run_server( # noqa: PLR0915 use_v2_migration_resolver: bool, reload: bool, ): + if cli_args: + if cli_args == ("xai-oauth", "login"): + from litellm.llms.xai.oauth import XAIOAuthAuthenticator + + authenticator = XAIOAuthAuthenticator() + auth_data = authenticator.login() + click.echo( + f"xAI OAuth login successful. Credentials saved to {authenticator.auth_file}." + ) + if auth_data.get("expires_at"): + click.echo(f"Access token expires at {auth_data['expires_at']}.") + return + raise click.UsageError(f"Unknown command: {' '.join(cli_args)}") + if setup: from litellm.setup_wizard import run_setup_wizard @@ -947,6 +969,7 @@ def run_server( # noqa: PLR0915 db_connection_timeout: Optional[Union[int, float]] = 60 db_connect_timeout: Optional[Union[int, float]] = None db_socket_timeout: Optional[Union[int, float]] = None + db_disable_prepared_statements: bool = False db_extra_connection_params: Optional[dict] = None general_settings = {} ### GET DB TOKEN FOR IAM AUTH ### @@ -1067,6 +1090,17 @@ def run_server( # noqa: PLR0915 ) db_connect_timeout = general_settings.get("database_connect_timeout") db_socket_timeout = general_settings.get("database_socket_timeout") + _disable_prepared_statements = general_settings.get( + "database_disable_prepared_statements", False + ) + if isinstance(_disable_prepared_statements, str): + from litellm.secret_managers.main import str_to_bool + + db_disable_prepared_statements = ( + str_to_bool(_disable_prepared_statements) is True + ) + else: + db_disable_prepared_statements = bool(_disable_prepared_statements) db_extra_connection_params = general_settings.get( "database_extra_connection_params" ) @@ -1114,6 +1148,7 @@ def run_server( # noqa: PLR0915 pool_timeout=db_connection_timeout, connect_timeout=db_connect_timeout, socket_timeout=db_socket_timeout, + disable_prepared_statements=db_disable_prepared_statements, extra_params=db_extra_connection_params, ) if os.getenv("DATABASE_URL", None) is not None: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e50e58838e6..96c9cd1e8fb 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -311,7 +311,10 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.container_endpoints.endpoints import router as container_router from litellm.proxy.credential_endpoints.endpoints import router as credential_router from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup -from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler +from litellm.proxy.db.exception_handler import ( + PrismaDBExceptionHandler, + call_with_db_reconnect_retry, +) from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router @@ -2263,7 +2266,7 @@ async def _reconcile_budget_reservation_for_counter_update( ) except Exception: verbose_proxy_logger.warning( - "Failed to reconcile budget reservation after persisted spend; invalidating reserved counters and continuing", + "Failed to reconcile budget reservation after persisted spend; invalidating reserved counters and falling back to direct increment", exc_info=True, ) try: @@ -2274,6 +2277,7 @@ async def _reconcile_budget_reservation_for_counter_update( verbose_proxy_logger.exception( "Failed to invalidate reserved counters after reservation reconciliation failed" ) + return set() return reserved_counter_keys @@ -4075,6 +4079,7 @@ class ProxyConfig: premium_user=premium_user, config_file_path=config_file_path, litellm_settings=litellm_settings, + callback_specific_params=callback_settings, ) elif key == "model_group_settings": @@ -4089,6 +4094,8 @@ class ProxyConfig: verbose_proxy_logger.debug( f"litellm.post_call_rules: {litellm.post_call_rules}" ) + elif key == "max_budget": + litellm.max_budget = float(value) elif key == "max_internal_user_budget": litellm.max_internal_user_budget = float(value) # type: ignore elif key == "default_max_internal_user_budget": @@ -5982,8 +5989,12 @@ class ProxyConfig: """ try: - sso_settings = await SSOConfigRepository(prisma_client).table.find_unique( - where={"id": "sso_config"} + sso_settings = await call_with_db_reconnect_retry( + prisma_client, + lambda: SSOConfigRepository(prisma_client).table.find_unique( + where={"id": "sso_config"} + ), + reason="init_sso_settings_in_db_lookup_failure", ) if sso_settings is not None: sso_settings.sso_settings.pop("role_mappings", None) @@ -6018,9 +6029,13 @@ class ProxyConfig: ) try: - db_record = await ConfigOverridesRepository( - prisma_client - ).table.find_unique(where={"config_type": "hashicorp_vault"}) + db_record = await call_with_db_reconnect_retry( + prisma_client, + lambda: ConfigOverridesRepository(prisma_client).table.find_unique( + where={"config_type": "hashicorp_vault"} + ), + reason="init_hashicorp_vault_config_override_lookup_failure", + ) if db_record is None or db_record.config_value is None: if self._last_hashicorp_vault_config is not None: diff --git a/litellm/proxy/search_endpoints/search_tool_registry.py b/litellm/proxy/search_endpoints/search_tool_registry.py index 588d71b77f9..2ec2533211b 100644 --- a/litellm/proxy/search_endpoints/search_tool_registry.py +++ b/litellm/proxy/search_endpoints/search_tool_registry.py @@ -7,6 +7,7 @@ from typing import List, Optional from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient from litellm.repositories.table_repositories import SearchToolsRepository from litellm.types.search import SearchTool @@ -180,10 +181,12 @@ class SearchToolRegistry: List of search tool configurations """ try: - search_tools_from_db = await SearchToolsRepository( - prisma_client - ).table.find_many( - order={"created_at": "desc"}, + search_tools_from_db = await call_with_db_reconnect_retry( + prisma_client, + lambda: SearchToolsRepository(prisma_client).table.find_many( + order={"created_at": "desc"}, + ), + reason="get_all_search_tools_from_db_lookup_failure", ) search_tools: List[SearchTool] = [] diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index f651e6e5f7b..aa85be6671a 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3405,7 +3405,7 @@ async def ui_view_session_spend_logs( session_id, status, mcp_namespaced_tool_name, agent_id FROM "LiteLLM_SpendLogs" WHERE session_id = $1 - ORDER BY "startTime" ASC + ORDER BY "startTime" DESC LIMIT $2 OFFSET $3 """ result = await prisma_client.db.query_raw( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 2bba8bbd604..ebd5b5d90cd 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3253,40 +3253,49 @@ class PrismaClient: self, sql_query: str, *args ) -> Optional[dict]: """ - Execute a query with automatic fallback for PostgreSQL cached plan errors. + Execute a query, recovering once from PostgreSQL's "cached plan must not + change result type" error. - This handles the "cached plan must not change result type" error that occurs - during rolling deployments when schema changes are applied while old pods - still have cached query plans expecting the old schema. + That error surfaces during rolling deployments when a schema change + invalidates the prepared-statement plans that pooled connections still + hold. Clearing only the server-side plans with DEALLOCATE ALL makes + things worse: Prisma's query engine keeps a per-connection client-side + cache of prepared-statement names, so once the server drops a plan the + engine re-sends a name PostgreSQL no longer recognizes and the + connection breaks with `prepared statement "sN" does not exist`. With a + small pool that connection stays poisoned and every auth lookup fails. - Args: - sql_query: SQL query string to execute + Recreating the Prisma client kills the engine subprocess and drops the + server-side plans and the engine's client-side name cache together, so + the retried query is prepared fresh. We reconnect through + `attempt_db_reconnect`, which is singleflight: when a schema change + poisons every pooled connection at once, the first cached-plan error + recreates the client and the concurrent waiters reuse that single + recreate instead of racing to kill each other's fresh engine. We then + retry the identical query exactly once. - Returns: - Query result or None + The retry reuses the original query byte-for-byte. Mutating the SQL + (e.g. injecting a unique comment) would defeat PostgreSQL's plan cache, + forcing a fresh plan on every request and pegging the database CPU. - Raises: - Original exception if not a cached plan error + If the reconnect is skipped because a recent reconnect is still within + its cooldown, the retry runs against the same connection and may fail + again; the get_data backoff decorator re-runs the lookup and a later + attempt reconnects once the cooldown elapses. """ try: return await self.db.query_first(sql_query, *args) except Exception as e: - error_str = str(e) - if "cached plan must not change result type" in error_str: - # Force PostgreSQL to re-plan by invalidating the cache - # Add a unique comment to make the query different - sql_query_retry = sql_query.replace( - "SELECT", - f"SELECT /* cache_invalidated_{int(time.time() * 1000)} */", - ) - verbose_proxy_logger.warning( - "PostgreSQL cached plan error detected for token lookup, " - "retrying with fresh plan. This may occur during rolling deployments " - "when schema changes are applied." - ) - return await self.db.query_first(sql_query_retry, *args) - else: + if "cached plan must not change result type" not in str(e): raise + verbose_proxy_logger.warning( + "PostgreSQL cached plan error detected for token lookup; " + "recreating the database connection and retrying with the same " + "query. This may occur during rolling deployments when schema " + "changes are applied." + ) + await self.attempt_db_reconnect(reason="postgres_cached_plan_error") + return await self.db.query_first(sql_query, *args) @backoff.on_exception( backoff.expo, diff --git a/litellm/router.py b/litellm/router.py index d0f4e5ff44d..d1c8e227bea 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1658,6 +1658,67 @@ class Router: f"Dictionary '{fallback_dict}' must have exactly one key, but has {len(fallback_dict)} keys." ) + def _add_encrypted_content_affinity_check( + self, enable_global_affinity: bool + ) -> None: + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + def _move_before_deployment_affinity( + callback_list: List[Any], + callback_to_move: EncryptedContentAffinityCheck, + ) -> None: + if callback_to_move not in callback_list: + return + callback_list.remove(callback_to_move) + insert_index = next( + ( + idx + for idx, callback in enumerate(callback_list) + if isinstance(callback, DeploymentAffinityCheck) + ), + len(callback_list), + ) + callback_list.insert(insert_index, callback_to_move) + + if ( + enable_global_affinity + or EncryptedContentAffinityCheck.has_model_group_affinity_enabled( + self.model_group_affinity_config + ) + ): + if self.optional_callbacks is None: + self.optional_callbacks = [] + + existing_ec_callback: Optional[EncryptedContentAffinityCheck] = None + for cb in self.optional_callbacks: + if isinstance(cb, EncryptedContentAffinityCheck): + existing_ec_callback = cb + break + + if existing_ec_callback is not None: + existing_ec_callback.router = self + existing_ec_callback.enable_global_affinity = ( + existing_ec_callback.enable_global_affinity + or enable_global_affinity + ) + existing_ec_callback.model_group_affinity_config = ( + self.model_group_affinity_config or {} + ) + ec_callback = existing_ec_callback + else: + ec_callback = EncryptedContentAffinityCheck( + router=self, + enable_global_affinity=enable_global_affinity, + model_group_affinity_config=self.model_group_affinity_config, + ) + self.optional_callbacks.append(ec_callback) + litellm.logging_callback_manager.add_litellm_callback(ec_callback) + + _move_before_deployment_affinity(self.optional_callbacks, ec_callback) + _move_before_deployment_affinity(litellm.callbacks, ec_callback) + def add_optional_pre_call_checks( self, optional_pre_call_checks: Optional[OptionalPreCallChecks] ): @@ -1721,22 +1782,11 @@ class Router: # --------------------------------------------------------------------- # Encrypted content affinity # --------------------------------------------------------------------- - if "encrypted_content_affinity" in optional_pre_call_checks: - from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( - EncryptedContentAffinityCheck, + self._add_encrypted_content_affinity_check( + enable_global_affinity=( + "encrypted_content_affinity" in optional_pre_call_checks ) - - if self.optional_callbacks is None: - self.optional_callbacks = [] - - already_registered = any( - isinstance(cb, EncryptedContentAffinityCheck) - for cb in self.optional_callbacks - ) - if not already_registered: - ec_callback = EncryptedContentAffinityCheck(router=self) - self.optional_callbacks.append(ec_callback) - litellm.logging_callback_manager.add_litellm_callback(ec_callback) + ) # --------------------------------------------------------------------- # Remaining optional pre-call checks @@ -7738,6 +7788,39 @@ class Router: return hash_object.hexdigest() + @staticmethod + def _inherit_builtin_cache_pricing( + model_info: dict, backend_model: str, custom_llm_provider: Optional[str] + ) -> None: + """Fill missing cache pricing on a custom-priced deployment entry from + the backend model's built-in cost map entry, so a deployment that + only spells out ``input_cost_per_token``/``output_cost_per_token`` + does not silently bill cache_read/cache_creation at 0. + + User-specified cache fields always win; only ``None``/missing entries + are inherited. No-op when the backend model has no canonical entry. + """ + cache_fields = ( + "cache_creation_input_token_cost", + "cache_creation_input_token_cost_above_1hr", + "cache_creation_input_token_cost_above_200k_tokens", + "cache_read_input_token_cost", + "cache_read_input_token_cost_above_200k_tokens", + ) + if all(model_info.get(f) is not None for f in cache_fields): + return + try: + backend_info = litellm.get_model_info( + model=backend_model, custom_llm_provider=custom_llm_provider + ) + except Exception: + return + for field in cache_fields: + if model_info.get(field) is None: + backend_value = backend_info.get(field) + if backend_value is not None: + model_info[field] = backend_value + def _create_deployment( self, deployment_info: dict, @@ -7766,6 +7849,13 @@ class Router: if deployment.litellm_params.get(field) is not None: _model_info[field] = deployment.litellm_params[field] + if _model_info.get("input_cost_per_token") is not None: + Router._inherit_builtin_cache_pricing( + model_info=_model_info, + backend_model=deployment.litellm_params.model, + custom_llm_provider=deployment.litellm_params.custom_llm_provider, + ) + ## REGISTER MODEL INFO IN LITELLM MODEL COST MAP model_id = deployment.model_info.id if model_id is not None: @@ -8471,6 +8561,13 @@ class Router: credential_values.get("api_key") or deployment.litellm_params.api_key ) + if api_key is None: + verbose_router_logger.debug( + "Skipping pass-through credential setup for deployment model=%s, custom_llm_provider=%s; no api_key set. Providers like bedrock resolve credentials at request time.", + model, + custom_llm_provider, + ) + return passthrough_endpoint_router.set_pass_through_credentials( custom_llm_provider=custom_llm_provider, api_base=api_base, @@ -8505,6 +8602,13 @@ class Router: if field_value is not None: _model_info_dict[field] = field_value + if _model_info_dict.get("input_cost_per_token") is not None: + Router._inherit_builtin_cache_pricing( + model_info=_model_info_dict, + backend_model=deployment.litellm_params.model, + custom_llm_provider=deployment.litellm_params.custom_llm_provider, + ) + # Register custom pricing in litellm.model_cost. # Mirrors _create_deployment() logic to ensure dynamically-added deployments # (e.g., loaded from DB) also have their custom pricing registered. diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index 148b7fce0ee..d3e7e2ffa34 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -39,7 +39,12 @@ class DeploymentAffinityCheck(CustomLogger): CACHE_KEY_PREFIX = "deployment_affinity:v1" VALID_FLAGS = frozenset( - {"deployment_affinity", "responses_api_deployment_check", "session_affinity"} + { + "deployment_affinity", + "responses_api_deployment_check", + "session_affinity", + "encrypted_content_affinity", + } ) def __init__( diff --git a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py index 4ed19c5cd26..5fd2be9c6dd 100644 --- a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py @@ -37,7 +37,7 @@ Safe to enable globally: """ import time -from typing import TYPE_CHECKING, Any, List, Optional, cast +from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast import httpx @@ -64,17 +64,45 @@ class EncryptedContentAffinityCheck(CustomLogger): The ``model_id`` is decoded directly from the litellm-encoded item IDs – no caching or TTL management needed. - Wired via ``Router(optional_pre_call_checks=["encrypted_content_affinity"])``. + Wired via ``Router(optional_pre_call_checks=["encrypted_content_affinity"])`` or + per-model group ``model_group_affinity_config``. """ - def __init__(self, router: Optional["Router"] = None) -> None: + def __init__( + self, + router: Optional["Router"] = None, + enable_global_affinity: bool = True, + model_group_affinity_config: Optional[Dict[str, List[str]]] = None, + ) -> None: super().__init__() self.router = router + self.enable_global_affinity = enable_global_affinity + self.model_group_affinity_config: Dict[str, List[str]] = ( + model_group_affinity_config or {} + ) # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ + @staticmethod + def has_model_group_affinity_enabled( + model_group_affinity_config: Optional[Dict[str, List[str]]], + ) -> bool: + if not model_group_affinity_config: + return False + + return any( + "encrypted_content_affinity" in checks + for checks in model_group_affinity_config.values() + ) + + def _is_enabled_for_model_group(self, model_group: str) -> bool: + group_checks = self.model_group_affinity_config.get(model_group) + return self.enable_global_affinity or ( + group_checks is not None and "encrypted_content_affinity" in group_checks + ) + @staticmethod def _extract_model_id_from_input(request_input: Any) -> Optional[str]: """ @@ -213,6 +241,8 @@ class EncryptedContentAffinityCheck(CustomLogger): """ request_kwargs = request_kwargs or {} typed_healthy_deployments = cast(List[dict], healthy_deployments) + if not self._is_enabled_for_model_group(model): + return typed_healthy_deployments # Signal to the response post-processor that encrypted item IDs should be # encoded in the output of this request. Only set the flag when diff --git a/litellm/setup_wizard.py b/litellm/setup_wizard.py index 862ca13e7ba..2f0cb1233ae 100644 --- a/litellm/setup_wizard.py +++ b/litellm/setup_wizard.py @@ -52,11 +52,12 @@ PROVIDERS: List[Dict] = [ { "id": "anthropic", "name": "Anthropic", - "description": "Claude Opus 4.8, Opus 4.7, Opus 4.6, Sonnet 4.6, Haiku 4.5", + "description": "Claude Fable 5, Opus 4.8, Opus 4.7, Opus 4.6, Sonnet 4.6, Haiku 4.5", "env_key": "ANTHROPIC_API_KEY", "key_hint": "sk-ant-...", "test_model": "claude-haiku-4-5-20251001", "models": [ + "claude-fable-5", "claude-opus-4-8", "claude-opus-4-7", "claude-opus-4-6", diff --git a/litellm/types/caching.py b/litellm/types/caching.py index f8050b292c7..10453c74a15 100644 --- a/litellm/types/caching.py +++ b/litellm/types/caching.py @@ -118,4 +118,5 @@ class CachedEmbedding(TypedDict): index: Optional[int] object: Optional[str] model: Optional[str] + prompt_tokens: Optional[int] prompt_tokens_details: Optional[dict] diff --git a/litellm/types/router.py b/litellm/types/router.py index ef7eb05d087..ed858557a61 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -220,6 +220,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): use_in_pass_through: Optional[bool] = False use_litellm_proxy: Optional[bool] = False use_chat_completions_api: Optional[bool] = None + use_xai_oauth: Optional[bool] = Field( + default=False, + description="Use stored xAI OAuth credentials when no xAI API key is configured.", + ) model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True) merge_reasoning_content_in_choices: Optional[bool] = False model_info: Optional[Dict] = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index c3ea605dd9e..21eb0c9a173 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -197,6 +197,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): ] # OpenAI priority service tier pricing cache_read_input_token_cost_above_200k_tokens: Optional[float] cache_read_input_token_cost_above_272k_tokens: Optional[float] + cache_read_input_token_cost_above_512k_tokens: Optional[float] input_cost_per_character: Optional[float] # only for vertex ai models input_cost_per_audio_token: Optional[float] input_cost_per_token_above_128k_tokens: Optional[float] # only for vertex ai models @@ -206,6 +207,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): input_cost_per_token_above_272k_tokens: Optional[ float ] # GPT-5.4/5.4-pro: prompts >272K priced at 2x input + input_cost_per_token_above_512k_tokens: Optional[ + float + ] # MiniMax-M3: prompts >512K priced at 2x input input_cost_per_character_above_128k_tokens: Optional[ float ] # only for vertex ai models @@ -239,6 +243,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): output_cost_per_token_above_272k_tokens: Optional[ float ] # GPT-5.4/5.4-pro: prompts >272K priced at 1.5x output + output_cost_per_token_above_512k_tokens: Optional[ + float + ] # MiniMax-M3: prompts >512K priced at 2x output output_cost_per_character_above_128k_tokens: Optional[ float ] # only for vertex ai models @@ -3217,6 +3224,7 @@ all_litellm_params = ( "search_tool_name", "order", "enable_json_schema_validation", + "use_xai_oauth", ] + list(StandardCallbackDynamicParams.__annotations__.keys()) + list(CustomPricingLiteLLMParams.model_fields.keys()) diff --git a/litellm/utils.py b/litellm/utils.py index 8d9d0a409c6..03c628b195f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2887,6 +2887,61 @@ def _convert_stringified_numbers(value): return value +_BEDROCK_REGION_PREFIXES = ( + "us.", + "eu.", + "apac.", + "jp.", + "au.", + "us-gov.", + "global.", + "ap-northeast-1.", +) + +_CACHE_PRICING_FIELDS = ( + "cache_creation_input_token_cost", + "cache_creation_input_token_cost_above_1hr", + "cache_creation_input_token_cost_above_200k_tokens", + "cache_read_input_token_cost", + "cache_read_input_token_cost_above_200k_tokens", +) + + +def _resolve_builtin_model_cost_entry( + key: str, provider: str +) -> Optional[Dict[str, Any]]: + """Best-effort lookup of a built-in ``model_cost`` entry for a custom key + whose shape ``get_model_info`` cannot resolve (double provider prefixes + like ``bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region aliases). + + Returns a copy of the matching entry so the caller can inherit its defaults + (most importantly cache pricing) without mutating the shared built-in. + Returns ``None`` when no safe match exists. + """ + candidates: List[str] = [] + segments = key.split("/") + idx = 0 + while idx < len(segments) - 1 and segments[idx] in LlmProvidersSet: + idx += 1 + candidates.append("/".join(segments[idx:])) + + base = candidates[-1] if candidates else key + for region_prefix in _BEDROCK_REGION_PREFIXES: + if base.startswith(region_prefix): + candidates.append(base[len(region_prefix) :]) + + if provider: + stripped = _strip_model_name(model=base, custom_llm_provider=provider) + if stripped != base: + candidates.append(stripped) + + for candidate in candidates: + entry = litellm.model_cost.get(candidate) + if entry is not None and entry.get("litellm_provider") is not None: + return dict(entry) + return None + + def register_model(model_cost: Union[str, dict]): # noqa: PLR0915 """ Register new / Override existing models (and their pricing) to specific providers. @@ -2933,6 +2988,26 @@ def register_model(model_cost: Union[str, dict]): # noqa: PLR0915 except Exception: existing_model = {} model_cost_key = key + builtin_entry = _resolve_builtin_model_cost_entry( + key=_key_str, provider=provider + ) + if builtin_entry is not None: + for field in _CACHE_PRICING_FIELDS: + if ( + value.get(field) is None + and builtin_entry.get(field) is not None + ): + existing_model[field] = builtin_entry[field] + elif ( + value.get("cache_creation_input_token_cost") is None + and value.get("cache_read_input_token_cost") is None + ): + verbose_logger.warning( + f"register_model: model={key} not in built-in cost map and no " + "prefix/region variant matched; cache cost fields will default " + "to 0. To track cache cost, add cache_creation_input_token_cost " + "and cache_read_input_token_cost to model_info" + ) # ``get_model_info`` returns ``litellm_provider: None`` when the # provider is unknown (e.g. custom deployments registered via # ``Router.add_deployment``). Persisting that None into @@ -5775,6 +5850,7 @@ def _get_model_info_helper( # noqa: PLR0915 ] split_model = potential_model_names["split_model"] custom_llm_provider = potential_model_names["custom_llm_provider"] + model_cost_custom_llm_provider = custom_llm_provider ######################### provider_config: Optional[BaseLLMModelInfo] = None if custom_llm_provider and custom_llm_provider in LlmProvidersSet: @@ -5840,7 +5916,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5849,7 +5926,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5858,7 +5936,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5867,7 +5946,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None if _model_info is None: @@ -5876,7 +5956,8 @@ def _get_model_info_helper( # noqa: PLR0915 key = _matched_key _model_info = _get_model_info_from_model_cost(key=cast(str, key)) if not _check_provider_match( - model_info=_model_info, custom_llm_provider=custom_llm_provider + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None @@ -5884,7 +5965,6 @@ def _get_model_info_helper( # noqa: PLR0915 raise ValueError( "This model isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" ) - _input_cost_per_token: Optional[float] = _model_info.get( "input_cost_per_token" ) @@ -5936,6 +6016,9 @@ def _get_model_info_helper( # noqa: PLR0915 cache_read_input_token_cost_above_272k_tokens=_model_info.get( "cache_read_input_token_cost_above_272k_tokens", None ), + cache_read_input_token_cost_above_512k_tokens=_model_info.get( + "cache_read_input_token_cost_above_512k_tokens", None + ), cache_read_input_token_cost_flex=_model_info.get( "cache_read_input_token_cost_flex", None ), @@ -5957,6 +6040,9 @@ def _get_model_info_helper( # noqa: PLR0915 input_cost_per_token_above_272k_tokens=_model_info.get( "input_cost_per_token_above_272k_tokens", None ), + input_cost_per_token_above_512k_tokens=_model_info.get( + "input_cost_per_token_above_512k_tokens", None + ), input_cost_per_query=_model_info.get("input_cost_per_query", None), input_cost_per_second=_model_info.get("input_cost_per_second", None), input_cost_per_audio_token=_model_info.get( @@ -6012,6 +6098,9 @@ def _get_model_info_helper( # noqa: PLR0915 output_cost_per_token_above_272k_tokens=_model_info.get( "output_cost_per_token_above_272k_tokens", None ), + output_cost_per_token_above_512k_tokens=_model_info.get( + "output_cost_per_token_above_512k_tokens", None + ), output_cost_per_second=_model_info.get("output_cost_per_second", None), output_cost_per_second_1080p=_model_info.get( "output_cost_per_second_1080p", None @@ -8922,14 +9011,33 @@ class ProviderConfigManager: elif litellm.LlmProviders.HOSTED_VLLM == provider: return litellm.HostedVLLMResponsesAPIConfig() elif litellm.LlmProviders.BEDROCK_MANTLE == provider: - # Only OpenAI gpt frontier models (gpt-5.x, and future gpt-6 etc.) are - # served on the /openai/v1/responses path. gpt-oss and every non-OpenAI - # model on Mantle (nvidia, mistral, google, zai, ...) are chat-completions - # only and 400 on that path, so they fall through to None to keep the - # chat-completions emulation (see litellm/responses/main.py "config is None"). - model_lower = model.lower() if model else "" - if "openai.gpt-" in model_lower and "gpt-oss" not in model_lower: - return litellm.BedrockMantleResponsesAPIConfig() + # Mantle serves Responses on two upstream paths. A model takes the + # /openai/v1/responses path when its price-map entry declares + # use_openai_responses_path (data-driven, so a non-gpt-named frontier + # model can be onboarded by JSON alone), or, as a fallback needing no + # price-map entry, when its name matches the openai.gpt- frontier + # convention (minus gpt-oss) -- this keeps a future gpt-6 routing + # correctly before its entry loads. Any other model declared + # mode=responses takes the standard /v1/responses path. Everything + # else returns None and keeps the chat-completions emulation (see + # responses/main.py "config is None"). + if not model: + return None + model_lower = model.lower() + entry = litellm.model_cost.get(f"bedrock_mantle/{model}", {}) + on_openai_path = entry.get("use_openai_responses_path") is True + name_is_frontier = ( + "openai.gpt-" in model_lower and "gpt-oss" not in model_lower + ) + if on_openai_path or name_is_frontier: + return litellm.BedrockMantleResponsesAPIConfig(use_openai_path=True) + try: + if get_model_info(model, "bedrock_mantle").get("mode") == "responses": + return litellm.BedrockMantleResponsesAPIConfig( + use_openai_path=False + ) + except Exception: + pass return None return None diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b0ffc66d03b..f0b2432ddc8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1156,6 +1156,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1202,6 +1203,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1233,6 +1235,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1264,6 +1267,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1295,6 +1299,139 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "global.anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "us.anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.375e-05, + "cache_creation_input_token_cost_above_1hr": 2.2e-05, + "cache_read_input_token_cost": 1.1e-06, + "input_cost_per_token": 1.1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": true, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh" + }, + "eu.anthropic.claude-fable-5": { + "cache_creation_input_token_cost": 1.375e-05, + "cache_creation_input_token_cost_above_1hr": 2.2e-05, + "cache_read_input_token_cost": 1.1e-06, + "input_cost_per_token": 1.1e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1327,6 +1464,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1359,6 +1497,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1391,6 +1530,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1423,6 +1563,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1455,6 +1596,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -1485,6 +1627,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -2208,6 +2351,37 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true + }, + "azure_ai/claude-fable-5": { + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "litellm_provider": "azure_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -2237,6 +2411,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -10170,6 +10345,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -10204,6 +10380,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -10214,6 +10391,40 @@ }, "supports_output_config": true }, + "claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true, + "provider_specific_entry": { + "us": 1.1 + }, + "supports_output_config": true + }, "claude-opus-4-8": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -10238,6 +10449,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -24180,9 +24392,12 @@ "max_output_tokens": 8192 }, "minimax/MiniMax-M3": { - "input_cost_per_token": 6e-07, - "output_cost_per_token": 2.4e-06, - "cache_read_input_token_cost": 1.2e-07, + "input_cost_per_token": 3e-07, + "input_cost_per_token_above_512k_tokens": 6e-07, + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_above_512k_tokens": 2.4e-06, + "cache_read_input_token_cost": 6e-08, + "cache_read_input_token_cost_above_512k_tokens": 1.2e-07, "litellm_provider": "minimax", "mode": "chat", "supports_function_calling": true, @@ -24191,7 +24406,7 @@ "supports_reasoning": true, "supports_system_messages": true, "supports_vision": true, - "max_input_tokens": 512000, + "max_input_tokens": 1000000, "max_output_tokens": 128000 }, "mistral.devstral-2-123b": { @@ -34044,6 +34259,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -34072,6 +34288,67 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true + }, + "vertex_ai/claude-fable-5": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_max_reasoning_effort": true + }, + "vertex_ai/claude-fable-5@default": { + "cache_creation_input_token_cost": 1.25e-05, + "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_read_input_token_cost": 1e-06, + "input_cost_per_token": 1e-05, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -34101,6 +34378,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -34130,6 +34408,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_sampling_params": false, "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true, @@ -41410,6 +41689,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", + "use_openai_responses_path": true, "supported_endpoints": ["/v1/responses"], "supported_modalities": ["text", "image"], "supported_output_modalities": ["text"], @@ -41429,6 +41709,7 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", + "use_openai_responses_path": true, "supported_endpoints": ["/v1/responses"], "supported_modalities": ["text", "image"], "supported_output_modalities": ["text"], @@ -41902,5 +42183,164 @@ "source": "https://soniox.com/pricing", "supported_endpoints": ["/v1/audio/transcriptions"], "supports_audio_input": true + }, + "tensormesh/Qwen/Qwen3.5-397B-A17B-FP8": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 6e-07, + "output_cost_per_token": 3.6e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 4.5e-07, + "output_cost_per_token": 1.8e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/Qwen/Qwen3.6-27B-FP8": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 3.2e-07, + "output_cost_per_token": 3.2e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/lukealonso/GLM-5.1-NVFP4-MTP": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.4e-06, + "output_cost_per_token": 4.4e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 202752, + "max_output_tokens": 202752, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/deepseek-ai/DeepSeek-V4-Flash": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 2.8e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/moonshotai/Kimi-K2.6": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 9.6e-07, + "output_cost_per_token": 4e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/MiniMaxAI/MiniMax-M2.5": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, + "cache_read_input_token_cost": 0, + "max_input_tokens": 196608, + "max_output_tokens": 196608, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/google/gemma-4-31B-it": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 5.6e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/openai/gpt-oss-120b": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" + }, + "tensormesh/openai/gpt-oss-20b": { + "litellm_provider": "tensormesh", + "mode": "chat", + "input_cost_per_token": 7e-08, + "output_cost_per_token": 2.8e-07, + "cache_read_input_token_cost": 0, + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_reasoning": true, + "source": "https://serverless.tensormesh.ai/v1/models/openrouter" } } diff --git a/packaging/homebrew/README.md b/packaging/homebrew/README.md new file mode 100644 index 00000000000..ef441ded304 --- /dev/null +++ b/packaging/homebrew/README.md @@ -0,0 +1,27 @@ +# Homebrew formula for the `lite` CLI + +[`lite.rb`](./lite.rb) is the canonical source for the Homebrew formula that installs the thin LiteLLM CLI (`litellm[cli]`). It lives here so it is versioned with the code, but Homebrew serves formulae from a tap, so it has to be published to the `BerriAI/homebrew-litellm` tap to be installable. + +Once published, end users install with + +```shell +brew install BerriAI/litellm/lite +``` + +which gives them the `lite` command (`lite login`, `lite claude`, `lite models list`, ...) without the proxy server runtime. For the full proxy server, they keep using pip/uv with `litellm[proxy]` or the Docker image. + +## Why a tap and not homebrew-core + +The formula builds the published `litellm` sdist with the `cli` extra and resolves that extra's dependencies from PyPI at build time. homebrew-core forbids network access during `install` and would require every transitive dependency declared as a pinned `resource`, regenerated on each release. For a fast-moving CLI that tradeoff is not worth it, so this stays a tap formula. + +## Release runbook + +The formula can only point at a published artifact, so it activates with the first `litellm` release that ships the `cli` extra (added in [pyproject.toml](../../pyproject.toml)). + +1. Cut a `litellm` release whose `pyproject.toml` includes the `cli` extra and confirm it is on PyPI. +2. Fetch the sdist URL and checksum for that version: `curl -fsSL https://pypi.org/pypi/litellm//json | jq -r '.urls[] | select(.packagetype=="sdist") | "\(.url)\n\(.digests.sha256)"'` +3. Set `url` and `sha256` in `lite.rb` to those values; `version` is parsed from `url`. +4. Copy `lite.rb` into the tap repo under `Formula/lite.rb`, then run `brew install --build-from-source ./Formula/lite.rb` and `brew test lite` to verify a clean build and that `lite --help` works. +5. Commit and push to `BerriAI/homebrew-litellm`. + +Keep `lite.rb` here in sync with the tap copy so the in-repo formula stays the source of truth. diff --git a/packaging/homebrew/lite.rb b/packaging/homebrew/lite.rb new file mode 100644 index 00000000000..d0d61bb5b43 --- /dev/null +++ b/packaging/homebrew/lite.rb @@ -0,0 +1,33 @@ +# Homebrew formula for the thin LiteLLM `lite` CLI (litellm[cli]). +# +# Ships in the BerriAI/homebrew-litellm tap, not homebrew-core: it builds the +# published litellm sdist with the `cli` extra into a dedicated virtualenv and +# pulls the extra's deps from PyPI. That is the low-maintenance path for a +# fast-moving Python CLI; the resource-stanza alternative would need every +# transitive dep re-pinned with a fresh sha256 on each release. +# +# RELEASE STEP (see README.md in this directory): point `url` + `sha256` at the +# PyPI sdist of the first litellm version that ships the `cli` extra. `version` +# is parsed from `url`, and the build installs exactly that version, so the three +# stay in lockstep automatically. +class Lite < Formula + include Language::Python::Virtualenv + + desc "Thin client for the LiteLLM proxy: lite login, lite claude/codex/opencode" + homepage "https://docs.litellm.ai/docs/proxy/management_cli" + url "https://files.pythonhosted.org/packages/source/l/litellm/litellm-REPLACE_AT_RELEASE.tar.gz" + sha256 "REPLACE_AT_RELEASE" + license "MIT" + + depends_on "python@3.13" + + def install + virtualenv_create(libexec, "python3.13") + system libexec/"bin/pip", "install", "#{buildpath}[cli]" + bin.install_symlink libexec/"bin/lite" + end + + test do + assert_match "login", shell_output("#{bin}/lite --help") + end +end diff --git a/pyproject.toml b/pyproject.toml index 28e6f48dc4c..b9d76379faf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -71,6 +71,14 @@ proxy = [ "pyroscope-io>=0.8.16,<1.0; sys_platform != 'win32'", "pydantic-settings>=2.14.1,<3.0", ] +# Thin client install for the `lite` CLI on developer laptops. The CLI's heavy +# imports (fastapi, cryptography, ...) are all guarded, so it runs on the base +# SDK plus just these three; none of the server runtime in `proxy` is pulled in. +cli = [ + "rich>=13.9.4,<14.0", + "pyyaml>=6.0.3,<7.0", + "requests>=2.32.0,<3.0", +] extra_proxy = [ "prisma>=0.11.0,<1.0", "azure-identity>=1.25.2,<2.0", @@ -132,6 +140,7 @@ proxy-runtime = [ [project.scripts] litellm = "litellm:run_server" +lite = "litellm.proxy.client.cli:cli" litellm-proxy = "litellm.proxy.client.cli:cli" [dependency-groups] diff --git a/scripts/install-cli.sh b/scripts/install-cli.sh new file mode 100755 index 00000000000..d147286fcac --- /dev/null +++ b/scripts/install-cli.sh @@ -0,0 +1,128 @@ +#!/usr/bin/env bash +# LiteLLM CLI Installer (the thin `lite` client) +# Usage: curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/install-cli.sh | sh +# +# Installs only litellm[cli]: the `lite` command for authenticating to a LiteLLM +# proxy and running coding agents (lite claude / codex / opencode) through it. +# None of the proxy server runtime is pulled in. To run a proxy server instead, +# use scripts/install.sh, which installs litellm[proxy]. +# +# Needs only curl: uv is bootstrapped if missing, and uv provisions a compatible +# Python itself (honouring litellm's requires-python), downloading a managed one +# when the host has no suitable interpreter. +# +# NOTE: set -e without pipefail for POSIX sh compatibility (dash on Ubuntu/Debian +# ignores the shebang when invoked as `sh` and does not support `pipefail`). +set -eu + +# NOTE: before merging, this must stay as "litellm[cli]" to install from PyPI. +LITELLM_PACKAGE="litellm[cli]" +UV_VERSION="0.10.9" + +# ── colours ──────────────────────────────────────────────────────────────── +if [ -t 1 ]; then + BOLD='\033[1m' + GREEN='\033[38;2;78;186;101m' + GREY='\033[38;2;153;153;153m' + RESET='\033[0m' +else + BOLD='' GREEN='' GREY='' RESET='' +fi + +info() { printf "${GREY} %s${RESET}\n" "$*"; } +success() { printf "${GREEN} ✔ %s${RESET}\n" "$*"; } +header() { printf "${BOLD} %s${RESET}\n" "$*"; } +die() { printf "\n Error: %s\n\n" "$*" >&2; exit 1; } + +# ── banner ───────────────────────────────────────────────────────────────── +echo "" +cat << 'EOF' + ██╗ ██╗████████╗███████╗ + ██║ ██║╚══██╔══╝██╔════╝ + ██║ ██║ ██║ █████╗ + ██║ ██║ ██║ ██╔══╝ + ███████╗██║ ██║ ███████╗ + ╚══════╝╚═╝ ╚═╝ ╚══════╝ +EOF +printf " ${BOLD}LiteLLM CLI Installer${RESET} ${GREY}the thin 'lite' client for your proxy${RESET}\n\n" + +# ── OS detection ─────────────────────────────────────────────────────────── +OS="$(uname -s)" +ARCH="$(uname -m)" + +case "$OS" in + Darwin) PLATFORM="macOS ($ARCH)" ;; + Linux) PLATFORM="Linux ($ARCH)" ;; + *) die "Unsupported OS: $OS. LiteLLM supports macOS and Linux." ;; +esac + +info "Platform: $PLATFORM" + +# ── uv detection / install ──────────────────────────────────────────────── +UV_BIN="" +CURRENT_UV_VERSION="" +for candidate in uv "$HOME/.local/bin/uv"; do + if command -v "$candidate" >/dev/null 2>&1; then + UV_BIN="$(command -v "$candidate")" + break + elif [ -x "$candidate" ]; then + UV_BIN="$candidate" + break + fi +done + +if [ -n "$UV_BIN" ]; then + CURRENT_UV_VERSION="$("$UV_BIN" --version 2>/dev/null | awk '{print $2}' | head -1 || true)" +fi + +if [ -z "$UV_BIN" ] || [ "${CURRENT_UV_VERSION:-}" != "$UV_VERSION" ]; then + header "Installing uv…" + if [ -n "${CURRENT_UV_VERSION:-}" ]; then + info "Upgrading uv from ${CURRENT_UV_VERSION} to ${UV_VERSION}" + fi + curl -LsSf "https://astral.sh/uv/${UV_VERSION}/install.sh" | env UV_NO_MODIFY_PATH=1 sh \ + || die "uv installation failed. Try manually: curl -LsSf https://astral.sh/uv/${UV_VERSION}/install.sh | sh" + UV_BIN="$HOME/.local/bin/uv" +fi + +# ── install ──────────────────────────────────────────────────────────────── +# --python-preference system: reuse a compatible system Python when present, +# otherwise download a managed one. Either way uv honours litellm's requires-python, +# so a too-old (3.9) or too-new (3.14+) system Python is skipped, not forced. +echo "" +header "Installing litellm[cli]…" +echo "" + +"$UV_BIN" tool install --python-preference system --force "${LITELLM_PACKAGE}" \ + || die "uv tool install failed. Try manually: $UV_BIN tool install '${LITELLM_PACKAGE}'" + +# ── find the lite binary installed by uv tool ────────────────────────────── +SCRIPTS_DIR="$("$UV_BIN" tool dir --bin)" +LITE_BIN="${SCRIPTS_DIR}/lite" + +if [ ! -x "$LITE_BIN" ]; then + die "lite binary not found after install. Try: $UV_BIN tool install '${LITELLM_PACKAGE}'" +fi + +# ── success banner ───────────────────────────────────────────────────────── +echo "" +success "LiteLLM CLI installed" + +installed_ver="$("$LITE_BIN" --version 2>&1 | grep -oE '[0-9]+\.[0-9]+\.[0-9]+' | head -1 || true)" +[ -n "$installed_ver" ] && info "Version: $installed_ver" + +# ── PATH hint ────────────────────────────────────────────────────────────── +if ! command -v lite >/dev/null 2>&1; then + info "Note: add lite to your PATH: export PATH=\"\$PATH:${SCRIPTS_DIR}\"" +fi + +# ── next steps ───────────────────────────────────────────────────────────── +echo "" +header "Next steps:" +echo "" +info " export LITELLM_PROXY_URL=https://your-proxy # point at your gateway" +info " lite login # authenticate via SSO" +info " lite claude # run Claude Code through the proxy" +echo "" +info "Docs: https://docs.litellm.ai/docs/proxy/management_cli" +echo "" diff --git a/scripts/install.sh b/scripts/install.sh index c28d7da872f..06e6249c9ba 100755 --- a/scripts/install.sh +++ b/scripts/install.sh @@ -2,13 +2,13 @@ # LiteLLM Installer # Usage: curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/install.sh | sh # +# Needs only curl: uv is bootstrapped if missing, and uv provisions a compatible +# Python itself (reusing a suitable system one, else downloading a managed build). +# # NOTE: set -e without pipefail for POSIX sh compatibility (dash on Ubuntu/Debian # ignores the shebang when invoked as `sh` and does not support `pipefail`). set -eu -MIN_PYTHON_MAJOR=3 -MIN_PYTHON_MINOR=9 - # NOTE: before merging, this must stay as "litellm[proxy]" to install from PyPI. LITELLM_PACKAGE="litellm[proxy]" UV_VERSION="0.10.9" @@ -52,27 +52,6 @@ esac info "Platform: $PLATFORM" -# ── Python detection ─────────────────────────────────────────────────────── -PYTHON_BIN="" -for candidate in python3 python; do - if command -v "$candidate" >/dev/null 2>&1; then - major="$("$candidate" -c 'import sys; print(sys.version_info.major)' 2>/dev/null || true)" - minor="$("$candidate" -c 'import sys; print(sys.version_info.minor)' 2>/dev/null || true)" - if [ "${major:-0}" -ge "$MIN_PYTHON_MAJOR" ] && [ "${minor:-0}" -ge "$MIN_PYTHON_MINOR" ]; then - PYTHON_BIN="$(command -v "$candidate")" - info "Python: $("$candidate" --version 2>&1)" - break - fi - fi -done - -if [ -z "$PYTHON_BIN" ]; then - die "Python ${MIN_PYTHON_MAJOR}.${MIN_PYTHON_MINOR}+ is required but not found. - Install it from https://python.org/downloads or via your package manager: - macOS: brew install python@3 - Ubuntu: sudo apt install python3" -fi - # ── uv detection / install ──────────────────────────────────────────────── UV_BIN="" CURRENT_UV_VERSION="" @@ -105,15 +84,18 @@ echo "" header "Installing litellm[proxy]…" echo "" -"$UV_BIN" tool install --python "$PYTHON_BIN" --force "${LITELLM_PACKAGE}" \ - || die "uv tool install failed. Try manually: $UV_BIN tool install --python '$PYTHON_BIN' '${LITELLM_PACKAGE}'" +# --python-preference system: reuse a compatible system Python when present, +# otherwise download a managed one. Either way uv honours litellm's requires-python, +# so a too-old (3.9) or too-new (3.14+) system Python is skipped, not forced. +"$UV_BIN" tool install --python-preference system --force "${LITELLM_PACKAGE}" \ + || die "uv tool install failed. Try manually: $UV_BIN tool install '${LITELLM_PACKAGE}'" # ── find the litellm binary installed by uv tool ─────────────────────────── SCRIPTS_DIR="$("$UV_BIN" tool dir --bin)" LITELLM_BIN="${SCRIPTS_DIR}/litellm" if [ ! -x "$LITELLM_BIN" ]; then - die "litellm binary not found after install. Try: $UV_BIN tool install --python '$PYTHON_BIN' '${LITELLM_PACKAGE}'" + die "litellm binary not found after install. Try: $UV_BIN tool install '${LITELLM_PACKAGE}'" fi # ── success banner ───────────────────────────────────────────────────────── diff --git a/tests/llm_translation/reasoning_effort_grid/grid_spec.py b/tests/llm_translation/reasoning_effort_grid/grid_spec.py index a08013cd439..83a2c286d64 100644 --- a/tests/llm_translation/reasoning_effort_grid/grid_spec.py +++ b/tests/llm_translation/reasoning_effort_grid/grid_spec.py @@ -1,7 +1,6 @@ from dataclasses import dataclass, field from typing import Dict, FrozenSet, List, Optional, Tuple - OMIT = object() @@ -136,6 +135,13 @@ _CAPS_NONE: FrozenSet[str] = frozenset() ANTHROPIC_DIRECT_MODELS: Tuple[ModelEntry, ...] = ( + ModelEntry( + alias="claude-fable-5", + model="anthropic/claude-fable-5", + mode="adaptive", + required_env=_ANTHROPIC_REQ, + caps=_CAPS_XHIGH_MAX, + ), ModelEntry( alias="claude-opus-4-8", model="anthropic/claude-opus-4-8", @@ -168,6 +174,19 @@ ANTHROPIC_DIRECT_MODELS: Tuple[ModelEntry, ...] = ( AZURE_AI_MODELS: Tuple[ModelEntry, ...] = ( + ModelEntry( + alias="azure-claude-fable-5", + model="azure_ai/claude-fable-5", + mode="adaptive", + required_env=_AZURE_FOUNDRY_REQ, + caps=_CAPS_XHIGH_MAX, + fail_reason=( + "claude-fable-5 has no deployment on the CI Microsoft Foundry " + "resource yet; Foundry returns DeploymentNotFound until someone " + "creates the fable-5 deployment, so this cell stays loud in CI. " + "Remove this fail_reason once the deployment exists." + ), + ), ModelEntry( alias="azure-claude-opus-4-8", model="azure_ai/claude-opus-4-8", @@ -213,6 +232,20 @@ AZURE_AI_MODELS: Tuple[ModelEntry, ...] = ( VERTEX_AI_MODELS: Tuple[ModelEntry, ...] = ( + ModelEntry( + alias="vertex-claude-fable-5", + model="vertex_ai/claude-fable-5", + mode="adaptive", + extra_params=(("vertex_location", "global"),), + required_env=_VERTEX_REQ, + caps=_CAPS_XHIGH_MAX, + fail_reason=( + "claude-fable-5 availability on the CI Vertex project is not yet " + "confirmed for this brand-new release, so this cell stays loud in " + "CI until verified. Remove this fail_reason once the model is " + "confirmed available on the global Vertex endpoint." + ), + ), ModelEntry( alias="vertex-claude-opus-4-8", model="vertex_ai/claude-opus-4-8", @@ -263,6 +296,23 @@ VERTEX_AI_MODELS: Tuple[ModelEntry, ...] = ( BEDROCK_CONVERSE_MODELS: Tuple[ModelEntry, ...] = ( + ModelEntry( + alias="bedrock-claude-fable-5", + model="bedrock/converse/us.anthropic.claude-fable-5", + mode="adaptive", + extra_params=(("aws_region_name", "us-east-1"),), + required_env=_BEDROCK_REQ, + caps=_CAPS_XHIGH_MAX, + bedrock_effort_ceiling="xhigh", + unavailable_error="is not available for this account", + fail_reason=( + "claude-fable-5 on Bedrock requires the account to opt in to " + "provider data sharing (data retention mode " + "'provider_data_sharing' via the Data Retention API); the CI " + "account has not opted in yet, so this cell stays loud in CI. " + "Remove this fail_reason once the opt-in is done." + ), + ), ModelEntry( alias="bedrock-claude-opus-4-8", model="bedrock/converse/us.anthropic.claude-opus-4-8", diff --git a/tests/llm_translation/reasoning_effort_grid/test_reasoning_effort_grid.py b/tests/llm_translation/reasoning_effort_grid/test_reasoning_effort_grid.py index 551ab8459d1..a5f16f928e5 100644 --- a/tests/llm_translation/reasoning_effort_grid/test_reasoning_effort_grid.py +++ b/tests/llm_translation/reasoning_effort_grid/test_reasoning_effort_grid.py @@ -15,7 +15,6 @@ from .grid_spec import ( all_cells, ) - _PROMPT_MESSAGES: List[Dict[str, str]] = [ {"role": "user", "content": "Step by step, calculate 47 * 53. Show your work."} ] @@ -201,8 +200,8 @@ async def test_reasoning_effort_grid( def test_grid_cell_count() -> None: - assert len(_PARAMS) == 25 * 11, ( - f"expected 275 cells (25 provider x model combos x 11 efforts), " + assert len(_PARAMS) == 29 * 11, ( + f"expected 319 cells (29 provider x model combos x 11 efforts), " f"got {len(_PARAMS)}" ) diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py index 01bcb1a247a..9a69f513069 100644 --- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py +++ b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py @@ -2425,6 +2425,42 @@ class TestConvertToModelResponseObjectCompletion: assert result.choices[0].message.content == "The answer is 4." assert result.choices[0].message.reasoning_content == "2+2=4" + def test_reasoning_content_not_mirrored_into_provider_specific_fields(self): + """Mirroring reasoning_content into provider_specific_fields made + cache-replayed messages diverge from live Anthropic messages, which + only set it top-level, breaking cache key stability (issue #27337).""" + response_object = { + "id": "chatcmpl-5", + "model": "claude-sonnet-4-5", + "choices": [ + { + "finish_reason": "stop", + "index": 0, + "message": { + "content": "The answer is 4.", + "role": "assistant", + "reasoning_content": "2+2=4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "2+2=4", + "signature": "sig", + } + ], + }, + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}, + } + + result = convert_to_model_response_object( + response_object=response_object, + model_response_object=ModelResponse(), + ) + message = result.choices[0].message + assert message.reasoning_content == "2+2=4" + assert "reasoning_content" not in (message.provider_specific_fields or {}) + def test_response_none_raises(self): with pytest.raises(Exception): convert_to_model_response_object( diff --git a/tests/local_testing/test_basic_python_version.py b/tests/local_testing/test_basic_python_version.py index 8308e0d6033..e31c3953714 100644 --- a/tests/local_testing/test_basic_python_version.py +++ b/tests/local_testing/test_basic_python_version.py @@ -92,6 +92,56 @@ def test_package_dependencies(): ) +def test_cli_extra_is_a_thin_client_install(): + """The `cli` extra must install a working `lite` client without dragging in the + proxy server runtime. It therefore has to declare the CLI's real third-party + deps (rich, pyyaml, requests) and must never contain a server-only dependency + from the `proxy` extra; a leak there silently re-bloats the laptop install. + """ + import pathlib + + import litellm + from packaging.requirements import Requirement + + try: + import tomllib as tomli + except ImportError: + try: + import tomli + except ImportError: + pytest.skip("tomli/tomllib not available - skipping dependency check") + + pyproject_path = pathlib.Path(litellm.__file__).parent.parent / "pyproject.toml" + with open(pyproject_path, "rb") as f: + optional_deps = tomli.load(f)["project"]["optional-dependencies"] + + assert "cli" in optional_deps, "Expected a `cli` extra for the thin lite install" + + cli_names = {Requirement(req).name.lower() for req in optional_deps["cli"]} + + missing = {"rich", "pyyaml", "requests"} - cli_names + assert not missing, f"`cli` extra is missing deps the lite CLI imports: {missing}" + + server_only = { + "fastapi", + "uvicorn", + "gunicorn", + "granian", + "starlette", + "boto3", + "polars", + "soundfile", + "mcp", + "cryptography", + "apscheduler", + "rq", + "litellm-enterprise", + "litellm-proxy-extras", + } + leaked = cli_names & server_only + assert not leaked, f"`cli` extra leaks proxy-server deps onto laptops: {leaked}" + + import os import subprocess import time diff --git a/tests/local_testing/test_router_debug_logs.py b/tests/local_testing/test_router_debug_logs.py index f1c7e9d722e..ad807539bf2 100644 --- a/tests/local_testing/test_router_debug_logs.py +++ b/tests/local_testing/test_router_debug_logs.py @@ -82,7 +82,9 @@ def test_async_fallbacks(caplog): asyncio.run(_make_request()) captured_logs = [rec.message for rec in caplog.records] - # on circle ci the captured logs get some async task exception logs - filter them out "Task exception was never retrieved" + # on circle ci the captured logs get async cleanup noise from the gc (leaked + # task warnings, plus aiohttp "Unclosed client session"/"Unclosed connector" + # warnings from cached clients other router tests evicted) - filter it out captured_logs = [ log for log in captured_logs @@ -90,6 +92,8 @@ def test_async_fallbacks(caplog): and "Task was destroyed but it is pending" not in log and "get_available_deployment" not in log and "in the Langfuse queue" not in log + and "Unclosed client session" not in log + and "Unclosed connector" not in log ] print("\n Captured caplog records - ", captured_logs) diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index 3f0afe2a5a6..c170972d984 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -1236,3 +1236,60 @@ async def test_init_containers_api_endpoints_managed_id_without_model_id_applies assert call_kw["container_id"] == "cfile_upstream_abc" assert call_kw["file_id"] == "cfile_xyz" assert call_kw["custom_llm_provider"] == "azure" + + +def test_router_model_group_encrypted_content_affinity_callback_registration(): + from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( + DeploymentAffinityCheck, + ) + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + model_group_affinity_config = { + model_group: ["encrypted_content_affinity"], + } + router = Router( + model_list=[ + { + "model_name": model_group, + "litellm_params": { + "model": "openai/gpt-5.1-codex", + "api_key": "mock-api-key", + }, + } + ], + model_group_affinity_config=model_group_affinity_config, + num_retries=0, + ) + + try: + callbacks = router.optional_callbacks or [] + encrypted_content_callbacks = [ + cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) + ] + deployment_callback = next( + cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) + ) + assert len(encrypted_content_callbacks) == 1 + assert encrypted_content_callbacks[0].enable_global_affinity is False + assert ( + encrypted_content_callbacks[0].model_group_affinity_config + == model_group_affinity_config + ) + assert callbacks.index(encrypted_content_callbacks[0]) < callbacks.index( + deployment_callback + ) + + router._add_encrypted_content_affinity_check(enable_global_affinity=True) + + callbacks = router.optional_callbacks or [] + encrypted_content_callbacks = [ + cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) + ] + assert len(encrypted_content_callbacks) == 1 + assert encrypted_content_callbacks[0].enable_global_affinity is True + assert encrypted_content_callbacks[0].router is router + finally: + router.discard() diff --git a/tests/test_litellm/caching/test_caching.py b/tests/test_litellm/caching/test_caching.py index 02d62a19152..20614103ed2 100644 --- a/tests/test_litellm/caching/test_caching.py +++ b/tests/test_litellm/caching/test_caching.py @@ -3,6 +3,7 @@ import re from litellm.caching.caching import Cache from litellm.types.caching import LiteLLMCacheType +from litellm.types.utils import Embedding, EmbeddingResponse, Usage def test_cache_key_debug_log_does_not_include_prompt_material(caplog): @@ -41,8 +42,37 @@ def test_cache_key_debug_log_does_not_include_prompt_material(caplog): assert re.fullmatch(r"[0-9a-f]{64}", cache_key) created_cache_key_logs = [ - record.getMessage() for record in caplog.records if "Created cache key:" in record.getMessage() + record.getMessage() + for record in caplog.records + if "Created cache key:" in record.getMessage() ] assert created_cache_key_logs assert all(prompt_marker not in message for message in created_cache_key_logs) assert any(cache_key in message for message in created_cache_key_logs) + + +def _embedding_response(prompt_tokens, num_items): + return EmbeddingResponse( + model="amazon.titan-embed-image-v1", + data=[ + Embedding(embedding=[0.0], index=i, object="embedding") + for i in range(num_items) + ], + usage=Usage( + prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens + ), + ) + + +def test_get_per_item_prompt_tokens_single_item_returns_full_value(): + cache = Cache(type=LiteLLMCacheType.LOCAL) + result = _embedding_response(prompt_tokens=0, num_items=1) + assert cache._get_per_item_prompt_tokens(result, 0) == 0 + + +def test_get_per_item_prompt_tokens_distributes_with_remainder(): + cache = Cache(type=LiteLLMCacheType.LOCAL) + result = _embedding_response(prompt_tokens=10, num_items=3) + per_item = [cache._get_per_item_prompt_tokens(result, i) for i in range(3)] + assert sum(per_item) == 10 # 4 + 3 + 3 + assert per_item == [4, 3, 3] diff --git a/tests/test_litellm/caching/test_caching_handler.py b/tests/test_litellm/caching/test_caching_handler.py index 3eb949d7f29..01327529410 100644 --- a/tests/test_litellm/caching/test_caching_handler.py +++ b/tests/test_litellm/caching/test_caching_handler.py @@ -436,3 +436,123 @@ def test_convert_cached_responses_legacy_stream_path(): ) assert isinstance(result, CachedResponsesAPIStreamingIterator) + + +@pytest.mark.asyncio +async def test_embedding_cache_restores_stored_prompt_tokens_for_image_input(): + """Image-embedding cache hit restores prompt_tokens=0 from the stored value + instead of recomputing a bogus count by tokenizing the base64 input.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + # base64-like blob — token_counter over this would return a large nonzero count + image_input = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk" * 50 + + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "amazon.titan-embed-image-v1", + "prompt_tokens": 0, + "prompt_tokens_details": {"image_count": 1}, + } + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "amazon.titan-embed-image-v1", "input": image_input}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="amazon.titan-embed-image-v1", + ) + + assert cache_hit + assert response.usage is not None + assert response.usage.prompt_tokens == 0 + assert response.usage.total_tokens == 0 + assert response.usage.prompt_tokens_details.image_count == 1 + + +@pytest.mark.asyncio +async def test_embedding_cache_sums_stored_prompt_tokens_across_items(): + """A multi-item cache hit sums the stored per-item prompt_tokens back to the total.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + cached_result = [ + { + "embedding": [-0.01], + "index": 0, + "object": "embedding", + "model": "text-embedding-3-small", + "prompt_tokens": 5, + }, + { + "embedding": [-0.02], + "index": 1, + "object": "embedding", + "model": "text-embedding-3-small", + "prompt_tokens": 4, + }, + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "text-embedding-3-small", "input": ["hello world", "foo bar"]}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="text-embedding-3-small", + ) + + assert cache_hit + assert response.usage.prompt_tokens == 9 + assert response.usage.total_tokens == 9 + + +@pytest.mark.asyncio +async def test_embedding_cache_falls_back_to_token_counter_for_legacy_entries(): + """Legacy cache entries with no stored prompt_tokens still recompute via token_counter + for str inputs (backward compatibility).""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + # No prompt_tokens key — pre-fix entry + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "text-embedding-ada-002", + }, + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "text-embedding-ada-002", "input": "hello world"}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="text-embedding-ada-002", + ) + + assert cache_hit + # token_counter over "hello world" yields a nonzero count — fallback path still runs + assert response.usage.prompt_tokens > 0 diff --git a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py index 56e5a94cd49..ffa81abf86c 100644 --- a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py +++ b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py @@ -32,6 +32,40 @@ def test_initialize_from_proxy_config(): assert logger.compression_target == 789 +def test_initialize_from_proxy_config_ignores_non_dict_callback_specific_params(): + """Regression (#29590): a non-dict value under + callback_settings.compression_interception must not crash initialization. + + Forwarding callback_settings as callback_specific_params activates this + branch; without the isinstance(dict) guard a non-dict value reached + from_config_yaml(...).get(...) and raised AttributeError at proxy startup. + The value is ignored and the logger falls back to defaults. + """ + logger = CompressionInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={"compression_interception": True}, + ) + + assert logger.enabled is True + assert logger.compression_trigger == 200_000 + + +def test_initialize_from_proxy_config_honors_dict_callback_specific_params(): + """A valid dict under callback_settings.compression_interception is applied.""" + logger = CompressionInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={ + "compression_interception": { + "enabled": False, + "compression_trigger": 12345, + } + }, + ) + + assert logger.enabled is False + assert logger.compression_trigger == 12345 + + @pytest.mark.asyncio async def test_pre_call_hook_compresses_messages_and_injects_tool(monkeypatch): """Test pre-call hook compresses and stores per-call cache.""" diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py index 10951265115..c2a502b34eb 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py @@ -34,6 +34,35 @@ def test_initialize_from_proxy_config(): assert logger.search_tool_name == "my-search" +def test_initialize_from_proxy_config_ignores_non_dict_callback_specific_params(): + """Regression (#29590): a non-dict value under + callback_settings.websearch_interception must not crash initialization. + + Forwarding callback_settings as callback_specific_params activates this + branch; without the isinstance(dict) guard a non-dict value reached + from_config_yaml(...).get(...) and raised AttributeError at proxy startup. + The value is ignored and the logger falls back to defaults. + """ + logger = WebSearchInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={"websearch_interception": True}, + ) + + assert logger.search_tool_name is None + + +def test_initialize_from_proxy_config_honors_dict_callback_specific_params(): + """A valid dict under callback_settings.websearch_interception is applied.""" + logger = WebSearchInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={ + "websearch_interception": {"search_tool_name": "ws-tool"} + }, + ) + + assert logger.search_tool_name == "ws-tool" + + @pytest.mark.asyncio async def test_async_should_run_agentic_loop(): """Test that agentic loop is NOT triggered for wrong provider or missing WebSearch tool""" diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 2b47a232262..fe49b930c10 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -328,6 +328,41 @@ def test_generic_cost_per_token_gpt54_above_272k_tokens(): assert round(completion_cost, 10) == round(expected_completion, 10) +def test_generic_cost_per_token_minimax_m3_above_512k_tokens(): + """MiniMax-M3: prompts >512K input tokens priced at 2x input, output, and cache read.""" + model = "minimax/MiniMax-M3" + custom_llm_provider = "minimax" + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model_cost_map = litellm.model_cost[model] + prompt_tokens = 600000 + cached_tokens = 100000 + completion_tokens = 1000 + usage = Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens), + ) + prompt_cost, completion_cost = generic_cost_per_token( + model=model, + usage=usage, + custom_llm_provider=custom_llm_provider, + ) + expected_prompt = ( + model_cost_map["input_cost_per_token_above_512k_tokens"] + * (prompt_tokens - cached_tokens) + + model_cost_map["cache_read_input_token_cost_above_512k_tokens"] + * cached_tokens + ) + expected_completion = ( + model_cost_map["output_cost_per_token_above_512k_tokens"] * completion_tokens + ) + assert round(prompt_cost, 10) == round(expected_prompt, 10) + assert round(completion_cost, 10) == round(expected_completion, 10) + + def test_generic_cost_per_token_gpt55(): """gpt-5.5: base pricing — $5/1M input, $30/1M output, $0.50/1M cached input.""" model = "gpt-5.5" diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 3bf3b04bf14..ed2dfc9440e 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -11,6 +11,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( BedrockImageProcessor, _bedrock_converse_messages_pt, _bedrock_tools_pt, + _rename_duplicate_bedrock_document_names, _convert_to_bedrock_tool_call_invoke, _convert_to_bedrock_tool_call_result, anthropic_messages_pt, @@ -2809,6 +2810,93 @@ def test_bedrock_converse_messages_pt_document_deterministic_name(): assert name1 == name2 +def test_bedrock_converse_messages_pt_renames_duplicate_document_names(): + """ + The same document in multiple turns must not produce duplicate names; + Bedrock rejects requests with "Messages can not contain duplicate + document names". The first occurrence keeps its hash-based name and + later occurrences get a deterministic positional suffix. + """ + document_block = { + "type": "document", + "source": { + "type": "base64", + "media_type": "application/pdf", + "data": "dGVzdA==", + }, + } + messages = [ + { + "role": "user", + "content": [document_block, {"type": "text", "text": "summarize this"}], + }, + {"role": "assistant", "content": "It says test."}, + { + "role": "user", + "content": [document_block, {"type": "text", "text": "summarize again"}], + }, + ] + + result1 = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) + result2 = _bedrock_converse_messages_pt( + messages, "anthropic.claude-sonnet-4-6", "bedrock" + ) + + names1 = [ + block["document"]["name"] + for message in result1 + for block in message["content"] + if "document" in block + ] + names2 = [ + block["document"]["name"] + for message in result2 + for block in message["content"] + if "document" in block + ] + + assert len(names1) == 2 + assert len(set(names1)) == 2 + assert names1[1] == f"{names1[0]}_2" + assert names1 == names2 + + single_turn = _bedrock_converse_messages_pt( + [messages[0]], "anthropic.claude-sonnet-4-6", "bedrock" + ) + assert names1[0] == single_turn[0]["content"][0]["document"]["name"] + + +def test_rename_duplicate_bedrock_document_names_skips_organic_suffixes(): + """ + A renamed duplicate must not collide with a document whose organic name + already carries the would-be suffix (e.g. an existing ``report_2``), + regardless of whether that document appears before or after the rename. + """ + + def _contents(names): + return [ + { + "role": "user", + "content": [{"document": {"name": name}} for name in names], + } + ] + + def _names(contents): + return [block["document"]["name"] for block in contents[0]["content"]] + + organic_first = _rename_duplicate_bedrock_document_names( + _contents(["report", "report_2", "report"]) + ) + assert _names(organic_first) == ["report", "report_2", "report_3"] + + organic_last = _rename_duplicate_bedrock_document_names( + _contents(["report", "report", "report_2"]) + ) + assert _names(organic_last) == ["report", "report_3", "report_2"] + + def test_bedrock_converse_messages_pt_document_rejects_url_source(): """Test that a URL-type document source raises a clear error instead of KeyError.""" messages = [ diff --git a/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py b/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py new file mode 100644 index 00000000000..03790b220eb --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py @@ -0,0 +1,81 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm import LlmProviders +from litellm.litellm_core_utils.get_litellm_params import get_litellm_params +from litellm.litellm_core_utils.get_llm_provider_logic import ( + _get_openai_compatible_provider_info, +) +from litellm.llms.xai.chat.transformation import XAIChatConfig +from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import ( + ProviderConfigManager, + get_optional_params, + validate_environment, +) + + +def test_xai_provider_config_routing(): + chat_config = ProviderConfigManager.get_provider_chat_config( + model="grok-3-mini", + provider=LlmProviders.XAI, + ) + responses_config = ProviderConfigManager.get_provider_responses_api_config( + model="grok-3-mini", + provider=LlmProviders.XAI, + ) + + assert isinstance(chat_config, XAIChatConfig) + assert isinstance(responses_config, XAIResponsesAPIConfig) + + +def test_xai_openai_compatible_provider_info(): + model, custom_llm_provider, dynamic_api_key, api_base = ( + _get_openai_compatible_provider_info( + model="xai/grok-3-mini", + api_base="https://api.x.ai/v1", + api_key="api-key", + dynamic_api_key=None, + ) + ) + + assert model == "grok-3-mini" + assert custom_llm_provider == "xai" + assert api_base == "https://api.x.ai/v1" + assert dynamic_api_key == "api-key" + + +def test_xai_get_model_info_uses_xai_pricing_metadata(): + model_info = litellm.get_model_info("xai/grok-3-mini") + + assert model_info["litellm_provider"] == "xai" + assert model_info["key"] == "xai/grok-3-mini" + assert model_info["mode"] == "chat" + + +def test_xai_validate_environment_reads_api_key(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "api-key") + + result = validate_environment(model="xai/grok-3-mini") + + assert result == {"keys_in_environment": True, "missing_keys": []} + + +def test_xai_oauth_flag_is_generic_litellm_param(): + litellm_params = GenericLiteLLMParams(use_xai_oauth=True) + runtime_params = get_litellm_params(use_xai_oauth=True) + result = get_optional_params( + model="grok-3-mini", + custom_llm_provider="xai", + temperature=0.2, + drop_params=True, + ) + + assert result["temperature"] == 0.2 + assert litellm_params.use_xai_oauth is True + assert runtime_params["use_xai_oauth"] is True + assert "use_xai_oauth" not in result diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index 75038574c63..abb162e9ddb 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -5261,6 +5261,8 @@ def test_should_strip_billing_metadata_by_provider( config_cls = getattr(importlib.import_module(module_path), class_name) assert config_cls().should_strip_billing_metadata() is expected_strip + + def test_namespace_tool_flat_nested_tools_are_extracted(): """Codex sends nested tools in flat format {type, name, description, parameters} with no 'function' wrapper. These must be normalized and mapped without raising KeyError: 'function'.""" @@ -5357,3 +5359,140 @@ def test_client_metadata_stripped_from_anthropic_request(): headers={}, ) assert "client_metadata" not in result + + +@pytest.mark.parametrize( + "model", + ["claude-fable-5", "claude-opus-4-7", "claude-opus-4-8-20260120"], +) +def test_sampling_params_dropped_for_models_that_removed_them(model): + """Fable 5 / Opus 4.7 / 4.8 reject temperature != 1 and any top_p with a + 400; with drop_params set they must be dropped, not forwarded (#30064).""" + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={"temperature": 0.5, "top_p": 0.9}, + optional_params={}, + model=model, + drop_params=True, + ) + + assert "temperature" not in result + assert "top_p" not in result + + +@pytest.mark.parametrize("params", [{"temperature": 0.5}, {"top_p": 0.9}, {"top_p": 1}]) +def test_sampling_params_raise_clean_error_without_drop_params(params, monkeypatch): + monkeypatch.setattr(litellm, "drop_params", False) + config = AnthropicConfig() + + with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): + config.map_openai_params( + non_default_params=params, + optional_params={}, + model="claude-fable-5", + drop_params=False, + ) + + +def test_temperature_1_forwarded_on_models_that_removed_sampling_params(): + """temperature=1 (the API default) is still accepted and must pass through.""" + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={"temperature": 1}, + optional_params={}, + model="claude-fable-5", + drop_params=False, + ) + + assert result["temperature"] == 1 + + +@pytest.mark.parametrize("model", ["claude-opus-4-6", "claude-sonnet-4-6"]) +def test_sampling_params_forwarded_on_models_that_accept_them(model): + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={"temperature": 0.5, "top_p": 0.9}, + optional_params={}, + model=model, + drop_params=True, + ) + + assert result["temperature"] == 0.5 + assert result["top_p"] == 0.9 + + +def test_sampling_param_gating_driven_by_model_map_flag(monkeypatch): + """The drop/raise decision must come from ``supports_sampling_params`` in + the model map, not just name matching: a flagged entry gates a model whose + name says nothing, and an explicit ``true`` overrides the name fallback.""" + monkeypatch.setitem( + litellm.model_cost, "claude-zeta-9", {"supports_sampling_params": False} + ) + monkeypatch.setitem( + litellm.model_cost, "claude-fable-5-test", {"supports_sampling_params": True} + ) + config = AnthropicConfig() + + flagged_off = config.map_openai_params( + non_default_params={"top_p": 0.9}, + optional_params={}, + model="claude-zeta-9", + drop_params=True, + ) + assert "top_p" not in flagged_off + + flagged_on = config.map_openai_params( + non_default_params={"top_p": 0.9}, + optional_params={}, + model="claude-fable-5-test", + drop_params=True, + ) + assert flagged_on["top_p"] == 0.9 + + +def test_top_k_dropped_at_transform_for_models_that_removed_it(): + """``top_k`` is a provider-specific kwarg that bypasses + ``map_openai_params``, so it must be stripped at the transform_request + boundary shared by the direct, invoke, Vertex, and Azure paths (#30064).""" + config = AnthropicConfig() + + result = config.transform_request( + model="claude-fable-5", + messages=[{"role": "user", "content": "hello"}], + optional_params={"max_tokens": 10, "top_k": 40}, + litellm_params={"drop_params": True}, + headers={}, + ) + + assert "top_k" not in result + + +def test_top_k_raises_at_transform_without_drop_params(monkeypatch): + monkeypatch.setattr(litellm, "drop_params", False) + config = AnthropicConfig() + + with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): + config.transform_request( + model="claude-fable-5", + messages=[{"role": "user", "content": "hello"}], + optional_params={"max_tokens": 10, "top_k": 40}, + litellm_params={}, + headers={}, + ) + + +def test_top_k_forwarded_at_transform_on_models_that_accept_it(): + config = AnthropicConfig() + + result = config.transform_request( + model="claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={"max_tokens": 10, "top_k": 40}, + litellm_params={"drop_params": True}, + headers={}, + ) + + assert result["top_k"] == 40 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py new file mode 100644 index 00000000000..f74c5b61300 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py @@ -0,0 +1,150 @@ +""" +Regression tests for fake-streamed providers routed through `/v1/messages`. + +A fake-streaming provider (e.g. Vertex AI Gemma `:predict`) collapses its whole +response into a single `MockResponseIterator` chunk that carries content text AND a +`finish_reason` together. `AnthropicStreamWrapper` previously dropped all content in +this case — `translate_streaming_openai_response_to_anthropic` sees the finish_reason +and emits only a `message_delta`. `_CombinedChunkSplitter` splits such chunks so the +content survives. +""" + +import asyncio +import json +from types import SimpleNamespace + +from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import ( + AnthropicStreamWrapper, + _CombinedChunkSplitter, +) +from litellm.llms.base_llm.base_model_iterator import MockResponseIterator +from litellm.types.utils import ( + Choices, + Delta, + Message, + ModelResponse, + ModelResponseStream, + StreamingChoices, + Usage, +) + + +def _build_fake_stream( + content: str, finish_reason: str = "stop" +) -> MockResponseIterator: + """Mimic a Vertex Gemma `:predict` fake stream: one collapsed chunk.""" + model_response = ModelResponse() + model_response.choices = [ + Choices( + index=0, + message=Message(role="assistant", content=content), + finish_reason=finish_reason, + ) + ] + model_response.usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15) + model_response.model = "gemma4" + return MockResponseIterator(model_response=model_response) + + +def _collect_async(wrapper: AnthropicStreamWrapper) -> str: + async def _run() -> str: + out = [] + async for raw in wrapper.async_anthropic_sse_wrapper(): + out.append(raw.decode() if isinstance(raw, bytes) else raw) + return "".join(out) + + return asyncio.run(_run()) + + +def test_fake_stream_content_reaches_anthropic_sse(): + """Content from a collapsed fake-stream chunk must be emitted as a delta.""" + wrapper = AnthropicStreamWrapper( + completion_stream=_build_fake_stream("Hello, the answer is 2."), + model="gemma4", + ) + sse = _collect_async(wrapper) + + assert "content_block_delta" in sse + assert "Hello, the answer is 2." in sse + assert "message_delta" in sse + assert "message_stop" in sse + + +def test_fake_stream_usage_preserved(): + """The finish chunk keeps usage so output_tokens is non-zero.""" + wrapper = AnthropicStreamWrapper( + completion_stream=_build_fake_stream("Two."), + model="gemma4", + ) + sse = _collect_async(wrapper) + + message_delta = next( + json.loads(line[len("data: ") :]) + for block in sse.split("\n\n") + for line in block.splitlines() + if line.startswith("data: ") and '"message_delta"' in line + ) + assert message_delta["usage"]["output_tokens"] == 5 + assert message_delta["usage"]["input_tokens"] == 10 + + +def test_splitter_passes_through_non_combined_chunks(): + """A chunk with content but no finish_reason is not split.""" + chunk = ModelResponseStream( + choices=[ + StreamingChoices( + index=0, delta=Delta(content="partial"), finish_reason=None + ) + ] + ) + chunks = list(_CombinedChunkSplitter(iter([chunk]))) + assert len(chunks) == 1 + assert chunks[0].choices[0].delta.content == "partial" + + +def test_splitter_splits_combined_chunk_into_content_then_finish(): + """A chunk with both content and finish_reason becomes two chunks.""" + chunk = ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="done"), finish_reason="stop") + ] + ) + content_chunk, finish_chunk = list(_CombinedChunkSplitter(iter([chunk]))) + + assert content_chunk.choices[0].delta.content == "done" + assert content_chunk.choices[0].finish_reason is None + + assert finish_chunk.choices[0].finish_reason == "stop" + assert finish_chunk.choices[0].delta.content is None + + +def test_is_combined_false_when_choices_empty(): + """A metadata-only chunk with no choices is never treated as combined.""" + assert _CombinedChunkSplitter._is_combined(SimpleNamespace(choices=[])) is False + + +def test_is_combined_false_when_delta_missing(): + """A finish chunk whose choice has no delta is not combined.""" + chunk = SimpleNamespace(choices=[SimpleNamespace(finish_reason="stop", delta=None)]) + assert _CombinedChunkSplitter._is_combined(chunk) is False + + +def test_split_clears_reasoning_and_thinking_on_finish_chunk(): + """When the combined delta carries reasoning/thinking, only the content + chunk keeps them — the finish chunk is cleared.""" + delta = SimpleNamespace( + content="hi", + tool_calls=None, + reasoning_content="some reasoning", + thinking_blocks=[{"type": "thinking"}], + ) + chunk = SimpleNamespace( + choices=[SimpleNamespace(finish_reason="stop", delta=delta)] + ) + + content_chunk, finish_chunk = _CombinedChunkSplitter._split(chunk) + + assert content_chunk.choices[0].delta.reasoning_content == "some reasoning" + assert content_chunk.choices[0].delta.thinking_blocks == [{"type": "thinking"}] + assert finish_chunk.choices[0].delta.reasoning_content is None + assert finish_chunk.choices[0].delta.thinking_blocks is None diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py new file mode 100644 index 00000000000..450f69fb87c --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py @@ -0,0 +1,79 @@ +""" +Tests for AnthropicResponsesStreamWrapper +(litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py) +""" + +import os +import sys + +sys.path.insert( + 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../..")) +) + +from litellm.llms.anthropic.experimental_pass_through.responses_adapters.streaming_iterator import ( + AnthropicResponsesStreamWrapper, +) + + +def _process_all(events: list) -> list: + wrapper = AnthropicResponsesStreamWrapper(responses_stream=None, model="m") + for event in events: + wrapper._process_event(event) + return list(wrapper._chunk_queue) + + +class TestProcessEventTextDeltaWithoutOutputItemAdded: + """Streams that skip response.output_item.added (e.g. LMStudio) must still + open a text block before any delta and never emit index -1.""" + + def test_process_event_synthesizes_content_block_start_before_delta(self): + chunks = _process_all( + [ + {"type": "response.output_text.delta", "item_id": "i1", "delta": "Hel"}, + {"type": "response.output_text.delta", "item_id": "i1", "delta": "lo"}, + ] + ) + assert [c["type"] for c in chunks] == [ + "content_block_start", + "content_block_delta", + "content_block_delta", + ] + assert chunks[0]["content_block"] == {"type": "text", "text": ""} + assert [c["index"] for c in chunks] == [0, 0, 0] + assert chunks[1]["delta"] == {"type": "text_delta", "text": "Hel"} + + def test_process_event_delta_without_item_id_never_yields_negative_index(self): + chunks = _process_all([{"type": "response.output_text.delta", "delta": "Hi"}]) + assert [(c["type"], c["index"]) for c in chunks] == [ + ("content_block_start", 0), + ("content_block_delta", 0), + ] + + def test_process_event_unregistered_item_id_opens_new_text_block(self): + chunks = _process_all( + [ + { + "type": "response.output_item.added", + "item": {"type": "reasoning", "id": "rs_1"}, + }, + {"type": "response.output_text.delta", "item_id": "m1", "delta": "Hi"}, + ] + ) + assert chunks[1]["type"] == "content_block_start" + assert chunks[1]["content_block"] == {"type": "text", "text": ""} + assert [c["index"] for c in chunks[1:]] == [1, 1] + + def test_process_event_registered_item_id_does_not_synthesize_start(self): + chunks = _process_all( + [ + { + "type": "response.output_item.added", + "item": {"type": "message", "id": "m1"}, + }, + {"type": "response.output_text.delta", "item_id": "m1", "delta": "Hi"}, + ] + ) + assert [(c["type"], c["index"]) for c in chunks] == [ + ("content_block_start", 0), + ("content_block_delta", 0), + ] diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index ed978113b8b..5c83f8b34f7 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -5267,3 +5267,122 @@ def test_transform_response_does_not_leak_body_on_parse_failure(): msg = str(exc_info.value) assert "secret content" not in msg assert "Error converting to valid response block" in msg + + +def test_converse_drops_sampling_params_for_models_that_removed_them(): + """Fable 5 / Opus 4.7 / 4.8 reject temperature != 1 and any top_p; with + drop_params set, converse must drop them instead of forwarding (#30064).""" + config = AmazonConverseConfig() + + result = config.map_openai_params( + non_default_params={"temperature": 0.5, "top_p": 0.9}, + optional_params={}, + model="us.anthropic.claude-fable-5", + drop_params=True, + ) + + assert "temperature" not in result + assert "topP" not in result + + +def test_converse_sampling_params_raise_without_drop_params(monkeypatch): + monkeypatch.setattr(litellm, "drop_params", False) + config = AmazonConverseConfig() + + with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): + config.map_openai_params( + non_default_params={"temperature": 0.5}, + optional_params={}, + model="global.anthropic.claude-opus-4-8-v1:0", + drop_params=False, + ) + + +def test_converse_sampling_params_forwarded_on_models_that_accept_them(): + config = AmazonConverseConfig() + + result = config.map_openai_params( + non_default_params={"temperature": 0.5, "top_p": 0.9}, + optional_params={}, + model="us.anthropic.claude-sonnet-4-6", + drop_params=True, + ) + + assert result["temperature"] == 0.5 + assert result["topP"] == 0.9 + + +def test_converse_top_k_dropped_for_models_that_removed_it(): + """``top_k`` reaches converse as a provider-specific kwarg destined for + ``additionalModelRequestFields``, bypassing ``map_openai_params``; the + transform must strip it for models that removed sampling params (#30064).""" + config = AmazonConverseConfig() + + result = config.transform_request( + model="us.anthropic.claude-fable-5", + messages=[{"role": "user", "content": "hello"}], + optional_params={"top_k": 40}, + litellm_params={"drop_params": True}, + headers={}, + ) + + assert "top_k" not in result.get("additionalModelRequestFields", {}) + + +def test_converse_top_k_raises_without_drop_params(monkeypatch): + monkeypatch.setattr(litellm, "drop_params", False) + config = AmazonConverseConfig() + + with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): + config.transform_request( + model="us.anthropic.claude-fable-5", + messages=[{"role": "user", "content": "hello"}], + optional_params={"top_k": 40}, + litellm_params={}, + headers={}, + ) + + +def test_converse_top_k_forwarded_on_models_that_accept_it(): + config = AmazonConverseConfig() + + result = config.transform_request( + model="us.anthropic.claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={"top_k": 40}, + litellm_params={"drop_params": True}, + headers={}, + ) + + assert result["additionalModelRequestFields"]["top_k"] == 40 + + +def test_converse_top_k_zero_raises_without_drop_params(monkeypatch): + """``top_k=0`` must hit the same gating as any other value; previously the + truthiness check let it silently disappear on models that removed sampling + params, diverging from the Anthropic boundary that treats ``0`` as present.""" + monkeypatch.setattr(litellm, "drop_params", False) + config = AmazonConverseConfig() + + with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): + config.transform_request( + model="us.anthropic.claude-fable-5", + messages=[{"role": "user", "content": "hello"}], + optional_params={"top_k": 0}, + litellm_params={}, + headers={}, + ) + + +def test_converse_top_k_zero_forwarded_on_models_that_accept_it(): + config = AmazonConverseConfig() + + result = config.transform_request( + model="us.anthropic.claude-sonnet-4-6", + messages=[{"role": "user", "content": "hello"}], + optional_params={"top_k": 0}, + litellm_params={"drop_params": True}, + headers={}, + ) + + assert result["additionalModelRequestFields"]["top_k"] == 0 diff --git a/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py b/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py index 0de3f833a37..6812f40829a 100644 --- a/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py +++ b/tests/test_litellm/llms/bedrock/count_tokens/test_bedrock_count_tokens_transformation.py @@ -1,10 +1,15 @@ +import base64 +import json import os import sys sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path -from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig +from litellm.llms.bedrock.count_tokens.transformation import ( + DEFAULT_ANTHROPIC_INVOKE_MODEL_MAX_TOKENS, + BedrockCountTokensConfig, +) def test_detect_input_type(): @@ -20,6 +25,71 @@ def test_detect_input_type(): assert config._detect_input_type(request_with_text) == "invokeModel" +def test_detect_input_type_anthropic_blocks_route_to_invoke_model(): + """Anthropic-shape content blocks must not go through the Converse path, + which Bedrock rejects with a 400 (and the caller then silently falls back + to the local tokenizer).""" + config = BedrockCountTokensConfig() + + request = { + "messages": [ + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Reading the file."}, + { + "type": "tool_use", + "id": "toolu_01", + "name": "read_file", + "input": {"path": "main.py"}, + }, + ], + }, + ], + } + assert config._detect_input_type(request) == "invokeModel" + + +def test_detect_input_type_converse_blocks_route_to_converse(): + """Converse-shape blocks (no "type" key) keep using the converse input.""" + config = BedrockCountTokensConfig() + + request = {"messages": [{"role": "user", "content": [{"text": "hi"}]}]} + assert config._detect_input_type(request) == "converse" + + +def test_transform_to_invoke_model_format_base64_encodes_body(): + """The CountTokens API expects invokeModel.body as a base64-encoded blob; + Anthropic Messages bodies additionally need anthropic_version/max_tokens + to pass Bedrock's InvokeModel schema validation.""" + config = BedrockCountTokensConfig() + + request = { + "model": "anthropic.claude-3-sonnet-20240229-v1:0", + "messages": [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}], + } + + result = config.transform_anthropic_to_bedrock_count_tokens(request) + + body = json.loads(base64.b64decode(result["input"]["invokeModel"]["body"])) + assert body["messages"] == request["messages"] + assert "model" not in body + assert body["anthropic_version"] == "bedrock-2023-05-31" + assert body["max_tokens"] == DEFAULT_ANTHROPIC_INVOKE_MODEL_MAX_TOKENS + + +def test_transform_to_invoke_model_format_raw_body_unchanged(): + """Non-messages bodies (e.g. Titan inputText) must not get Anthropic fields.""" + config = BedrockCountTokensConfig() + + result = config.transform_anthropic_to_bedrock_count_tokens( + {"model": "amazon.titan-text-express-v1", "inputText": "hello"} + ) + + body = json.loads(base64.b64decode(result["input"]["invokeModel"]["body"])) + assert body == {"inputText": "hello"} + + def test_transform_anthropic_to_bedrock_request(): """Test basic request transformation""" config = BedrockCountTokensConfig() diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 92b5ca7b10b..e83992b6bde 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -1,11 +1,13 @@ """ Unit tests for Amazon Bedrock Mantle Responses API configuration. -Mantle's gpt-5.5 / gpt-5.4 are served ONLY on the non-standard -`/openai/v1/responses` path. These tests lock the URL construction and -Bearer auth that make that routing work. +Mantle serves Responses on two paths: gpt frontier models on +`/openai/v1/responses` and other Responses-capable models (e.g. gpt-oss) on the +standard `/v1/responses`. These tests lock the per-model path selection in the +gate, the URL construction for both paths, and the shared Bearer auth. """ +import copy import os import sys @@ -89,6 +91,42 @@ class TestBedrockMantleResponsesURL: url = cfg.get_complete_url(api_base=None, litellm_params={}) assert url == "https://bedrock-mantle.us-east-1.api.aws/openai/v1/responses" + def test_standard_path_uses_region_from_env(self, monkeypatch): + monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-east-2") + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + cfg = BedrockMantleResponsesAPIConfig(use_openai_path=False) + url = cfg.get_complete_url(api_base=None, litellm_params={}) + assert url == "https://bedrock-mantle.us-east-2.api.aws/v1/responses" + assert "/openai/v1/responses" not in url + + def test_standard_path_normalizes_v1_base(self, monkeypatch): + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + cfg = BedrockMantleResponsesAPIConfig(use_openai_path=False) + url = cfg.get_complete_url( + api_base="https://bedrock-mantle.us-east-2.api.aws/v1", + litellm_params={}, + ) + assert url == "https://bedrock-mantle.us-east-2.api.aws/v1/responses" + assert url.count("/responses") == 1 + assert "/v1/v1/responses" not in url + + def test_standard_path_full_endpoint_base_not_doubled(self, monkeypatch): + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + cfg = BedrockMantleResponsesAPIConfig(use_openai_path=False) + url = cfg.get_complete_url( + api_base="https://bedrock-mantle.us-east-2.api.aws/v1/responses", + litellm_params={}, + ) + assert url == "https://bedrock-mantle.us-east-2.api.aws/v1/responses" + assert url.count("/responses") == 1 + + def test_default_construction_keeps_openai_path(self, monkeypatch): + monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-east-2") + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + cfg = BedrockMantleResponsesAPIConfig() + url = cfg.get_complete_url(api_base=None, litellm_params={}) + assert url == "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" + class TestBedrockMantleResponsesAuth: def test_config_api_key_takes_priority(self, monkeypatch): @@ -158,6 +196,36 @@ class TestBedrockMantleResponsesAuth: is True ) + def test_standard_path_still_uses_bearer_auth(self, monkeypatch): + monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-key") + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + cfg = BedrockMantleResponsesAPIConfig(use_openai_path=False) + headers = cfg.validate_environment( + headers={}, + model="openai.gpt-oss-120b", + litellm_params=GenericLiteLLMParams(), + ) + assert headers["Authorization"] == "Bearer env-key" + + def test_standard_path_opts_out_of_native_features(self): + cfg = BedrockMantleResponsesAPIConfig(use_openai_path=False) + assert cfg.supports_native_file_search() is False + assert cfg.supports_native_websocket() is False + + +class TestBedrockMantleResponsesRequestBody: + def test_standard_path_outbound_body_carries_bare_model(self): + cfg = BedrockMantleResponsesAPIConfig(use_openai_path=False) + body = cfg.transform_responses_api_request( + model="openai.gpt-oss-120b", + input="hello", + response_api_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert body["model"] == "openai.gpt-oss-120b" + assert "input" in body + class TestBedrockMantleResponsesRegistry: def test_registry_returns_config_for_gpt_5_5(self): @@ -168,6 +236,7 @@ class TestBedrockMantleResponsesRegistry: model="openai.gpt-5.5", ) assert isinstance(cfg, BedrockMantleResponsesAPIConfig) + assert cfg.use_openai_path is True def test_registry_returns_config_for_gpt_5_4_enum(self): from litellm.utils import ProviderConfigManager @@ -177,6 +246,7 @@ class TestBedrockMantleResponsesRegistry: model="openai.gpt-5.4", ) assert isinstance(cfg, BedrockMantleResponsesAPIConfig) + assert cfg.use_openai_path is True def test_registry_returns_none_for_gpt_oss(self): # Regression guard: gpt-oss must NOT get the native Responses config; it @@ -199,9 +269,10 @@ class TestBedrockMantleResponsesRegistry: assert cfg is None def test_registry_returns_config_for_future_frontier_model(self): - # Forward-compatibility: an unseen OpenAI gpt frontier model (e.g. gpt-6) must - # get the native Responses config without a code change. The gate allow-lists - # the openai.gpt- family (minus gpt-oss), so gpt-6 matches automatically. + # Forward-compatibility: an unseen OpenAI gpt frontier model (e.g. gpt-6), + # not yet in the price map, must get the openai-path Responses config with + # no code or JSON change. The name-convention fallback (openai.gpt- minus + # gpt-oss) catches it before any price-map entry exists. from litellm.utils import ProviderConfigManager cfg = ProviderConfigManager.get_provider_responses_api_config( @@ -209,6 +280,48 @@ class TestBedrockMantleResponsesRegistry: model="openai.gpt-6", ) assert isinstance(cfg, BedrockMantleResponsesAPIConfig) + assert cfg.use_openai_path is True + + def test_price_map_flag_routes_non_gpt_name_to_openai_path( + self, restore_model_cost + ): + # Data-driven onboarding: a frontier model whose name does NOT match the + # openai.gpt- convention can still be routed to /openai/v1/responses by + # declaring use_openai_responses_path in its price-map entry, with no code + # change. The string fallback alone could never catch this name. + from litellm.utils import ProviderConfigManager, register_model + + register_model( + { + "bedrock_mantle/somelab.frontier-x": { + "litellm_provider": "bedrock_mantle", + "mode": "responses", + "use_openai_responses_path": True, + } + } + ) + cfg = ProviderConfigManager.get_provider_responses_api_config( + provider="bedrock_mantle", + model="somelab.frontier-x", + ) + assert isinstance(cfg, BedrockMantleResponsesAPIConfig) + assert cfg.use_openai_path is True + + def test_gpt_5_5_price_map_declares_openai_responses_path(self, local_cost_map): + # The gpt-5.x entries must carry the data-driven flag so frontier routing + # does not rely on the name-string fallback alone. + assert ( + litellm.model_cost["bedrock_mantle/openai.gpt-5.5"].get( + "use_openai_responses_path" + ) + is True + ) + assert ( + litellm.model_cost["bedrock_mantle/openai.gpt-5.4"].get( + "use_openai_responses_path" + ) + is True + ) @pytest.mark.parametrize( "model", @@ -243,6 +356,129 @@ class TestBedrockMantleResponsesRegistry: ) assert cfg is None + def test_declared_responses_non_openai_routes_to_standard_path( + self, restore_model_cost + ): + # New feature: a non-OpenAI model declared mode=responses (e.g. via a + # user's proxy model_info block) must route to the STANDARD /v1/responses + # path, not the frontier /openai/v1/responses path. Fails before the + # path-aware gate exists (old gate returned None for non-gpt models). + from litellm.utils import ProviderConfigManager, register_model + + register_model( + { + "bedrock_mantle/somelab.future-model": { + "litellm_provider": "bedrock_mantle", + "mode": "responses", + } + } + ) + cfg = ProviderConfigManager.get_provider_responses_api_config( + provider="bedrock_mantle", + model="somelab.future-model", + ) + assert isinstance(cfg, BedrockMantleResponsesAPIConfig) + assert cfg.use_openai_path is False + + def test_gpt_oss_opt_in_routes_to_standard_path(self, restore_model_cost): + # When a user opts gpt-oss into native Responses via model_info mode, + # it must take the STANDARD /v1/responses path (gpt-oss Responses is on + # /v1/responses, NOT the frontier /openai/v1/responses path). + from litellm.utils import ProviderConfigManager, register_model + + register_model( + { + "bedrock_mantle/openai.gpt-oss-120b": { + "litellm_provider": "bedrock_mantle", + "mode": "responses", + } + } + ) + cfg = ProviderConfigManager.get_provider_responses_api_config( + provider="bedrock_mantle", + model="openai.gpt-oss-120b", + ) + assert isinstance(cfg, BedrockMantleResponsesAPIConfig) + assert cfg.use_openai_path is False + + def test_unmapped_model_degrades_to_none_without_crashing(self, restore_model_cost): + # A non-frontier model that is not in model_cost makes get_model_info + # raise; the gate must swallow it and return None rather than crash. + from litellm.utils import ProviderConfigManager + + litellm.model_cost.pop("bedrock_mantle/somelab.unmapped-model", None) + litellm.get_model_info.cache_clear() + cfg = ProviderConfigManager.get_provider_responses_api_config( + provider="bedrock_mantle", + model="somelab.unmapped-model", + ) + assert cfg is None + + def test_register_model_restore_undoes_existing_key_overwrite(self): + # Self-contained guard for the deepcopy requirement of restore_model_cost. + # register_model overwrites an existing key by mutating its nested dict in + # place, so the snapshot must be a deepcopy: a shallow dict() copy would + # share that nested dict and leave mode=responses after restore, making + # the final assertion fail. The in-place clear+update mirrors the fixture. + from litellm.utils import ProviderConfigManager, register_model + + snapshot = copy.deepcopy(litellm.model_cost) + litellm.get_model_info.cache_clear() + try: + register_model( + { + "bedrock_mantle/openai.gpt-oss-120b": { + "litellm_provider": "bedrock_mantle", + "mode": "responses", + } + } + ) + during = ProviderConfigManager.get_provider_responses_api_config( + provider="bedrock_mantle", model="openai.gpt-oss-120b" + ) + assert isinstance(during, BedrockMantleResponsesAPIConfig) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(snapshot) + litellm.get_model_info.cache_clear() + after = ProviderConfigManager.get_provider_responses_api_config( + provider="bedrock_mantle", model="openai.gpt-oss-120b" + ) + assert after is None + + +@pytest.fixture +def restore_model_cost(): + """Snapshot litellm.model_cost so register_model edits don't leak across tests. + + register_model mutates the global litellm.model_cost, and get_model_info is + lru_cached, so without restore + cache_clear a registered model would bleed + into sibling tests in the same process. + + Two subtleties make this fixture non-obvious: + + 1. The snapshot must be a deepcopy. register_model overwrites an existing key + via `litellm.model_cost.setdefault(key, {}).update(...)`, mutating the + nested dict in place; a shallow copy would share those nested dicts and + could not capture the pre-mutation values of an existing entry. + 2. The restore must be in place (clear + update the SAME dict object), not a + reassignment. The conftest autouse `isolate_litellm_state` fixture + snapshots `litellm.model_cost` by reference and restores that reference on + its teardown, which runs after this one. Reassigning `litellm.model_cost` + to a fresh dict here is undone when conftest reinstalls its (in-place + mutated) reference, so the registered mode would leak and poison + TestBedrockMantleResponsesPricing. Mutating the original object in place + restores the contents conftest's reference points at. + """ + original_model_cost = copy.deepcopy(litellm.model_cost) + litellm.get_model_info.cache_clear() + try: + yield + finally: + litellm.model_cost.clear() + litellm.model_cost.update(original_model_cost) + litellm.get_model_info.cache_clear() + @pytest.fixture def local_cost_map(monkeypatch): diff --git a/tests/test_litellm/llms/databricks/test_databricks_streaming_utils.py b/tests/test_litellm/llms/databricks/test_databricks_streaming_utils.py new file mode 100644 index 00000000000..5612864a841 --- /dev/null +++ b/tests/test_litellm/llms/databricks/test_databricks_streaming_utils.py @@ -0,0 +1,64 @@ +""" +Regression test for the databricks streaming chunk parser. + +OpenAI-compatible servers (e.g. Vertex AI Model Garden vLLM endpoints) send a final +usage-only chunk with an empty `choices` list when `stream_options.include_usage` is +set. `chunk_parser` previously did `choices[0]` unconditionally, raising +`IndexError` -> `MidStreamFallbackError` and crashing the stream. +""" + +from litellm.llms.databricks.streaming_utils import ModelResponseIterator + + +def test_chunk_parser_handles_empty_choices_usage_chunk(): + """A usage-only final chunk (empty choices) must not raise IndexError.""" + iterator = ModelResponseIterator(streaming_response=None, sync_stream=True) + usage_only_chunk = { + "id": "chatcmpl-x", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [], + "usage": {"prompt_tokens": 20, "completion_tokens": 8, "total_tokens": 28}, + } + + result = iterator.chunk_parser(chunk=usage_only_chunk) + + assert result["text"] == "" + assert result["is_finished"] is False + assert result["usage"] is not None + assert result["usage"]["prompt_tokens"] == 20 + assert result["usage"]["completion_tokens"] == 8 + + +def test_chunk_parser_empty_choices_without_usage(): + """An empty-choices chunk with no usage block returns usage=None, no error.""" + iterator = ModelResponseIterator(streaming_response=None, sync_stream=True) + chunk = { + "id": "chatcmpl-x", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [], + } + + result = iterator.chunk_parser(chunk=chunk) + + assert result["text"] == "" + assert result["usage"] is None + + +def test_chunk_parser_normal_content_chunk_still_works(): + """A regular content chunk is unaffected by the empty-choices guard.""" + iterator = ModelResponseIterator(streaming_response=None, sync_stream=True) + chunk = { + "id": "chatcmpl-x", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [{"index": 0, "delta": {"content": "hi"}, "finish_reason": None}], + } + + result = iterator.chunk_parser(chunk=chunk) + + assert result["text"] == "hi" diff --git a/tests/test_litellm/llms/openai/completion/test_completion_handler.py b/tests/test_litellm/llms/openai/completion/test_completion_handler.py new file mode 100644 index 00000000000..c6af96fa375 --- /dev/null +++ b/tests/test_litellm/llms/openai/completion/test_completion_handler.py @@ -0,0 +1,93 @@ +""" +Tests that client headers are forwarded to the provider on the OpenAI +text completion path. + +Regression tests for https://github.com/BerriAI/litellm/issues/27410 +""" + +import os +import sys + +import pytest +import respx +from httpx import Response + +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm +from litellm import atext_completion, text_completion + + +@pytest.fixture(autouse=True) +def setup_env(monkeypatch): + monkeypatch.setenv("OPENAI_API_KEY", "sk-test-fake-key") + + +@pytest.fixture +def mock_completions_endpoint(): + return respx.post("https://api.openai.com/v1/completions").mock( + return_value=Response( + 200, + json={ + "id": "cmpl-test123", + "object": "text_completion", + "created": 1677652288, + "model": "gpt-3.5-turbo-instruct", + "choices": [ + { + "text": "hi", + "index": 0, + "logprobs": None, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + }, + ) + ) + + +@respx.mock +def test_completion_forwards_client_headers_to_provider(mock_completions_endpoint): + text_completion( + model="gpt-3.5-turbo-instruct", + prompt="hello", + max_tokens=5, + headers={"x-mycorp-llmcall-id": "abc-123"}, + ) + + request_headers = mock_completions_endpoint.calls.last.request.headers + assert request_headers["x-mycorp-llmcall-id"] == "abc-123" + + +@respx.mock +def test_completion_forwards_extra_headers_to_provider(mock_completions_endpoint): + text_completion( + model="gpt-3.5-turbo-instruct", + prompt="hello", + max_tokens=5, + extra_headers={"x-mycorp-llmcall-id": "abc-123"}, + ) + + request_headers = mock_completions_endpoint.calls.last.request.headers + assert request_headers["x-mycorp-llmcall-id"] == "abc-123" + + +@respx.mock +async def test_acompletion_forwards_client_headers_to_provider( + mock_completions_endpoint, monkeypatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + await atext_completion( + model="gpt-3.5-turbo-instruct", + prompt="hello", + max_tokens=5, + headers={"x-mycorp-llmcall-id": "abc-123"}, + ) + + request_headers = mock_completions_endpoint.calls.last.request.headers + assert request_headers["x-mycorp-llmcall-id"] == "abc-123" diff --git a/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py b/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py index f81f1c00a7b..09248a779c5 100644 --- a/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py +++ b/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py @@ -2,8 +2,23 @@ Tests for Tensormesh provider configuration and integration. """ +import pytest + import litellm +TENSORMESH_MODELS = [ + "tensormesh/Qwen/Qwen3.5-397B-A17B-FP8", + "tensormesh/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8", + "tensormesh/Qwen/Qwen3.6-27B-FP8", + "tensormesh/lukealonso/GLM-5.1-NVFP4-MTP", + "tensormesh/deepseek-ai/DeepSeek-V4-Flash", + "tensormesh/moonshotai/Kimi-K2.6", + "tensormesh/MiniMaxAI/MiniMax-M2.5", + "tensormesh/google/gemma-4-31B-it", + "tensormesh/openai/gpt-oss-120b", + "tensormesh/openai/gpt-oss-20b", +] + class TestTensormeshProviderConfig: """Test Tensormesh provider configuration""" @@ -82,3 +97,60 @@ class TestTensormeshProviderConfig: assert len(router.model_list) == 1 assert router.model_list[0]["model_name"] == "tensormesh-chat" + + +class TestTensormeshCostMap: + """The serverless models are registered in the cost map so LiteLLM can + price requests and unblock tool-calling params on the JSON provider path.""" + + @pytest.fixture(autouse=True) + def _use_local_model_cost_map(self, monkeypatch): + original_model_cost = litellm.model_cost + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.get_model_info.cache_clear() + try: + yield + finally: + litellm.model_cost = original_model_cost + litellm.get_model_info.cache_clear() + + def test_models_registered_with_capabilities(self): + for model in TENSORMESH_MODELS: + info = litellm.get_model_info(model) + assert info["litellm_provider"] == "tensormesh" + assert info["mode"] == "chat" + assert litellm.supports_function_calling(model) is True, model + assert litellm.supports_response_schema(model) is True, model + assert litellm.model_cost[model]["supports_tool_choice"] is True, model + assert litellm.model_cost[model]["supports_prompt_caching"] is True, model + + def test_reasoning_flag_matches_expected_set(self): + reasoning_models = { + "tensormesh/deepseek-ai/DeepSeek-V4-Flash", + "tensormesh/Qwen/Qwen3.5-397B-A17B-FP8", + "tensormesh/Qwen/Qwen3.6-27B-FP8", + "tensormesh/lukealonso/GLM-5.1-NVFP4-MTP", + "tensormesh/MiniMaxAI/MiniMax-M2.5", + "tensormesh/moonshotai/Kimi-K2.6", + "tensormesh/openai/gpt-oss-120b", + "tensormesh/openai/gpt-oss-20b", + "tensormesh/google/gemma-4-31B-it", + } + for model in TENSORMESH_MODELS: + assert litellm.supports_reasoning(model) is (model in reasoning_models), model + + def test_cost_is_wired_and_cache_reads_are_free(self): + prompt_cost, completion_cost = litellm.cost_per_token( + model="tensormesh/openai/gpt-oss-120b", + prompt_tokens=1_000_000, + completion_tokens=1_000_000, + ) + assert prompt_cost == pytest.approx(0.15) + assert completion_cost == pytest.approx(0.60) + assert ( + litellm.model_cost["tensormesh/openai/gpt-oss-120b"][ + "cache_read_input_token_cost" + ] + == 0 + ) diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 74888e6cd9e..cc8b14e5514 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -86,7 +86,7 @@ class TestContextCachingEndpoints: cached_content = "cached_content_123" optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = self.context_caching.check_and_create_cache( @@ -129,7 +129,7 @@ class TestContextCachingEndpoints: mock_separate.return_value = ([], self.sample_messages) # No cached messages optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = self.context_caching.check_and_create_cache( @@ -177,7 +177,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = self.context_caching.check_and_create_cache( @@ -254,7 +254,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = self.context_caching.check_and_create_cache( @@ -324,7 +324,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute and Assert with pytest.raises(VertexAIError) as exc_info: @@ -364,7 +364,7 @@ class TestContextCachingEndpoints: cached_content = "cached_content_123" optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -404,7 +404,7 @@ class TestContextCachingEndpoints: mock_separate.return_value = ([], self.sample_messages) optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -453,7 +453,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -535,7 +535,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -606,7 +606,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute and Assert with pytest.raises(VertexAIError) as exc_info: @@ -648,7 +648,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Mock the check_cache to return existing cache so we don't make HTTP calls with patch.object( @@ -694,7 +694,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = self.context_caching.check_and_create_cache( @@ -735,7 +735,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -778,7 +778,7 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() test_project = "test_project" - test_location = "test_location" + test_location = "us-central1" # Mock the async_check_cache to return existing cache so we don't make HTTP calls with patch.object( @@ -837,7 +837,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -870,7 +870,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -908,7 +908,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -942,7 +942,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1002,7 +1002,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1072,7 +1072,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1138,7 +1138,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1205,7 +1205,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1280,7 +1280,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1336,7 +1336,7 @@ class TestContextCachingEndpoints: logging_obj=self.mock_logging, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="vertext_test_token", ) @@ -1390,7 +1390,7 @@ class TestContextCachingEndpoints: cached_content=None, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="test_token", ) @@ -1441,7 +1441,7 @@ class TestContextCachingEndpoints: cached_content=None, custom_llm_provider=custom_llm_provider, vertex_project="test_project", - vertex_location="test_location", + vertex_location="us-central1", vertex_auth_header="test_token", ) diff --git a/tests/test_litellm/llms/xai/test_xai_oauth.py b/tests/test_litellm/llms/xai/test_xai_oauth.py new file mode 100644 index 00000000000..45fa6a405f2 --- /dev/null +++ b/tests/test_litellm/llms/xai/test_xai_oauth.py @@ -0,0 +1,801 @@ +import base64 +import hashlib +import json +import os +import threading +import time +from urllib.parse import parse_qs, urlparse +from unittest.mock import MagicMock + +import httpx +import litellm +import pytest +from click.testing import CliRunner + +import litellm.llms.xai.oauth as xai_oauth_module +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.llms.xai.oauth import ( + XAI_OAUTH_CLIENT_ID, + XAI_OAUTH_SCOPE, + XAIOAuthError, + XAIOAuthAuthenticator, + XAIOAuthLoginRequiredError, +) +from litellm.llms.xai.chat.transformation import XAIChatConfig +from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import get_optional_params, validate_environment + + +def _write_auth_file(tmp_path, payload): + token_dir = tmp_path / "xai_oauth" + token_dir.mkdir() + auth_file = token_dir / "auth.json" + auth_file.write_text(json.dumps(payload)) + return token_dir, auth_file + + +def test_get_access_token_uses_fresh_local_token(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "fresh-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + assert XAIOAuthAuthenticator().get_access_token() == "fresh-token" + + +def test_get_access_token_refreshes_and_preserves_refresh_token(tmp_path, monkeypatch): + token_dir, auth_file = _write_auth_file( + tmp_path, + { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() - 1, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + def handler(request: httpx.Request) -> httpx.Response: + body = dict(item.split("=") for item in request.content.decode().split("&")) + assert body["grant_type"] == "refresh_token" + assert body["refresh_token"] == "refresh-token" + assert body["client_id"] == XAI_OAUTH_CLIENT_ID + return httpx.Response( + 200, + json={ + "access_token": "new-token", + "expires_in": 3600, + "token_type": "Bearer", + }, + ) + + client = httpx.Client(transport=httpx.MockTransport(handler)) + + assert XAIOAuthAuthenticator(http_client=client).get_access_token() == "new-token" + stored = json.loads(auth_file.read_text()) + assert stored["access_token"] == "new-token" + assert stored["refresh_token"] == "refresh-token" + + +def test_get_access_token_reuses_token_refreshed_by_parallel_request(): + expired_auth_data = { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() - 1, + } + refreshed_auth_data = { + "access_token": "already-refreshed-token", + "refresh_token": "rotated-refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() + 3600, + } + authenticator = XAIOAuthAuthenticator() + authenticator._read_auth_file = MagicMock( + side_effect=[expired_auth_data, refreshed_auth_data] + ) + authenticator._refresh_tokens = MagicMock() + + assert authenticator.get_access_token() == "already-refreshed-token" + authenticator._refresh_tokens.assert_not_called() + + +def test_get_access_token_requires_login_without_auth_file(tmp_path, monkeypatch): + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing")) + + with pytest.raises(XAIOAuthLoginRequiredError): + XAIOAuthAuthenticator().get_access_token() + + +def test_get_access_token_ignores_invalid_auth_file(tmp_path, monkeypatch): + token_dir = tmp_path / "xai_oauth" + token_dir.mkdir() + (token_dir / "auth.json").write_text("{not-json") + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + with pytest.raises(XAIOAuthLoginRequiredError): + XAIOAuthAuthenticator().get_access_token() + + +def test_refresh_failure_surfaces_oauth_error(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "token_endpoint": "https://auth.x.ai/oauth/token", + "expires_at": time.time() - 1, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + client = httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(401, text="invalid_grant", request=request) + ) + ) + + with pytest.raises(XAIOAuthError) as exc_info: + XAIOAuthAuthenticator(http_client=client).get_access_token() + + assert "401 invalid_grant" in str(exc_info.value) + + +def test_build_auth_record_requires_access_and_refresh_tokens(): + authenticator = XAIOAuthAuthenticator() + + with pytest.raises(XAIOAuthError, match="access_token"): + authenticator._build_auth_record( + {"refresh_token": "refresh-token"}, + "https://auth.x.ai/oauth/token", + ) + + with pytest.raises(XAIOAuthError, match="refresh_token"): + authenticator._build_auth_record( + {"access_token": "access-token"}, + "https://auth.x.ai/oauth/token", + ) + + +def test_build_auth_record_defaults_expiry_and_token_type(): + authenticator = XAIOAuthAuthenticator() + + auth_data = authenticator._build_auth_record( + { + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_in": "not-a-number", + }, + "https://auth.x.ai/oauth/token", + ) + + assert auth_data["token_type"] == "Bearer" + assert auth_data["expires_at"] > time.time() + + +def test_is_expired_treats_missing_or_invalid_expiry_as_expired(): + authenticator = XAIOAuthAuthenticator() + + assert authenticator._is_expired({}) is True + assert authenticator._is_expired({"expires_at": "not-a-number"}) is True + + +def test_write_auth_file_creates_private_file(tmp_path, monkeypatch): + token_dir = tmp_path / "xai_oauth" + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + authenticator = XAIOAuthAuthenticator() + old_umask = os.umask(0o022) + replace_calls = [] + real_replace = os.replace + + def assert_private_temp_file(src, dst): + replace_calls.append((src, dst)) + assert oct(os.stat(src).st_mode & 0o777) == "0o600" + with open(src) as f: + assert json.load(f)["refresh_token"] == "refresh-token" + real_replace(src, dst) + + monkeypatch.setattr(os, "replace", assert_private_temp_file) + + try: + authenticator._write_auth_file( + { + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + } + ) + finally: + os.umask(old_umask) + + stored = json.loads((token_dir / "auth.json").read_text()) + assert stored["access_token"] == "access-token" + assert replace_calls + assert oct(os.stat(token_dir).st_mode & 0o777) == "0o700" + assert oct(os.stat(token_dir / "auth.json").st_mode & 0o777) == "0o600" + + +def test_discovery_rejects_unexpected_endpoint(): + authenticator = XAIOAuthAuthenticator() + + with pytest.raises(XAIOAuthError, match="unexpected endpoint"): + authenticator._validate_xai_endpoint("https://evil.example.com/oauth/token") + + with pytest.raises(XAIOAuthError, match="unexpected endpoint"): + authenticator._validate_xai_endpoint("http://auth.x.ai/oauth/token") + + +def test_discover_returns_validated_xai_endpoints(): + def handler(request: httpx.Request) -> httpx.Response: + assert request.url == "https://auth.x.ai/.well-known/openid-configuration" + return httpx.Response( + 200, + json={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + }, + ) + + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client(transport=httpx.MockTransport(handler)) + ) + + assert authenticator._discover() == { + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + } + + +def test_discover_requires_authorization_and_token_endpoints(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport(lambda request: httpx.Response(200, json={})) + ) + ) + + with pytest.raises(XAIOAuthError, match="missing endpoints"): + authenticator._discover() + + +def test_discover_wraps_http_errors(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response( + 500, text="discovery failed", request=request + ) + ) + ) + ) + + with pytest.raises(XAIOAuthError) as exc_info: + authenticator._discover() + + assert "xAI OAuth discovery request failed: 500 discovery failed" in str( + exc_info.value + ) + + +def test_discover_wraps_invalid_json_response(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(200, text="not-json") + ) + ) + ) + + with pytest.raises(XAIOAuthError, match="discovery response was not valid JSON"): + authenticator._discover() + + +def test_refresh_discovers_token_endpoint_when_auth_file_is_legacy( + tmp_path, monkeypatch +): + token_dir, auth_file = _write_auth_file( + tmp_path, + { + "access_token": "expired-token", + "refresh_token": "refresh-token", + "expires_at": time.time() - 1, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + def handler(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response( + 200, + json={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + }, + ) + return httpx.Response( + 200, + json={ + "access_token": "discovered-token", + "refresh_token": "new-refresh-token", + "expires_in": 3600, + }, + ) + + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client(transport=httpx.MockTransport(handler)) + ) + + assert authenticator.get_access_token() == "discovered-token" + stored = json.loads(auth_file.read_text()) + assert stored["token_endpoint"] == "https://auth.x.ai/oauth/token" + + +def test_exchange_token_rejects_non_object_response(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(200, json=["not", "an", "object"]) + ) + ) + ) + + with pytest.raises(XAIOAuthError, match="was not an object"): + authenticator._exchange_token("https://auth.x.ai/oauth/token", {}) + + +def test_exchange_token_wraps_invalid_json_response(): + authenticator = XAIOAuthAuthenticator( + http_client=httpx.Client( + transport=httpx.MockTransport( + lambda request: httpx.Response(200, text="not-json") + ) + ) + ) + + with pytest.raises(XAIOAuthError, match="token response was not valid JSON"): + authenticator._exchange_token("https://auth.x.ai/oauth/token", {}) + + +def test_start_callback_server_falls_back_to_ephemeral_port(monkeypatch): + calls = [] + real_server = xai_oauth_module._CallbackServer + + class FirstPortFailsCallbackServer(real_server): + def __init__(self, server_address, handler_class): + calls.append(server_address[1]) + if server_address[1] == xai_oauth_module.XAI_OAUTH_REDIRECT_PORT: + raise OSError("port unavailable") + super().__init__(server_address, handler_class) + + monkeypatch.setattr( + xai_oauth_module, "_CallbackServer", FirstPortFailsCallbackServer + ) + + server, redirect_uri = XAIOAuthAuthenticator()._start_callback_server("state-value") + try: + assert calls == [xai_oauth_module.XAI_OAUTH_REDIRECT_PORT, 0] + assert redirect_uri.startswith("http://127.0.0.1:") + assert redirect_uri.endswith("/callback") + finally: + server.server_close() + + +def test_wait_for_callback_times_out_and_closes_server(monkeypatch): + server, _ = XAIOAuthAuthenticator()._start_callback_server("state-value") + monkeypatch.setattr(xai_oauth_module, "XAI_OAUTH_CALLBACK_TIMEOUT_SECONDS", 0) + + with pytest.raises(XAIOAuthError, match="Timed out"): + XAIOAuthAuthenticator()._wait_for_callback(server) + + +def test_callback_handler_records_success_and_rejects_state_mismatch(): + authenticator = XAIOAuthAuthenticator() + server, redirect_uri = authenticator._start_callback_server("expected-state") + thread = threading.Thread(target=server.handle_request) + thread.start() + response = httpx.get(f"{redirect_uri}?code=auth-code&state=expected-state") + thread.join(timeout=5) + + assert response.status_code == 200 + assert server.callback_result == { + "code": "auth-code", + "state": "expected-state", + "error": None, + "error_description": None, + } + + server, redirect_uri = authenticator._start_callback_server("expected-state") + thread = threading.Thread(target=server.handle_request) + thread.start() + response = httpx.get(f"{redirect_uri}?code=auth-code&state=wrong-state") + thread.join(timeout=5) + + assert response.status_code == 400 + assert server.callback_result["state"] == "wrong-state" + + +def test_login_exchanges_authorization_code_and_persists_auth_record(monkeypatch): + authenticator = XAIOAuthAuthenticator() + fake_server = MagicMock() + written_records = [] + + class FakeUUID: + def __init__(self, value): + self.hex = value + + monkeypatch.setattr( + xai_oauth_module.uuid, + "uuid4", + MagicMock(side_effect=[FakeUUID("state-value"), FakeUUID("nonce-value")]), + ) + authenticator._read_auth_file = MagicMock(return_value=None) + authenticator._discover = MagicMock( + return_value={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + } + ) + authenticator._pkce_pair = MagicMock(return_value=("verifier", "challenge")) + authenticator._start_callback_server = MagicMock( + return_value=(fake_server, "http://127.0.0.1:56121/callback") + ) + authenticator._wait_for_callback = MagicMock( + return_value={"state": "state-value", "code": "auth-code"} + ) + authenticator._exchange_token = MagicMock( + return_value={ + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_in": 3600, + } + ) + authenticator._write_auth_file = MagicMock(side_effect=written_records.append) + + auth_data = authenticator.login(no_browser=True) + + authenticator._exchange_token.assert_called_once_with( + "https://auth.x.ai/oauth/token", + { + "grant_type": "authorization_code", + "code": "auth-code", + "redirect_uri": "http://127.0.0.1:56121/callback", + "client_id": XAI_OAUTH_CLIENT_ID, + "code_verifier": "verifier", + }, + ) + assert auth_data["access_token"] == "access-token" + assert written_records == [auth_data] + + +def test_login_raises_on_callback_error_or_missing_code(monkeypatch): + authenticator = XAIOAuthAuthenticator() + + class FakeUUID: + hex = "state-value" + + monkeypatch.setattr( + xai_oauth_module.uuid, "uuid4", MagicMock(return_value=FakeUUID()) + ) + authenticator._read_auth_file = MagicMock(return_value=None) + authenticator._discover = MagicMock( + return_value={ + "authorization_endpoint": "https://auth.x.ai/oauth/authorize", + "token_endpoint": "https://auth.x.ai/oauth/token", + } + ) + authenticator._pkce_pair = MagicMock(return_value=("verifier", "challenge")) + authenticator._start_callback_server = MagicMock( + return_value=(MagicMock(), "http://127.0.0.1:56121/callback") + ) + authenticator._wait_for_callback = MagicMock( + return_value={ + "state": "state-value", + "error": "access_denied", + "error_description": "denied", + } + ) + + with pytest.raises(XAIOAuthError, match="denied"): + authenticator.login(no_browser=True) + + authenticator._wait_for_callback = MagicMock(return_value={"state": "state-value"}) + + with pytest.raises(XAIOAuthError, match="no code returned"): + authenticator.login(no_browser=True) + + +def test_pkce_pair_generates_s256_challenge(): + verifier, challenge = XAIOAuthAuthenticator()._pkce_pair() + expected = ( + base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()) + .rstrip(b"=") + .decode() + ) + + assert challenge == expected + assert "=" not in verifier + assert "=" not in challenge + + +def test_build_authorize_url_contains_xai_oauth_parameters(): + authorize_url = XAIOAuthAuthenticator()._build_authorize_url( + authorization_endpoint="https://auth.x.ai/oauth/authorize", + redirect_uri="http://127.0.0.1:56121/callback", + challenge="pkce-challenge", + state="state-value", + nonce="nonce-value", + ) + parsed = urlparse(authorize_url) + params = parse_qs(parsed.query) + + assert parsed.scheme == "https" + assert parsed.netloc == "auth.x.ai" + assert params["response_type"] == ["code"] + assert params["client_id"] == [XAI_OAUTH_CLIENT_ID] + assert params["scope"] == [XAI_OAUTH_SCOPE] + assert params["code_challenge"] == ["pkce-challenge"] + assert params["code_challenge_method"] == ["S256"] + assert params["state"] == ["state-value"] + assert params["nonce"] == ["nonce-value"] + + +def test_get_llm_provider_uses_single_xai_provider(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "api-key") + + model, provider, api_key, api_base = get_llm_provider("xai/grok-4") + + assert model == "grok-4" + assert provider == "xai" + assert api_key == "api-key" + assert api_base == "https://api.x.ai/v1" + + +def test_xai_oauth_alias_is_not_a_provider(): + with pytest.raises(Exception): + get_llm_provider("xai_oauth/grok-4") + + +def test_chat_config_wraps_flagged_oauth_errors_as_authentication_error( + tmp_path, monkeypatch +): + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing")) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + + assert exc_info.value.llm_provider == "xai" + assert "litellm xai-oauth login" in str(exc_info.value) + + +def test_chat_config_injects_flagged_oauth_token(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "chat-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + + assert headers["Authorization"] == "Bearer chat-token" + + +def test_chat_config_ignores_api_base_override_for_flagged_oauth(monkeypatch): + monkeypatch.setenv("XAI_OAUTH_API_BASE", "https://api.x.ai/v1") + + url = XAIChatConfig().get_complete_url( + api_base="https://attacker.example.com/v1", + api_key=None, + model="grok-4", + optional_params={}, + litellm_params={"use_xai_oauth": True}, + ) + + assert url == "https://api.x.ai/v1/chat/completions" + + +def test_chat_config_treats_blank_api_key_as_absent_for_flagged_oauth( + tmp_path, monkeypatch +): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "stored-oauth-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key="", + ) + + assert headers["Authorization"] == "Bearer stored-oauth-token" + + +def test_chat_config_allows_api_base_override_with_caller_api_key(): + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key="caller-api-key", + ) + url = XAIChatConfig().get_complete_url( + api_base="https://custom.example.com/v1", + api_key="caller-api-key", + model="grok-4", + optional_params={}, + litellm_params={"use_xai_oauth": True}, + ) + + assert headers["Authorization"] == "Bearer caller-api-key" + assert url == "https://custom.example.com/v1/chat/completions" + + +def test_chat_config_prioritizes_env_api_key_over_oauth_flag(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "env-api-key") + + headers = XAIChatConfig().validate_environment( + headers={}, + model="grok-4", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + url = XAIChatConfig().get_complete_url( + api_base="https://custom.example.com/v1", + api_key=None, + model="grok-4", + optional_params={}, + litellm_params={"use_xai_oauth": True}, + ) + + assert headers["Authorization"] == "Bearer env-api-key" + assert url == "https://custom.example.com/v1/chat/completions" + + +def test_validate_environment_still_reports_xai_api_key(monkeypatch): + monkeypatch.setenv("XAI_API_KEY", "env-api-key") + + assert validate_environment("xai/grok-4") == { + "keys_in_environment": True, + "missing_keys": [], + } + + +def test_xai_oauth_flag_uses_xai_optional_param_mapping(): + litellm_params = GenericLiteLLMParams(use_xai_oauth=True) + optional_params = get_optional_params( + model="grok-4", + custom_llm_provider="xai", + temperature=0.2, + max_tokens=8, + ) + + assert optional_params["temperature"] == 0.2 + assert optional_params["max_tokens"] == 8 + assert litellm_params.use_xai_oauth is True + assert "use_xai_oauth" not in optional_params + + +def test_responses_config_injects_flagged_oauth_bearer_token(tmp_path, monkeypatch): + token_dir, _ = _write_auth_file( + tmp_path, + { + "access_token": "responses-token", + "refresh_token": "refresh-token", + "expires_at": time.time() + 3600, + }, + ) + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(token_dir)) + + headers = XAIResponsesAPIConfig().validate_environment( + headers={}, + model="grok-4", + litellm_params=GenericLiteLLMParams(use_xai_oauth=True), + ) + + assert headers["Authorization"] == "Bearer responses-token" + + +def test_responses_config_endpoint_url_uses_oauth_authenticator(monkeypatch): + monkeypatch.setenv("XAI_OAUTH_API_BASE", "https://xai.example.com/v1/") + config = XAIResponsesAPIConfig() + + assert config.get_complete_url( + api_base=None, litellm_params={"use_xai_oauth": True} + ) == ("https://xai.example.com/v1/responses") + assert ( + config.get_complete_url( + api_base="https://custom.example.com/v1/", + litellm_params={"use_xai_oauth": True}, + ) + == "https://xai.example.com/v1/responses" + ) + assert ( + config.get_complete_url( + api_base="https://custom.example.com/v1/", + litellm_params={"api_key": "", "use_xai_oauth": True}, + ) + == "https://xai.example.com/v1/responses" + ) + assert ( + config.get_complete_url( + api_base="https://custom.example.com/v1/", + litellm_params={"api_key": "caller-api-key"}, + ) + == "https://custom.example.com/v1/responses" + ) + + +def test_responses_config_wraps_flagged_oauth_errors_as_authentication_error( + tmp_path, monkeypatch +): + monkeypatch.setenv("XAI_OAUTH_TOKEN_DIR", str(tmp_path / "missing")) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + XAIResponsesAPIConfig().validate_environment( + headers={}, + model="grok-4", + litellm_params=GenericLiteLLMParams(use_xai_oauth=True), + ) + + assert XAIResponsesAPIConfig().custom_llm_provider.value == "xai" + assert exc_info.value.llm_provider == "xai" + + +def test_proxy_cli_xai_oauth_login_uses_single_authenticator(monkeypatch): + from litellm.proxy.proxy_cli import run_server + + instances = [] + + class FakeAuthenticator: + auth_file = "/tmp/xai-oauth-auth.json" + + def __init__(self): + instances.append(self) + + def login(self): + return {"expires_at": 1234567890} + + monkeypatch.setattr( + "litellm.llms.xai.oauth.XAIOAuthAuthenticator", FakeAuthenticator + ) + + result = CliRunner().invoke(run_server, ["xai-oauth", "login"]) + + assert result.exit_code == 0 + assert len(instances) == 1 + assert "Credentials saved to /tmp/xai-oauth-auth.json" in result.output + assert "Access token expires at 1234567890" in result.output diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 19065ff816b..a846ca24739 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -872,6 +872,7 @@ def _mock_env_vars_prisma(row=None): prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[]) prisma.db.litellm_mcpuserenvvars.upsert = AsyncMock() prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock() + prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock() return prisma @@ -1252,6 +1253,64 @@ async def test_delete_mcp_server_succeeds_when_orphan_cleanup_fails(): prisma.db.litellm_mcpuserenvvars.delete_many.assert_awaited_once() +@pytest.mark.asyncio +async def test_delete_mcp_server_removes_orphaned_user_credentials(): + """Deleting a server must also drop every user's stored BYOK/OAuth credential + rows for it; there is no FK cascade, so skipping this leaves encrypted secrets + pointing at a now-missing server.""" + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.db import delete_mcp_server + + prisma = _mock_env_vars_prisma() + prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=object()) + + await delete_mcp_server(prisma, "srv-1") + + prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once() + call = prisma.db.litellm_mcpusercredentials.delete_many.call_args + assert call.kwargs["where"] == {"server_id": "srv-1"} + + +@pytest.mark.asyncio +async def test_delete_mcp_server_skips_credential_cleanup_when_server_missing(): + """A no-op delete (server not found) must not touch the credential table.""" + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.db import delete_mcp_server + + prisma = _mock_env_vars_prisma() + prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=None) + + result = await delete_mcp_server(prisma, "srv-1") + + assert result is None + prisma.db.litellm_mcpusercredentials.delete_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_delete_mcp_server_credential_cleanup_failure_still_cleans_env_vars(): + """Each per-user table is cleaned independently: a failure dropping credential + rows must not skip the env var cleanup (or vice versa), and the delete must + still succeed for the caller.""" + from unittest.mock import AsyncMock + + from litellm.proxy._experimental.mcp_server.db import delete_mcp_server + + deleted = object() + prisma = _mock_env_vars_prisma() + prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=deleted) + prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock( + side_effect=Exception("connection pool exhausted") + ) + + result = await delete_mcp_server(prisma, "srv-1") + + assert result is deleted + prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once() + prisma.db.litellm_mcpuserenvvars.delete_many.assert_awaited_once() + + # ── DB helpers: global env vars encrypted at rest ───────────────────────── diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 0b1240f8bac..b6550fee6b9 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -4970,17 +4970,143 @@ class TestGatewayCreateInitializationOptions: try: from litellm.proxy._experimental.mcp_server.mcp_context import ( _mcp_gateway_initialize_instructions, + _mcp_gateway_server_name, ) from litellm.proxy._experimental.mcp_server.server import server except ImportError: pytest.skip("MCP server not available") - tok = _mcp_gateway_initialize_instructions.set(None) + instructions_token = _mcp_gateway_initialize_instructions.set(None) + server_name_token = _mcp_gateway_server_name.set(None) try: opts = server.create_initialization_options() assert getattr(opts, "instructions", None) is None + assert opts.server_name == "litellm-mcp-server" finally: - _mcp_gateway_initialize_instructions.reset(tok) + _mcp_gateway_initialize_instructions.reset(instructions_token) + _mcp_gateway_server_name.reset(server_name_token) + + @pytest.mark.asyncio + async def test_scoped_request_uses_configured_server_alias(self): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _gateway_initialize_instructions_request_scope, + global_mcp_server_manager, + server, + ) + except ImportError: + pytest.skip("MCP server not available") + + scoped_server = MCPServer( + server_id="server-123", + name="upstream-server", + alias="grafana", + transport=MCPTransport.http, + url="https://example.com/mcp", + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[scoped_server], + ), + patch.object( + global_mcp_server_manager, + "_ensure_upstream_initialize_instructions_cached", + new_callable=AsyncMock, + ), + ): + async with _gateway_initialize_instructions_request_scope( + user_api_key_auth=None, + mcp_servers=["grafana"], + client_ip=None, + scoped_server_endpoint=True, + ): + assert server.create_initialization_options().server_name == "grafana" + + assert ( + server.create_initialization_options().server_name == "litellm-mcp-server" + ) + + @pytest.mark.asyncio + async def test_sse_handler_scopes_server_name_from_single_server_path(self): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import ( + global_mcp_server_manager, + handle_sse_mcp, + server, + ) + except ImportError: + pytest.skip("MCP server not available") + + scoped_server = MCPServer( + server_id="server-123", + name="upstream-server", + alias="grafana", + transport=MCPTransport.http, + url="https://example.com/mcp", + ) + captured = {} + + async def record_request(scope, receive, send): + captured["server_name"] = server.create_initialization_options().server_name + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/grafana", + "headers": [], + } + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + UserAPIKeyAuth(api_key="sk-test"), + None, + ["grafana"], + None, + None, + None, + ), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[scoped_server], + ), + patch.object( + global_mcp_server_manager, + "_ensure_upstream_initialize_instructions_cached", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch.object( + mcp_server.sse_session_manager, + "handle_request", + side_effect=record_request, + ), + ): + await handle_sse_mcp(scope, AsyncMock(), AsyncMock()) + + assert captured["server_name"] == "grafana" + assert ( + server.create_initialization_options().server_name == "litellm-mcp-server" + ) def test_contextvar_set_injects_instructions(self): """When ContextVar has a value, it appears in InitializationOptions.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 48c09f6e456..1b815b7a1c9 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -30,6 +30,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _deserialize_json_dict, _deserialize_json_list, + _normalize_mcp_server_cost_info, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -257,6 +258,69 @@ class TestMCPServerManager: assert server.alias == "friendly_alias" assert server.server_name == "validserver" + @pytest.mark.asyncio + async def test_load_servers_from_config_coerces_cost_string_to_float(self): + """YAML 1.1 parses `7e-05` as a string; ingest must coerce it to float.""" + manager = MCPServerManager() + config = { + "google_maps": { + "url": "https://example.com/mcp", + "transport": MCPTransport.http, + "mcp_info": { + "mcp_server_cost_info": { + "default_cost_per_query": "7e-05", + "tool_name_to_cost_per_query": {"geocode": "1e-3"}, + } + }, + } + } + + await manager.load_servers_from_config(config) + + server = next(iter(manager.config_mcp_servers.values())) + cost_info = server.mcp_info["mcp_server_cost_info"] + assert cost_info["default_cost_per_query"] == 7e-05 + assert isinstance(cost_info["default_cost_per_query"], float) + assert cost_info["tool_name_to_cost_per_query"]["geocode"] == 1e-3 + assert isinstance(cost_info["tool_name_to_cost_per_query"]["geocode"], float) + + def test_normalize_mcp_server_cost_info_preserves_float_values(self): + mcp_info = { + "server_name": "maps", + "mcp_server_cost_info": { + "default_cost_per_query": 0.01, + "tool_name_to_cost_per_query": {"search": 0.05}, + }, + } + + _normalize_mcp_server_cost_info(mcp_info) + + cost_info = mcp_info["mcp_server_cost_info"] + assert cost_info["default_cost_per_query"] == 0.01 + assert cost_info["tool_name_to_cost_per_query"] == {"search": 0.05} + + def test_normalize_mcp_server_cost_info_drops_non_numeric_values(self): + mcp_info = { + "server_name": "maps", + "mcp_server_cost_info": { + "default_cost_per_query": "not-a-number", + "tool_name_to_cost_per_query": {"search": "free", "geocode": "2e-4"}, + }, + } + + _normalize_mcp_server_cost_info(mcp_info) + + cost_info = mcp_info["mcp_server_cost_info"] + assert "default_cost_per_query" not in cost_info + assert cost_info["tool_name_to_cost_per_query"] == {"geocode": 2e-4} + + def test_normalize_mcp_server_cost_info_leaves_missing_cost_info_alone(self): + mcp_info = {"server_name": "maps"} + + _normalize_mcp_server_cost_info(mcp_info) + + assert "mcp_server_cost_info" not in mcp_info + def test_warns_when_custom_separator_invalid(self, monkeypatch, caplog): """Invalid MCP_TOOL_PREFIX_SEPARATOR values should log a warning.""" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 42c76c4671d..e14ef05bd43 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -2409,6 +2409,8 @@ async def test_virtual_key_budget_check_fallback_no_counter(): assert exc_info.value.current_cost == 15.0 + + @pytest.mark.asyncio async def test_team_budget_check_reads_from_spend_counter(): """Team budget check should use get_current_spend when counter exists.""" diff --git a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py index 27f6015e6f4..11e6f483e35 100644 --- a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py +++ b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py @@ -112,6 +112,166 @@ async def test_handle_authentication_error_data_layer_errors_do_not_fall_back( ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "db_error", + [ + ConnectionError("connection refused"), + TimeoutError("timed out"), + asyncio.TimeoutError(), + OSError("network is unreachable"), + HTTPClientClosedError(), + PrismaError("can't reach database server"), + RawQueryError( + data={ + "user_facing_error": { + "message": "cached plan must not change result type", + "meta": {"table": "t"}, + } + } + ), + ], +) +async def test_handle_authentication_error_db_infra_error_returns_503(db_error): + """Regression for the outage where valid keys got 401 for 4 hours: an + infrastructure-level DB failure during auth must surface as 503 (the DB + could not confirm the key), never as 401 ("Invalid API key").""" + handler = UserAPIKeyAuthExceptionHandler() + + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy.auth.auth_exception_handler.seed_request_identity", + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException) as exc_info: + await handler._handle_authentication_error( + db_error, + MagicMock(), + {}, + "/v1/chat/completions", + None, + "sk-valid-but-db-down", + ) + + assert int(exc_info.value.code) == status.HTTP_503_SERVICE_UNAVAILABLE + assert exc_info.value.type == ProxyErrorTypes.no_db_connection + assert "Invalid API key" not in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_handle_authentication_error_prisma_engine_teardown_returns_503(): + """Regression for the first-request-of-an-outage edge case: at the instant + the DB socket drops, the prisma query engine returns a malformed error + payload and prisma-client-py crashes with a bare + ``AttributeError: 'NoneType' object has no attribute 'get'`` before it can + raise P1001. That AttributeError reached auth and fell through to 401. It + must surface as 503 like every other infra failure during the outage.""" + from prisma.engine import utils as prisma_engine_utils + + malformed_payload = [ + { + "error": "Can't reach database server", + "user_facing_error": { + "error_code": "P1001", + "message": "Can't reach database server at `localhost`:`5503`", + "meta": None, + }, + } + ] + try: + prisma_engine_utils.handle_response_errors(None, malformed_payload) + raise AssertionError("expected prisma to raise AttributeError") + except AttributeError as e: + teardown_error = e + + handler = UserAPIKeyAuthExceptionHandler() + + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy.auth.auth_exception_handler.seed_request_identity", + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException) as exc_info: + await handler._handle_authentication_error( + teardown_error, + MagicMock(), + {}, + "/v1/chat/completions", + None, + "sk-valid-but-db-down", + ) + + assert int(exc_info.value.code) == status.HTTP_503_SERVICE_UNAVAILABLE + assert exc_info.value.type == ProxyErrorTypes.no_db_connection + assert "Invalid API key" not in str(exc_info.value.message) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth_error", + [ + # DB returned no row -> get_key_object raises this exact 401. + ProxyException( + message="Authentication Error, Invalid proxy server token passed.", + type=ProxyErrorTypes.token_not_found_in_db, + param="key", + code=status.HTTP_401_UNAUTHORIZED, + ), + # A bare auth failure raised as a plain Exception (e.g. master-key-only + # route) must keep returning 401, not get reclassified as 503. + Exception("Invalid proxy server token passed"), + ], +) +async def test_handle_authentication_error_genuine_auth_failure_stays_401(auth_error): + """Guard against the 503 conversion being too broad: a genuine auth + failure (missing key / wrong key) must still be 401.""" + handler = UserAPIKeyAuthExceptionHandler() + + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy.auth.auth_exception_handler.seed_request_identity", + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException) as exc_info: + await handler._handle_authentication_error( + auth_error, + MagicMock(), + {}, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED + + @pytest.mark.asyncio async def test_handle_authentication_error_budget_exceeded(): handler = UserAPIKeyAuthExceptionHandler() diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 0236646c796..80f12d4459f 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1696,7 +1696,9 @@ class TestJWTOAuth2Coexistence: assert mock_auto_register.call_args.kwargs["team_id"] == "validated-team" assert mock_auto_register.call_args.kwargs["user_id"] == "validated-user" assert mock_auto_register.call_args.kwargs["org_id"] == "validated-org" - assert mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user" + assert ( + mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user" + ) assert result.org_id == "validated-org" @pytest.mark.asyncio @@ -3608,3 +3610,118 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() finally: for k, v in originals.items(): setattr(_proxy_server_mod, k, v) + + +def _proxy_attrs_for_db_lookup(): + """Minimal proxy_server attributes for driving the real + ``_user_api_key_auth_builder`` down to the DB key lookup.""" + proxy_logging_obj = MagicMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + return { + "prisma_client": MagicMock(), + "user_api_key_cache": DualCache(), + "proxy_logging_obj": proxy_logging_obj, + "master_key": "sk-test-master", + "general_settings": {"allow_requests_on_db_unavailable": False}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + + +async def _run_builder_with_key_lookup(get_key_object_mock): + """Drive the real auth builder with ``get_key_object`` replaced by the + given mock. Returns the builder result. Patches ``seed_request_identity`` + so the failure path doesn't touch OTEL.""" + from fastapi import Request + from starlette.datastructures import URL + + import litellm.proxy.proxy_server as _proxy_server_mod + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + attrs = _proxy_attrs_for_db_lookup() + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", + get_key_object_mock, + ), + patch( + "litellm.proxy.auth.auth_exception_handler.seed_request_identity", + ), + ): + return await _user_api_key_auth_builder( + request=request, + api_key="Bearer sk-db-lookup-test", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + +@pytest.mark.asyncio +async def test_builder_returns_503_when_db_lookup_raises_infra_error(): + """End-to-end: a DB infrastructure failure during the key lookup must + propagate past the ``except ProxyException`` guard and surface as 503, + not the 401 that masked the 4-hour outage. Killing the new 503 branch + flips this to 401 and fails the test.""" + get_key_object = AsyncMock(side_effect=ConnectionError("connection refused")) + + with pytest.raises(ProxyException) as exc_info: + await _run_builder_with_key_lookup(get_key_object) + + assert int(exc_info.value.code) == status.HTTP_503_SERVICE_UNAVAILABLE + assert exc_info.value.type == ProxyErrorTypes.no_db_connection + assert "Invalid API key" not in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_builder_returns_401_when_db_lookup_reports_missing_key(): + """Regression guard: a genuinely missing key (DB returned no row, which + ``get_key_object`` raises as a 401 ProxyException) must still be 401.""" + missing_key_error = ProxyException( + message="Authentication Error, Invalid proxy server token passed. key=..., not found in db.", + type=ProxyErrorTypes.token_not_found_in_db, + param="key", + code=status.HTTP_401_UNAUTHORIZED, + ) + get_key_object = AsyncMock(side_effect=missing_key_error) + + with pytest.raises(ProxyException) as exc_info: + await _run_builder_with_key_lookup(get_key_object) + + assert int(exc_info.value.code) == status.HTTP_401_UNAUTHORIZED + + +@pytest.mark.asyncio +async def test_builder_succeeds_when_db_lookup_returns_valid_token(): + """Regression guard: a valid key still authenticates. Proves the 503 + conversion only fires on the failure path and never intercepts success.""" + valid_token = UserAPIKeyAuth(api_key="sk-db-lookup-test", token="hashed-valid") + get_key_object = AsyncMock(return_value=valid_token) + + with patch( + "litellm.proxy.auth.user_api_key_auth._return_user_api_key_auth_obj", + new_callable=AsyncMock, + return_value=valid_token, + ) as mock_return: + result = await _run_builder_with_key_lookup(get_key_object) + + assert isinstance(result, UserAPIKeyAuth) + # Reaching the success-assembly return (never the exception handler) + # proves a valid key is unaffected by the 503 conversion. + mock_return.assert_awaited_once() diff --git a/tests/test_litellm/proxy/client/cli/test_agents.py b/tests/test_litellm/proxy/client/cli/test_agents.py new file mode 100644 index 00000000000..afd1696a89f --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/test_agents.py @@ -0,0 +1,475 @@ +import os +import sys +from unittest.mock import patch + +import click +import pytest +import requests +from click.testing import CliRunner + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + + +from litellm.proxy.client.cli.commands.agents import ( + AgentRunError, + agent_commands, + agent_launch_args, + agent_profile, + build_agent_env, + run_agent, + verify_proxy_key, +) + +AGENTS_MODULE = "litellm.proxy.client.cli.commands.agents" + + +def _agent_command(name): + return next(c for c in agent_commands() if c.name == name) + + +class _FakeResponse: + def __init__(self, status_code): + self.status_code = status_code + + +class TestAgentProfile: + def test_claude_is_anthropic(self): + name, profiles = agent_profile("claude") + assert name == "Claude Code" + assert profiles == frozenset({"anthropic"}) + + def test_claude_full_path_uses_basename(self): + name, profiles = agent_profile("/usr/local/bin/claude") + assert name == "Claude Code" + assert profiles == frozenset({"anthropic"}) + + def test_codex_and_opencode_are_openai(self): + assert agent_profile("codex") == ("Codex", frozenset({"openai"})) + assert agent_profile("opencode") == ("OpenCode", frozenset({"openai"})) + + def test_unknown_command_gets_both_profiles(self): + name, profiles = agent_profile("mytool") + assert name == "mytool" + assert profiles == frozenset({"anthropic", "openai"}) + + +class TestBuildAgentEnv: + def test_anthropic_profile_uses_bare_root_and_bearer(self): + env = build_agent_env( + {}, "http://localhost:4000/", "sk-key", frozenset({"anthropic"}) + ) + assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000" + assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-key" + assert "OPENAI_BASE_URL" not in env + assert "OPENAI_API_KEY" not in env + + def test_anthropic_profile_drops_existing_api_key(self): + env = build_agent_env( + {"ANTHROPIC_API_KEY": "real-key"}, + "http://localhost:4000", + "sk-key", + frozenset({"anthropic"}), + ) + assert "ANTHROPIC_API_KEY" not in env + + def test_openai_profile_appends_v1(self): + env = build_agent_env( + {}, "http://localhost:4000/", "sk-key", frozenset({"openai"}) + ) + assert env["OPENAI_BASE_URL"] == "http://localhost:4000/v1" + assert env["OPENAI_API_KEY"] == "sk-key" + assert "ANTHROPIC_BASE_URL" not in env + + def test_both_profiles_set_everything(self): + env = build_agent_env( + {}, "http://localhost:4000", "sk-key", frozenset({"anthropic", "openai"}) + ) + assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000" + assert env["OPENAI_BASE_URL"] == "http://localhost:4000/v1" + assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-key" + assert env["OPENAI_API_KEY"] == "sk-key" + + def test_preserves_unrelated_env_and_does_not_mutate_input(self): + base = {"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "real-key"} + env = build_agent_env( + base, "http://localhost:4000", "sk-key", frozenset({"anthropic"}) + ) + assert env["PATH"] == "/usr/bin" + assert base == {"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "real-key"} + + +class TestAgentLaunchArgs: + def test_claude_and_opencode_get_no_extra_args(self): + assert agent_launch_args("claude", "http://localhost:4000") == [] + assert agent_launch_args("opencode", "http://localhost:4000") == [] + + def test_unknown_agent_gets_no_extra_args(self): + assert agent_launch_args("mytool", "http://localhost:4000") == [] + + def test_codex_points_provider_at_proxy_over_http(self): + args = agent_launch_args("codex", "http://localhost:4000/") + joined = " ".join(args) + assert 'model_provider="litellm"' in args + assert 'model_providers.litellm.base_url="http://localhost:4000/v1"' in args + assert 'model_providers.litellm.env_key="OPENAI_API_KEY"' in args + assert 'model_providers.litellm.wire_api="responses"' in args + assert "model_providers.litellm.supports_websockets=false" in args + assert joined.count("-c") == 6 + + def test_codex_uses_basename(self): + assert agent_launch_args("/usr/local/bin/codex", "http://localhost:4000") == ( + agent_launch_args("codex", "http://localhost:4000") + ) + + +class TestVerifyProxyKey: + def test_ok_status_passes_and_uses_models_endpoint(self): + captured = {} + + def fake_get(url, headers, timeout): + captured["url"] = url + captured["headers"] = headers + return _FakeResponse(200) + + verify_proxy_key("http://localhost:4000/", "sk-key", get=fake_get) + + assert captured["url"] == "http://localhost:4000/v1/models" + assert captured["headers"] == {"Authorization": "Bearer sk-key"} + + @pytest.mark.parametrize("status", [401, 403]) + def test_rejected_key_raises(self, status): + with pytest.raises(AgentRunError, match="rejected your key"): + verify_proxy_key( + "http://localhost:4000", + "sk-key", + get=lambda *a, **k: _FakeResponse(status), + ) + + def test_unreachable_proxy_raises(self): + def boom(*a, **k): + raise requests.ConnectionError("refused") + + with pytest.raises(AgentRunError, match="Could not reach"): + verify_proxy_key("http://localhost:4000", "sk-key", get=boom) + + def test_other_non_2xx_is_tolerated(self): + verify_proxy_key( + "http://localhost:4000", + "sk-key", + get=lambda *a, **k: _FakeResponse(500), + ) + + +class TestRunAgent: + def test_wires_env_and_launches_resolved_binary(self): + calls = {} + + def fake_launcher(path, args, env): + calls["path"] = path + calls["args"] = tuple(args) + calls["env"] = dict(env) + + run_agent( + "http://localhost:4000", + "sk-key", + ["claude", "--resume"], + base_env={"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "leaked"}, + which=lambda name: "/usr/local/bin/claude", + verify=lambda *a: None, + launcher=fake_launcher, + ) + + assert calls["path"] == "/usr/local/bin/claude" + assert calls["args"] == ("claude", "--resume") + env = calls["env"] + assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000" + assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-key" + assert "ANTHROPIC_API_KEY" not in env + assert "OPENAI_BASE_URL" not in env + + def test_codex_gets_openai_env(self): + calls = {} + run_agent( + "http://localhost:4000", + "sk-key", + ["codex"], + base_env={}, + which=lambda name: "/usr/local/bin/codex", + verify=lambda *a: None, + launcher=lambda p, a, e: calls.update(env=dict(e)), + ) + assert calls["env"]["OPENAI_BASE_URL"] == "http://localhost:4000/v1" + assert calls["env"]["OPENAI_API_KEY"] == "sk-key" + assert "ANTHROPIC_BASE_URL" not in calls["env"] + + def test_codex_injects_proxy_provider_args_before_user_args(self): + calls = {} + run_agent( + "http://localhost:4000", + "sk-key", + ["codex", "exec", "do a thing"], + base_env={}, + which=lambda name: "/usr/local/bin/codex", + verify=lambda *a: None, + launcher=lambda p, a, e: calls.update(args=tuple(a)), + ) + args = calls["args"] + assert args[0] == "codex" + assert args[-2:] == ("exec", "do a thing") + assert 'model_provider="litellm"' in args + assert 'model_providers.litellm.base_url="http://localhost:4000/v1"' in args + # overrides must precede the codex subcommand so codex parses them + assert args.index('model_provider="litellm"') < args.index("exec") + + def test_claude_launches_without_injected_args(self): + calls = {} + run_agent( + "http://localhost:4000", + "sk-key", + ["claude", "--resume"], + base_env={}, + which=lambda name: "/usr/local/bin/claude", + verify=lambda *a: None, + launcher=lambda p, a, e: calls.update(args=tuple(a)), + ) + assert calls["args"] == ("claude", "--resume") + + def test_missing_binary_raises_with_install_hint(self): + with pytest.raises(AgentRunError, match="claude.*Install it first"): + run_agent( + "http://localhost:4000", + "sk-key", + ["claude"], + base_env={}, + which=lambda name: None, + verify=lambda *a: None, + launcher=lambda *a: None, + ) + + def test_skip_verify_does_not_call_verify(self): + verified = [] + launched = [] + run_agent( + "http://localhost:4000", + "sk-key", + ["claude"], + skip_verify=True, + base_env={}, + which=lambda name: "/usr/local/bin/claude", + verify=lambda *a: verified.append(a), + launcher=lambda *a: launched.append(a), + ) + assert verified == [] + assert len(launched) == 1 + + def test_verify_failure_aborts_before_launch(self): + launched = [] + + def boom(*a): + raise AgentRunError("rejected") + + with pytest.raises(AgentRunError): + run_agent( + "http://localhost:4000", + "sk-key", + ["claude"], + base_env={}, + which=lambda name: "/usr/local/bin/claude", + verify=boom, + launcher=lambda *a: launched.append(a), + ) + assert launched == [] + + def test_empty_command_raises(self): + with pytest.raises(AgentRunError): + run_agent("http://localhost:4000", "sk-key", []) + + def test_reattach_terminal_runs_just_before_launch(self): + order = [] + run_agent( + "http://localhost:4000", + "sk-key", + ["claude"], + skip_verify=True, + base_env={}, + which=lambda name: "/usr/local/bin/claude", + launcher=lambda *a: order.append("launch"), + reattach_terminal=lambda: order.append("reattach"), + ) + assert order == ["reattach", "launch"] + + def test_no_reattach_terminal_by_default(self): + order = [] + run_agent( + "http://localhost:4000", + "sk-key", + ["claude"], + skip_verify=True, + base_env={}, + which=lambda name: "/usr/local/bin/claude", + launcher=lambda *a: order.append("launch"), + ) + assert order == ["launch"] + + +class TestAgentCommands: + def setup_method(self): + self.runner = CliRunner() + + def test_one_command_per_known_agent(self): + assert {c.name for c in agent_commands()} == {"claude", "codex", "opencode"} + + def test_claude_launches_with_stored_key_and_forwards_args(self): + captured = {} + + def fake_run_agent(base_url, api_key, command, **kwargs): + captured["base_url"] = base_url + captured["api_key"] = api_key + captured["command"] = list(command) + captured["skip_verify"] = kwargs.get("skip_verify") + + with patch(f"{AGENTS_MODULE}.run_agent", side_effect=fake_run_agent): + result = self.runner.invoke( + _agent_command("claude"), + ["--resume", "-p", "hi"], + obj={"base_url": "http://localhost:4000", "api_key": "sk-key"}, + ) + + assert result.exit_code == 0, result.output + assert captured["api_key"] == "sk-key" + assert captured["command"] == ["claude", "--resume", "-p", "hi"] + assert captured["skip_verify"] is False + assert ( + "routing Claude Code through proxy at http://localhost:4000" + in result.output + ) + + def test_codex_shows_friendly_name(self): + captured = {} + with patch( + f"{AGENTS_MODULE}.run_agent", + side_effect=lambda b, k, c, **kw: captured.update(command=list(c)), + ): + result = self.runner.invoke( + _agent_command("codex"), + ["exec", "do a thing"], + obj={"base_url": "http://localhost:4000", "api_key": "sk-key"}, + ) + assert result.exit_code == 0, result.output + assert captured["command"] == ["codex", "exec", "do a thing"] + assert "routing Codex through proxy" in result.output + + def test_skip_verify_is_consumed_not_forwarded(self): + captured = {} + + def fake_run_agent(base_url, api_key, command, **kwargs): + captured["command"] = list(command) + captured["skip_verify"] = kwargs.get("skip_verify") + + with patch(f"{AGENTS_MODULE}.run_agent", side_effect=fake_run_agent): + result = self.runner.invoke( + _agent_command("claude"), + ["--skip-verify", "--resume"], + obj={"base_url": "http://localhost:4000", "api_key": "sk-key"}, + ) + + assert result.exit_code == 0, result.output + assert captured["skip_verify"] is True + assert captured["command"] == ["claude", "--resume"] + + def test_non_interactive_without_key_errors_clearly(self): + with ( + patch(f"{AGENTS_MODULE}._is_interactive", return_value=False), + patch(f"{AGENTS_MODULE}.run_agent") as mock_run, + ): + result = self.runner.invoke( + _agent_command("claude"), + [], + obj={"base_url": "http://localhost:4000", "api_key": None}, + ) + assert result.exit_code != 0 + assert "LITELLM_PROXY_API_KEY" in result.output + mock_run.assert_not_called() + + def test_interactive_without_key_logs_in_then_launches(self): + captured = {} + + @click.command() + def fake_login(): + pass + + with ( + patch(f"{AGENTS_MODULE}._is_interactive", return_value=True), + patch(f"{AGENTS_MODULE}.login", fake_login), + patch( + f"{AGENTS_MODULE}.get_stored_api_key", return_value="sk-after-login" + ) as mock_get, + patch( + f"{AGENTS_MODULE}.run_agent", + side_effect=lambda base_url, api_key, command, **k: captured.update( + api_key=api_key + ), + ), + ): + result = self.runner.invoke( + _agent_command("claude"), + [], + obj={"base_url": "http://localhost:4000", "api_key": None}, + ) + + assert result.exit_code == 0, result.output + assert captured["api_key"] == "sk-after-login" + mock_get.assert_called_once_with(expected_base_url="http://localhost:4000") + + def test_agent_run_error_becomes_click_error(self): + with patch( + f"{AGENTS_MODULE}.run_agent", + side_effect=AgentRunError("could not reach proxy"), + ): + result = self.runner.invoke( + _agent_command("claude"), + [], + obj={"base_url": "http://localhost:4000", "api_key": "sk-key"}, + ) + assert result.exit_code != 0 + assert "could not reach proxy" in result.output + + def test_interactive_session_reattaches_terminal_before_handoff(self): + from litellm.proxy.client.cli.commands.agents import ( + _restore_controlling_terminal, + ) + + captured = {} + with ( + patch(f"{AGENTS_MODULE}._is_interactive", return_value=True), + patch( + f"{AGENTS_MODULE}.run_agent", + side_effect=lambda b, k, c, **kw: captured.update(kw), + ), + ): + result = self.runner.invoke( + _agent_command("claude"), + [], + obj={"base_url": "http://localhost:4000", "api_key": "sk-key"}, + ) + assert result.exit_code == 0, result.output + assert captured["reattach_terminal"] is _restore_controlling_terminal + + def test_non_interactive_agent_mode_leaves_stdin_alone(self): + captured = {} + with ( + patch(f"{AGENTS_MODULE}._is_interactive", return_value=False), + patch( + f"{AGENTS_MODULE}.run_agent", + side_effect=lambda b, k, c, **kw: captured.update(kw), + ), + ): + result = self.runner.invoke( + _agent_command("claude"), + [], + obj={"base_url": "http://localhost:4000", "api_key": "sk-key"}, + ) + assert result.exit_code == 0, result.output + assert captured["reattach_terminal"] is None diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 2e738ff900d..4ee8b502aa2 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -517,7 +517,7 @@ class TestWhoamiCommand: assert result.exit_code == 0 assert "❌ Not authenticated" in result.output - assert "Run 'litellm-proxy login'" in result.output + assert "Run 'lite login'" in result.output def test_whoami_old_token(self): """Test whoami with old token showing warning""" diff --git a/tests/test_litellm/proxy/common_utils/test_callback_utils.py b/tests/test_litellm/proxy/common_utils/test_callback_utils.py index d328d68dcd4..36ff3f3c399 100644 --- a/tests/test_litellm/proxy/common_utils/test_callback_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_callback_utils.py @@ -1,7 +1,9 @@ import copy import sys import os -from types import SimpleNamespace +from types import ModuleType, SimpleNamespace + +import pytest sys.path.insert( 0, os.path.abspath("../../..") @@ -309,3 +311,98 @@ def test_encrypt_callback_vars_only_encrypts_credential_fields(monkeypatch): assert cv["langfuse_host"] == "https://cloud.langfuse.com" assert cv["langsmith_project"] == "my-proj" assert cv["langsmith_base_url"] == "https://smith.example" + + +def test_initialize_callbacks_on_proxy_lakera_ignores_non_dict_callback_settings( + monkeypatch, +): + """Regression: a non-dict value under callback_settings.lakera_prompt_injection + must not crash initialize_callbacks_on_proxy. + + Forwarding callback_settings as callback_specific_params (so callbacks like + DatadogCostManagementLogger receive their init params) exposes the lakera + branch, which previously did lakeraAI_Moderation(**callback_specific_params[ + "lakera_prompt_injection"]) with no isinstance(dict) guard. For a config like + {"lakera_prompt_injection": "x"} that is `**"x"` -> TypeError: argument after + ** must be a mapping, not str. The branch now guards on isinstance(dict), + matching the presidio / datadog_cost_management branches. + """ + captured = {} + + class _DummyLakera: + def __init__(self, **kwargs): + captured["kwargs"] = kwargs + + # Inject a fake lakera_ai module so the branch's + # `from ...lakera_ai import lakeraAI_Moderation` resolves to our stub without + # importing the real module (which imports proxy_server symbols not present + # under the stubbed proxy_server below). + fake_lakera = ModuleType("litellm.proxy.guardrails.guardrail_hooks.lakera_ai") + fake_lakera.lakeraAI_Moderation = _DummyLakera + monkeypatch.setitem( + sys.modules, + "litellm.proxy.guardrails.guardrail_hooks.lakera_ai", + fake_lakera, + ) + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + SimpleNamespace(prisma_client=None), + ) + + original_callbacks = ( + list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] + ) + litellm.callbacks = [] + try: + # A non-dict value must be ignored (init_params stays {}), not **-unpacked. + initialize_callbacks_on_proxy( + value=["lakera_prompt_injection"], + premium_user=False, + config_file_path=".", + litellm_settings={}, + callback_specific_params={"lakera_prompt_injection": "any-string"}, + ) + assert captured["kwargs"] == {} + assert any(isinstance(c, _DummyLakera) for c in litellm.callbacks) + finally: + litellm.callbacks = original_callbacks + + +@pytest.mark.parametrize("bad_root", [None, True]) +def test_initialize_callbacks_on_proxy_non_dict_callback_specific_params_root( + monkeypatch, bad_root +): + """Regression: a blank `callback_settings:` key in YAML loads as None (and + `callback_settings: true` as a bool); load_config forwards that value + verbatim as callback_specific_params. Membership tests like + `"compression_interception" in callback_specific_params` then raise + TypeError and abort proxy startup. A non-dict root must be normalized to {} + so the callback initializes with its defaults. + """ + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + SimpleNamespace(prisma_client=None), + ) + from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, + ) + + original_callbacks = ( + list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] + ) + litellm.callbacks = [] + try: + initialize_callbacks_on_proxy( + value=["compression_interception"], + premium_user=False, + config_file_path=".", + litellm_settings={}, + callback_specific_params=bad_root, + ) + assert any( + isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks + ) + finally: + litellm.callbacks = original_callbacks diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 9dcf5df4aeb..6021c221426 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -107,6 +107,201 @@ def test_is_database_connection_generic_errors(): ) +@pytest.mark.parametrize( + "error", + [ + ConnectionError("connection refused"), + TimeoutError("timed out"), + OSError("network is unreachable"), + asyncio.TimeoutError(), + HTTPClientClosedError(), + ClientNotConnectedError(), + PrismaError("can't reach database server"), + PrismaError(), + ], +) +def test_is_database_service_unavailable_error_infra_failures(error): + """Infrastructure-level failures (socket/connection/timeout, prisma + transport, unknown PrismaError) mean the DB could not answer, so auth + must surface 503 instead of treating a valid key as invalid.""" + assert PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is True + + +def test_is_database_service_unavailable_error_prisma_p1001_masquerades_as_dataerror(): + """Real-world regression: prisma-client-py raises the P1001 "can't reach + database server" connectivity failure as a DataError (a data-layer type). + A type-only check would miss it and return 401 during a genuine outage; + the message keyword must still classify it as service-unavailable -> 503.""" + p1001_as_dataerror = DataError( + data={ + "user_facing_error": { + "message": "Can't reach database server at `127.0.0.1`:`5499`", + "meta": {"table": "t"}, + } + } + ) + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error( + p1001_as_dataerror + ) + is True + ) + + +def test_is_database_service_unavailable_error_cached_plan_escapes_as_503(): + """Composes with the cached-plan retry: when that recovery fails and the + Postgres "cached plan must not change result type" error escapes (raised by + prisma as a data-layer RawQueryError), it is a transient stale-DB-state + condition, not an invalid key, so it must classify as service-unavailable + -> 503 rather than fall through to 401.""" + cached_plan_error = RawQueryError( + data={ + "user_facing_error": { + "message": "cached plan must not change result type", + "meta": {"table": "t"}, + } + } + ) + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error( + cached_plan_error + ) + is True + ) + + +def test_is_database_service_unavailable_error_prisma_engine_malformed_payload(): + """Real-world regression: at the instant the DB socket drops, the prisma + query engine returns a malformed error payload (``user_facing_error.meta`` + is ``null``). prisma-client-py's ``handle_response_errors`` then crashes + with ``AttributeError: 'NoneType' object has no attribute 'get'`` before it + can raise the proper P1001 error. That bare AttributeError has no + connection keyword, so without the prisma-engine-origin check it falls + through to 401 on the first request of an outage. Reproduce the exact + prisma crash and assert it classifies as service-unavailable -> 503.""" + from prisma.engine import utils as prisma_engine_utils + + malformed_payload = [ + { + "error": "Can't reach database server", + "user_facing_error": { + "error_code": "P1001", + "message": "Can't reach database server at `localhost`:`5503`", + "meta": None, + }, + } + ] + with pytest.raises(AttributeError) as exc_info: + prisma_engine_utils.handle_response_errors(None, malformed_payload) + + assert "no attribute 'get'" in str(exc_info.value) + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) + is True + ) + + +def test_is_prisma_engine_internal_error_excludes_application_attributeerror(): + """The prisma-engine-origin check must stay narrow: a genuine AttributeError + raised by application code (a real bug) must NOT be classified as + service-unavailable, otherwise real bugs would silently become 503s.""" + + def application_bug(): + none_value = None + return none_value.get("oops") + + with pytest.raises(AttributeError) as exc_info: + application_bug() + + assert ( + PrismaDBExceptionHandler.is_prisma_engine_internal_error(exc_info.value) + is False + ) + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) + is False + ) + + +def test_is_prisma_engine_internal_error_excludes_data_layer_prisma_error(): + """A data-layer ``PrismaError`` (the DB IS reachable and rejected the data) + must stay 401. These are always raised from prisma internals, so the check + excludes any ``PrismaError`` by type before inspecting the traceback.""" + data_layer_error = UniqueViolationError( + data={"user_facing_error": {"meta": {"table": "t"}}} + ) + try: + raise data_layer_error + except UniqueViolationError as e: + assert PrismaDBExceptionHandler.is_prisma_engine_internal_error(e) is False + + +@pytest.mark.parametrize( + "error", + [ + DataError(data={"user_facing_error": {"meta": {"table": "t"}}}), + UniqueViolationError(data={"user_facing_error": {"meta": {"table": "t"}}}), + RecordNotFoundError(data={"user_facing_error": {"meta": {"table": "t"}}}), + Exception("some unrelated error"), + ValueError("bad value"), + ], +) +def test_is_database_service_unavailable_error_excludes_non_infra(error): + """Data-layer errors (the DB IS reachable and answered) and generic + non-DB errors must NOT be classified as service-unavailable, otherwise a + genuine 401 would be masked as a transient 503.""" + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False + ) + + +def test_is_database_service_unavailable_error_asyncpg(monkeypatch): + """asyncpg connection/interface errors map to service-unavailable. asyncpg + is not a hard dependency, so inject a stand-in module to exercise the + branch deterministically regardless of the install environment.""" + import sys + import types + + fake_asyncpg = types.ModuleType("asyncpg") + fake_exceptions = types.ModuleType("asyncpg.exceptions") + + class PostgresConnectionError(Exception): + pass + + class InterfaceError(Exception): + pass + + class UniqueViolationError(Exception): # data-layer, must stay False + pass + + fake_exceptions.PostgresConnectionError = PostgresConnectionError + fake_exceptions.InterfaceError = InterfaceError + fake_exceptions.UniqueViolationError = UniqueViolationError + fake_asyncpg.exceptions = fake_exceptions + + monkeypatch.setitem(sys.modules, "asyncpg", fake_asyncpg) + monkeypatch.setitem(sys.modules, "asyncpg.exceptions", fake_exceptions) + + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error( + PostgresConnectionError("connection reset") + ) + is True + ) + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error( + InterfaceError("connection was closed") + ) + is True + ) + assert ( + PrismaDBExceptionHandler.is_database_service_unavailable_error( + UniqueViolationError("duplicate key") + ) + is False + ) + + # Test should_allow_request_on_db_unavailable method @patch( "litellm.proxy.proxy_server.general_settings", diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py index 8074871c3dd..7bf1ffda4fe 100644 --- a/tests/test_litellm/proxy/db/test_tool_registry_writer.py +++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py @@ -291,3 +291,77 @@ async def test_tool_policy_registry_not_initialized_returns_untrusted(): assert not registry.is_initialized() result = registry.get_effective_policies(["unknown_tool"]) assert result == {"unknown_tool": "untrusted"} + + +@pytest.mark.asyncio +async def test_sync_tool_policy_from_db_retries_on_transport_error_first_read(): + """`ToolPolicyRegistry.sync_tool_policy_from_db` self-heals across one + ClientNotConnectedError on the tools read — the perms read still fires + after the recovery and the registry initializes cleanly.""" + import prisma as prisma_pkg + + registry = ToolPolicyRegistry() + invocations: list = [] + + async def _flaky_find_many(): + invocations.append(None) + if len(invocations) == 1: + raise prisma_pkg.errors.ClientNotConnectedError() + return [] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_tooltable.find_many = AsyncMock( + side_effect=_flaky_find_many + ) + mock_prisma_client.db.litellm_objectpermissiontable.find_many = AsyncMock( + return_value=[] + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + + await registry.sync_tool_policy_from_db(mock_prisma_client) + + assert len(invocations) == 2 + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs + assert ( + reconnect_kwargs["reason"] + == "sync_tool_policy_from_db_tools_lookup_failure" + ) + assert registry.is_initialized() + + +@pytest.mark.asyncio +async def test_sync_tool_policy_from_db_retries_on_transport_error_second_read(): + """Same as above but the blip happens on the perms read — distinct reason + tag in telemetry confirms the second wrap is wired separately.""" + import prisma as prisma_pkg + + registry = ToolPolicyRegistry() + perms_invocations: list = [] + + async def _flaky_perms_find_many(): + perms_invocations.append(None) + if len(perms_invocations) == 1: + raise prisma_pkg.errors.ClientNotConnectedError() + return [] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_tooltable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_objectpermissiontable.find_many = AsyncMock( + side_effect=_flaky_perms_find_many + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + + await registry.sync_tool_policy_from_db(mock_prisma_client) + + assert len(perms_invocations) == 2 + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs + assert ( + reconnect_kwargs["reason"] + == "sync_tool_policy_from_db_perms_lookup_failure" + ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 565bf83c6a2..8a5eeeff367 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -2398,11 +2398,9 @@ async def test_anonymize_text_uses_correct_positions_no_parse_pii(): ) expected = "My name is , my email is , phone " - assert result == expected, ( - f"anonymize_text produced garbled output with PII remnants.\n" - f"Expected: {expected!r}\n" - f"Got: {result!r}" - ) + assert ( + result == expected + ), f"anonymize_text produced garbled output with PII remnants.\nExpected: {expected!r}\nGot: {result!r}" assert masked_entity_count == { "PERSON": 1, "EMAIL_ADDRESS": 1, @@ -2495,3 +2493,157 @@ async def test_anonymize_text_uses_correct_positions_with_parse_pii(): assert pii_tokens.get("") == "John Smith" assert pii_tokens.get("") == "john@example.com" assert pii_tokens.get("") == "555-867-5309" + + +def test_unmask_sse_bytes_chunk_replaces_text_delta(): + import json + + pii_tokens = {"": "Bobby"} + event = { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "Hello , how are you?"}, + } + chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") + + result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + + decoded = result.decode("utf-8") + parsed = json.loads(decoded.split("data: ", 1)[1].strip()) + assert parsed["delta"]["text"] == "Hello Bobby, how are you?" + + +def test_unmask_sse_bytes_chunk_ignores_non_text_delta(): + import json + + pii_tokens = {"": "Bobby"} + + # message_start event — no delta + event = {"type": "message_start", "message": {"id": "msg_01", "role": "assistant"}} + chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") + result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + assert result == chunk + + # input_json_delta — should not be touched + event2 = { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "input_json_delta", "partial_json": '{"name": ""}'}, + } + chunk2 = ("data: " + json.dumps(event2) + "\n\n").encode("utf-8") + result2 = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk2, pii_tokens) + assert result2 == chunk2 + + +def test_unmask_sse_bytes_chunk_handles_malformed_json(): + chunk = b"data: {not valid json}\n\n" + result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk( + chunk, {"": "Bobby"} + ) + assert result == chunk + + +def test_unmask_sse_bytes_chunk_handles_unicode_decode_error(): + chunk = b"\xff\xfe invalid utf-8" + result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk( + chunk, {"": "Bobby"} + ) + assert result == chunk + + +def test_unmask_sse_bytes_chunk_non_ascii_pii_not_escaped(): + import json + + pii_tokens = {"": "José"} + event = { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "Hello !"}, + } + chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8") + + result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens) + + decoded = result.decode("utf-8") + assert "Jos\\u" not in decoded + parsed = json.loads(decoded.split("data: ", 1)[1].strip()) + assert parsed["delta"]["text"] == "Hello José!" + + +def test_unmask_sse_bytes_chunk_handles_crlf_line_endings(): + import json + + pii_tokens = {"": "Bobby"} + event = { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "Hi !"}, + } + crlf_chunk = ("data: " + json.dumps(event) + "\r\ndata: [DONE]\r\n").encode("utf-8") + + result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk( + crlf_chunk, pii_tokens + ) + + decoded = result.decode("utf-8") + parsed = json.loads(decoded.split("data: ", 1)[1].split("\n")[0].strip()) + assert parsed["delta"]["text"] == "Hi Bobby!" + assert "data: [DONE]" in decoded + + +@pytest.mark.asyncio +async def test_stream_pii_unmasking_unmaskes_bytes_chunks(mock_user_api_key): + import json + + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + ) + + pii_tokens = {"": "Bobby"} + request_data = {"metadata": {"pii_tokens": pii_tokens}} + + def _make_sse_chunk(text: str) -> bytes: + event = { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": text}, + } + return ("data: " + json.dumps(event) + "\n\n").encode("utf-8") + + async def mock_stream(): + yield _make_sse_chunk("Hello !") + yield _make_sse_chunk(" How can I help?") + + chunks = [] + async for chunk in guardrail._stream_pii_unmasking(mock_stream(), request_data): + chunks.append(chunk) + + assert len(chunks) == 2 + first = chunks[0].decode("utf-8") + first_event = json.loads(first.split("data: ", 1)[1].strip()) + assert first_event["delta"]["text"] == "Hello Bobby!" + + second = chunks[1].decode("utf-8") + second_event = json.loads(second.split("data: ", 1)[1].strip()) + assert second_event["delta"]["text"] == " How can I help?" + + +@pytest.mark.asyncio +async def test_stream_pii_unmasking_passthrough_when_no_tokens(mock_user_api_key): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + output_parse_pii=True, + ) + + raw_chunk = b"data: {}\n\n" + request_data: dict = {"metadata": {}} + + async def mock_stream(): + yield raw_chunk + + chunks = [] + async for chunk in guardrail._stream_pii_unmasking(mock_stream(), request_data): + chunks.append(chunk) + + assert chunks == [raw_chunk] diff --git a/tests/test_litellm/proxy/hooks/litellm_skills/test_main.py b/tests/test_litellm/proxy/hooks/litellm_skills/test_main.py new file mode 100644 index 00000000000..f716a8533d8 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/litellm_skills/test_main.py @@ -0,0 +1,67 @@ +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.proxy.hooks.litellm_skills.main import SkillsInjectionHook + +SKILL_TOOL_NAME = "litellm_skill_e2b8dca8_031a_4481_b034_b9ec7d4eb7bf" + + +def _request_data(): + return { + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "run the skill"}], + "litellm_metadata": { + "_litellm_code_execution_enabled": True, + "_skill_files": {SKILL_TOOL_NAME: {"main.py": b"print('hi')"}}, + }, + } + + +def _tool_use_response(tool_name): + return { + "stop_reason": "tool_use", + "content": [ + {"type": "tool_use", "id": "toolu_1", "name": tool_name, "input": {}} + ], + } + + +@pytest.mark.asyncio +async def test_post_call_success_hook_executes_litellm_skill_tool(): + """DB skill tool names carry the litellm_skill_ prefix and must trigger the execution loop.""" + hook = SkillsInjectionHook() + response = _tool_use_response(SKILL_TOOL_NAME) + + with patch.object( + hook, "_execute_code_loop_messages_api", new=AsyncMock(return_value=response) + ) as mock_loop: + result = await hook.async_post_call_success_deployment_hook( + request_data=_request_data(), response=response, call_type=None + ) + + mock_loop.assert_awaited_once() + assert result is response + + +@pytest.mark.asyncio +async def test_execute_code_loop_dispatches_litellm_skill_tool(): + """The agentic loop must route litellm_skill_ tool calls to _execute_skill_tool.""" + hook = SkillsInjectionHook() + final_response = {"stop_reason": "end_turn", "content": []} + + with ( + patch.object( + hook, "_execute_skill_tool", new=AsyncMock(return_value="skill ran") + ) as mock_exec, + patch("litellm.anthropic.acreate", new=AsyncMock(return_value=final_response)), + ): + result = await hook._execute_code_loop_messages_api( + data=_request_data(), + response=_tool_use_response(SKILL_TOOL_NAME), + skill_files={"main.py": b"print('hi')"}, + ) + + mock_exec.assert_awaited_once() + assert mock_exec.await_args.args[0] == SKILL_TOOL_NAME + assert result is final_response diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index ae71d10b378..a6f6e651487 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -713,7 +713,8 @@ async def test_count_input_file_usage_decodes_model_embedded_file_id(): @pytest.mark.asyncio async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias(): """After replace_model_in_jsonl, body.model is the provider id (e.g. gpt-5.5). - Auth must check the proxy model_name the key was granted, not the stripped id.""" + Auth must check target_model_names from the unified file id, not reverse-map + the stripped id.""" from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter rate_limiter = _PROXY_BatchRateLimiter( @@ -732,7 +733,6 @@ async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias( ) mock_router = MagicMock() mock_router.model_list = [] - mock_router.resolve_model_name_from_model_id.return_value = proxy_alias can_key_call_model = AsyncMock(return_value=True) with ( @@ -745,10 +745,105 @@ async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias( await rate_limiter._enforce_batch_file_model_access( user_api_key_dict=user, file_content_as_dict=file_dict, + target_model_names=[proxy_alias], ) can_key_call_model.assert_awaited_once() assert can_key_call_model.await_args.kwargs["model"] == proxy_alias + mock_router.resolve_model_name_from_model_id.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "model_list_order", + [ + [ + "openai/openai/gpt-5.5", + "openai/openai/gpt-5.5-batch", + "us/azure/openai/gpt-5.5", + ], + [ + "us/azure/openai/gpt-5.5", + "openai/openai/gpt-5.5", + "openai/openai/gpt-5.5-batch", + ], + [ + "openai/openai/gpt-5.5-batch", + "us/azure/openai/gpt-5.5", + "openai/openai/gpt-5.5", + ], + ], +) +async def test_pre_call_uses_target_model_names_not_stripped_reverse_lookup( + model_list_order, +): + """LIT-3593: three deployments strip to gpt-5.5; auth must use the upload + target alias from target_model_names, not first-match reverse lookup.""" + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=MagicMock(), + ) + batch_alias = "openai/openai/gpt-5.5-batch" + deployment_templates = { + "openai/openai/gpt-5.5": { + "model_name": "openai/openai/gpt-5.5", + "litellm_params": {"model": "openai/gpt-5.5"}, + "model_info": {"id": "openai/openai/gpt-5.5", "mode": "chat"}, + }, + "openai/openai/gpt-5.5-batch": { + "model_name": "openai/openai/gpt-5.5-batch", + "litellm_params": {"model": "openai/gpt-5.5"}, + "model_info": {"id": "openai/openai/gpt-5.5-batch", "mode": "batch"}, + }, + "us/azure/openai/gpt-5.5": { + "model_name": "us/azure/openai/gpt-5.5", + "litellm_params": {"model": "azure/gpt-5.5"}, + "model_info": {"id": "openai/openai/gpt-5.5", "mode": "chat"}, + }, + } + mock_router = MagicMock() + mock_router.model_list = [deployment_templates[name] for name in model_list_order] + + def _resolve(model_id): + for deployment in mock_router.model_list: + actual_model = deployment.get("litellm_params", {}).get("model") + if actual_model == model_id or ( + actual_model and actual_model.endswith(f"/{model_id}") + ): + return deployment.get("model_name") + return None + + mock_router.resolve_model_name_from_model_id.side_effect = _resolve + + file_dict = [ + {"body": {"model": "gpt-5.5", "messages": [{"role": "user", "content": "x"}]}} + ] + user = UserAPIKeyAuth( + api_key="sk-ok", + user_id="alice", + models=[batch_alias], + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + can_key_call_model = AsyncMock(return_value=True) + + with ( + patch( + "litellm.proxy.auth.auth_checks.can_key_call_model", + new=can_key_call_model, + ), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + ): + await rate_limiter._enforce_batch_file_model_access( + user_api_key_dict=user, + file_content_as_dict=file_dict, + target_model_names=[batch_alias], + ) + + can_key_call_model.assert_awaited_once() + assert can_key_call_model.await_args.kwargs["model"] == batch_alias + mock_router.resolve_model_name_from_model_id.assert_not_called() @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py new file mode 100644 index 00000000000..0e2683dcbfd --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -0,0 +1,86 @@ +""" +Unit Tests for the max parallel request limiter v1 for the proxy +""" + +from datetime import datetime + +import pytest + +from litellm.caching.caching import DualCache +from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, +) +from litellm.proxy.utils import InternalUsageCache, hash_token +from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage + + +@pytest.mark.parametrize( + "response_obj", + [ + EmbeddingResponse( + model="text-embedding-3-small", + usage=Usage(prompt_tokens=50, completion_tokens=0, total_tokens=50), + ), + TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + usage=Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50), + ), + ], +) +@pytest.mark.asyncio +async def test_async_log_success_event_counts_non_chat_response_tokens(response_obj): + """ + Embedding and text completion responses must increment the per key, user, + team, and end user TPM counters, not just chat completion ModelResponse + objects. + """ + _api_key = hash_token("sk-12345") + user_id = "ishaan" + team_id = "litellm-team" + end_user_id = "customer-1" + + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + + current_date = datetime.now().strftime("%Y-%m-%d") + current_hour = datetime.now().strftime("%H") + current_minute = datetime.now().strftime("%M") + precise_minute = f"{current_date}-{current_hour}-{current_minute}" + + scope_ids = [_api_key, user_id, team_id, end_user_id] + for scope_id in scope_ids: + await parallel_request_handler.internal_usage_cache.async_set_cache( + key=f"{scope_id}::{precise_minute}::request_count", + value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1}, + litellm_parent_otel_span=None, + ) + + kwargs = { + "litellm_params": { + "metadata": { + "user_api_key": _api_key, + "user_api_key_user_id": user_id, + "user_api_key_team_id": team_id, + "user_api_key_model_max_budget": {}, + } + }, + "user": end_user_id, + } + + await parallel_request_handler.async_log_success_event( + kwargs=kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + for scope_id in scope_ids: + current = await parallel_request_handler.internal_usage_cache.async_get_cache( + key=f"{scope_id}::{precise_minute}::request_count", + litellm_parent_otel_span=None, + ) + assert current["current_tpm"] == 50, ( + f"expected 50 tokens counted for {scope_id}, " + f"got {current['current_tpm']}" + ) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 676f623a5dd..d10311b9f41 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -20,7 +20,12 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token -from litellm.types.utils import ModelResponse, Usage +from litellm.types.utils import ( + EmbeddingResponse, + ModelResponse, + TextCompletionResponse, + Usage, +) class TimeController: @@ -547,6 +552,68 @@ async def test_token_rate_limit_type_respected_v3(monkeypatch, token_rate_limit_ ), f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}" +@pytest.mark.parametrize( + "response_obj", + [ + EmbeddingResponse( + model="text-embedding-3-small", + usage=Usage(prompt_tokens=50, completion_tokens=0, total_tokens=50), + ), + TextCompletionResponse( + model="gpt-3.5-turbo-instruct", + usage=Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50), + ), + ], +) +@pytest.mark.asyncio +async def test_async_log_success_event_counts_non_chat_response_tokens( + monkeypatch, response_obj +): + """ + Embedding and text completion responses must increment the TPM counter, + not just chat completion ModelResponse objects. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + + _api_key = hash_token("sk-12345") + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + monkeypatch.setattr( + parallel_request_handler, "get_rate_limit_type", lambda: "total" + ) + + mock_kwargs = { + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, + "model": response_obj.model, + } + + captured_operations = [] + + async def mock_increment_pipeline(increment_list, **kwargs): + captured_operations.extend(increment_list) + return True + + monkeypatch.setattr( + parallel_request_handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + mock_increment_pipeline, + ) + + await parallel_request_handler.async_log_success_event( + kwargs=mock_kwargs, + response_obj=response_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + tpm_operation = next( + (op for op in captured_operations if op["key"].endswith(":tokens")), None + ) + assert tpm_operation is not None, "Should have a TPM increment operation" + assert tpm_operation["increment_value"] == 50 + + @pytest.mark.asyncio async def test_async_log_failure_event_v3(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py index ea7e5591f18..f2ccfcd0155 100644 --- a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py +++ b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py @@ -611,6 +611,45 @@ async def test_list_search_tools_db_masking_sensitive_values(monkeypatch): app.dependency_overrides.pop(user_api_key_auth, None) +@pytest.mark.asyncio +async def test_get_all_search_tools_from_db_retries_on_transport_error(): + """`SearchToolRegistry.get_all_search_tools_from_db` self-heals across one + ClientNotConnectedError via call_with_db_reconnect_retry.""" + import prisma + from litellm.proxy.search_endpoints.search_tool_registry import ( + SearchToolRegistry, + ) + + invocations: list = [] + + async def _flaky_find_many(**kwargs): + invocations.append(None) + if len(invocations) == 1: + raise prisma.errors.ClientNotConnectedError() + return [] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_searchtoolstable.find_many = AsyncMock( + side_effect=_flaky_find_many + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + + result = await SearchToolRegistry.get_all_search_tools_from_db( + prisma_client=mock_prisma_client + ) + + assert result == [] + assert len(invocations) == 2 + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs + assert ( + reconnect_kwargs["reason"] + == "get_all_search_tools_from_db_lookup_failure" + ) + + @contextlib.contextmanager def _mock_search_tool_backend(db_tools): """Patch the DB registry, prisma client, and config so /search_tools/list diff --git a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py index b892c4e556d..4bdef2e8f96 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py @@ -259,6 +259,41 @@ class TestCacheSettingsManager: mock_proxy_config._init_cache.assert_not_called() mock_proxy_config.switch_on_llm_response_caching.assert_not_called() + @pytest.mark.asyncio + async def test_init_cache_settings_in_db_retries_on_transport_error(self): + """`CacheSettingsManager.init_cache_settings_in_db` self-heals across one + ClientNotConnectedError via call_with_db_reconnect_retry.""" + import prisma + + invocations: list = [] + + async def _flaky_find_unique(**kwargs): + invocations.append(None) + if len(invocations) == 1: + raise prisma.errors.ClientNotConnectedError() + return None # No config → function returns early after retry. + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_cacheconfig.find_unique = AsyncMock( + side_effect=_flaky_find_unique + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + mock_proxy_config = MagicMock() + + await CacheSettingsManager.init_cache_settings_in_db( + prisma_client=mock_prisma_client, proxy_config=mock_proxy_config + ) + + assert len(invocations) == 2 + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs + assert ( + reconnect_kwargs["reason"] + == "init_cache_settings_in_db_lookup_failure" + ) + # ── Audit-log emission for /cache/settings ──────────────────────────────────── diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 3c212d86e65..473d61f8a85 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -6496,6 +6496,9 @@ async def test_reset_key_spend_success(monkeypatch): patch( "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" ) as mock_delete_cache, + patch( + "litellm.proxy.proxy_server._invalidate_spend_counter" + ) as mock_invalidate, ): mock_hash_token.return_value = hashed_key mock_check_admin.return_value = None @@ -6520,6 +6523,76 @@ async def test_reset_key_spend_success(monkeypatch): assert response["max_budget"] == 200.0 mock_prisma_client.db.litellm_verificationtoken.update.assert_called_once() mock_delete_cache.assert_awaited_once() + mock_invalidate.assert_awaited_once_with(counter_key=f"spend:key:{hashed_key}") + + +@pytest.mark.asyncio +async def test_update_key_spend_invalidates_counter(monkeypatch): + """ + Test that updating a key's spend via update_key_fn immediately invalidates the spend counter. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = AsyncMock() + mock_proxy_logging_obj = MagicMock() + + hashed_key = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" + key_in_db = LiteLLM_VerificationToken( + token=hashed_key, + user_id="test-user", + spend=10.0, + max_budget=200.0, + litellm_budget_table=None, + ) + + mock_prisma_client.get_data = AsyncMock(return_value=key_in_db) + mock_prisma_client.update_data = AsyncMock(return_value={"data": {"spend": 0.0}}) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_in_db + ) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache, + patch( + "litellm.proxy.proxy_server._invalidate_spend_counter" + ) as mock_invalidate, + ): + mock_delete_cache.return_value = None + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-user", + ) + + mock_request = MagicMock() + mock_request.query_params = {} + + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key="sk-test-key", spend=0.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + mock_delete_cache.assert_awaited_once() + mock_invalidate.assert_awaited_once_with(counter_key=f"spend:key:{hashed_key}") @pytest.mark.asyncio @@ -11668,3 +11741,84 @@ async def test_ghsa_q775_default_team_id_does_not_grant_session_token_exemption( msg = str(getattr(err, "detail", "")) + str(getattr(err, "message", "")) assert str(code) == "400" assert "cannot exceed" in msg.lower() + + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_budget_duration_null_clears_fields(): + """ + When budget_duration is explicitly set to null, prepare_key_update_data + should produce budget_duration=None and budget_reset_at=None so Prisma + clears them in the DB. + """ + existing_key = LiteLLM_VerificationToken( + token="test-token", + key_alias="test-key", + models=[], + user_id="test-user", + team_id=None, + metadata={}, + ) + + update_request = UpdateKeyRequest(key="test-token", budget_duration=None) + + result = await prepare_key_update_data( + data=update_request, existing_key_row=existing_key + ) + + assert "budget_duration" in result + assert result["budget_duration"] is None + assert "budget_reset_at" in result + assert result["budget_reset_at"] is None + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_budget_duration_not_sent_excluded(): + """ + When budget_duration is NOT sent in the request (unset), it should not + appear in the result dict at all — the existing DB value stays unchanged. + """ + existing_key = LiteLLM_VerificationToken( + token="test-token", + key_alias="test-key", + models=[], + user_id="test-user", + team_id=None, + metadata={}, + ) + + update_request = UpdateKeyRequest(key="test-token", models=["gpt-4"]) + + result = await prepare_key_update_data( + data=update_request, existing_key_row=existing_key + ) + + assert "budget_duration" not in result + assert "budget_reset_at" not in result + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_budget_duration_valid_sets_reset(): + """ + When budget_duration is set to a valid duration string, both + budget_duration and budget_reset_at should be populated. + """ + existing_key = LiteLLM_VerificationToken( + token="test-token", + key_alias="test-key", + models=[], + user_id="test-user", + team_id=None, + metadata={}, + ) + + update_request = UpdateKeyRequest(key="test-token", budget_duration="30d") + + result = await prepare_key_update_data( + data=update_request, existing_key_row=existing_key + ) + + assert result["budget_duration"] == "30d" + assert result["budget_reset_at"] is not None + + diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index b750b6d022c..d4bc3841668 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -8886,3 +8886,329 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): ) assert str(exc.value.code) == "403" assert "allowed_passthrough_routes" in str(exc.value.message) + + +def test_set_budget_reset_at_clears_when_budget_duration_null(): + """ + When budget_duration is explicitly set to null, _set_budget_reset_at + should set budget_reset_at=None in updated_kv so Prisma clears it in the DB. + """ + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import _set_budget_reset_at + + data = UpdateTeamRequest(team_id="test-team", budget_duration=None) + updated_kv = {"team_id": "test-team", "budget_duration": None} + + _set_budget_reset_at(data, updated_kv) + + assert "budget_reset_at" in updated_kv + assert updated_kv["budget_reset_at"] is None + + +def test_set_budget_reset_at_noop_when_budget_duration_not_sent(): + """ + When budget_duration is NOT sent (unset), _set_budget_reset_at should + not add budget_reset_at to updated_kv. + """ + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import _set_budget_reset_at + + data = UpdateTeamRequest(team_id="test-team") + updated_kv = {"team_id": "test-team"} + + _set_budget_reset_at(data, updated_kv) + + assert "budget_reset_at" not in updated_kv + + +def test_set_budget_reset_at_sets_value_when_budget_duration_provided(): + """ + When budget_duration is set to a valid string, _set_budget_reset_at + should compute and set budget_reset_at. + """ + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import _set_budget_reset_at + + data = UpdateTeamRequest(team_id="test-team", budget_duration="30d") + updated_kv = {"team_id": "test-team", "budget_duration": "30d"} + + _set_budget_reset_at(data, updated_kv) + + assert "budget_reset_at" in updated_kv + assert updated_kv["budget_reset_at"] is not None + + +@pytest.mark.asyncio +async def test_clear_team_member_budget_duration_calls_update_budget(): + """ + When team_member_budget_duration is explicitly null and a budget row + exists, clear_team_member_budget_fields should call update_budget + with budget_duration=None and budget_reset_at=None. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="admin-user", + ) + + team_table = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "budget-123"}, + members_with_roles=[], + ) + + updated_kv = { + "team_id": "test-team", + "team_member_budget_duration": None, + } + + with patch( + "litellm.proxy.management_endpoints.budget_management_endpoints.update_budget", + new_callable=AsyncMock, + ) as mock_update_budget: + result = await TeamMemberBudgetHandler.clear_team_member_budget_fields( + team_table=team_table, + user_api_key_dict=mock_user_api_key_dict, + updated_kv=updated_kv, + explicitly_set_fields={"team_member_budget_duration"}, + ) + + mock_update_budget.assert_awaited_once() + budget_request = mock_update_budget.call_args.kwargs["budget_obj"] + assert budget_request.budget_id == "budget-123" + assert "budget_duration" in budget_request.model_fields_set + assert budget_request.budget_duration is None + assert "budget_reset_at" in budget_request.model_fields_set + assert budget_request.budget_reset_at is None + assert "team_member_budget_duration" not in result + + +@pytest.mark.asyncio +async def test_clear_team_member_budget_clears_max_budget(): + """ + When team_member_budget is explicitly null, clear_team_member_budget_fields + should call update_budget with max_budget=None. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="admin-user", + ) + + team_table = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "budget-456"}, + members_with_roles=[], + ) + + updated_kv = { + "team_id": "test-team", + "team_member_budget": None, + } + + with patch( + "litellm.proxy.management_endpoints.budget_management_endpoints.update_budget", + new_callable=AsyncMock, + ) as mock_update_budget: + result = await TeamMemberBudgetHandler.clear_team_member_budget_fields( + team_table=team_table, + user_api_key_dict=mock_user_api_key_dict, + updated_kv=updated_kv, + explicitly_set_fields={"team_member_budget"}, + ) + + mock_update_budget.assert_awaited_once() + budget_request = mock_update_budget.call_args.kwargs["budget_obj"] + assert budget_request.budget_id == "budget-456" + assert "max_budget" in budget_request.model_fields_set + assert budget_request.max_budget is None + assert "team_member_budget" not in result + + +@pytest.mark.asyncio +async def test_clear_team_member_rpm_tpm_limits(): + """ + When team_member_rpm_limit and team_member_tpm_limit are explicitly null, + clear_team_member_budget_fields should clear both on the budget row. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="admin-user", + ) + + team_table = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "budget-789"}, + members_with_roles=[], + ) + + updated_kv = { + "team_id": "test-team", + "team_member_rpm_limit": None, + "team_member_tpm_limit": None, + } + + with patch( + "litellm.proxy.management_endpoints.budget_management_endpoints.update_budget", + new_callable=AsyncMock, + ) as mock_update_budget: + result = await TeamMemberBudgetHandler.clear_team_member_budget_fields( + team_table=team_table, + user_api_key_dict=mock_user_api_key_dict, + updated_kv=updated_kv, + explicitly_set_fields={"team_member_rpm_limit", "team_member_tpm_limit"}, + ) + + mock_update_budget.assert_awaited_once() + budget_request = mock_update_budget.call_args.kwargs["budget_obj"] + assert budget_request.budget_id == "budget-789" + assert "rpm_limit" in budget_request.model_fields_set + assert budget_request.rpm_limit is None + assert "tpm_limit" in budget_request.model_fields_set + assert budget_request.tpm_limit is None + assert "team_member_rpm_limit" not in result + assert "team_member_tpm_limit" not in result + + +@pytest.mark.asyncio +async def test_clear_all_team_member_fields_at_once(): + """ + When all team_member fields are explicitly null, all corresponding + budget row fields should be cleared in a single update. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="admin-user", + ) + + team_table = LiteLLM_TeamTable( + team_id="test-team", + metadata={"team_member_budget_id": "budget-all"}, + members_with_roles=[], + ) + + updated_kv = { + "team_id": "test-team", + "team_member_budget": None, + "team_member_budget_duration": None, + "team_member_rpm_limit": None, + "team_member_tpm_limit": None, + } + + all_fields = { + "team_member_budget", + "team_member_budget_duration", + "team_member_rpm_limit", + "team_member_tpm_limit", + } + + with patch( + "litellm.proxy.management_endpoints.budget_management_endpoints.update_budget", + new_callable=AsyncMock, + ) as mock_update_budget: + result = await TeamMemberBudgetHandler.clear_team_member_budget_fields( + team_table=team_table, + user_api_key_dict=mock_user_api_key_dict, + updated_kv=updated_kv, + explicitly_set_fields=all_fields, + ) + + mock_update_budget.assert_awaited_once() + budget_request = mock_update_budget.call_args.kwargs["budget_obj"] + assert budget_request.budget_id == "budget-all" + assert budget_request.max_budget is None + assert budget_request.budget_duration is None + assert budget_request.budget_reset_at is None + assert budget_request.rpm_limit is None + assert budget_request.tpm_limit is None + for field in all_fields: + assert field not in result + + +@pytest.mark.asyncio +async def test_team_member_budget_duration_not_sent_does_not_update(): + """ + When team_member_budget_duration is NOT sent in the request, no budget + update should occur and the field should not appear in updated_kv. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + updated_kv = {"team_id": "test-team", "max_budget": 200} + + _team_member_fields_in_request = { + field + for field in [ + "team_member_budget", + "team_member_rpm_limit", + "team_member_tpm_limit", + "team_member_budget_duration", + ] + if field in updated_kv + } + + assert len(_team_member_fields_in_request) == 0 + + TeamMemberBudgetHandler._clean_team_member_fields(updated_kv) + + assert "team_member_budget_duration" not in updated_kv + assert "team_member_budget" not in updated_kv + + +@pytest.mark.asyncio +async def test_clear_team_member_budget_fields_no_budget_row_skips_update(): + from litellm.proxy.management_endpoints.team_endpoints import ( + TeamMemberBudgetHandler, + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="admin-user", + ) + + team_table = LiteLLM_TeamTable( + team_id="test-team", + metadata=None, + members_with_roles=[], + ) + + updated_kv = { + "team_id": "test-team", + "team_member_budget": None, + "team_member_rpm_limit": None, + } + + with patch( + "litellm.proxy.management_endpoints.budget_management_endpoints.update_budget", + new_callable=AsyncMock, + ) as mock_update_budget: + result = await TeamMemberBudgetHandler.clear_team_member_budget_fields( + team_table=team_table, + user_api_key_dict=mock_user_api_key_dict, + updated_kv=updated_kv, + explicitly_set_fields={"team_member_budget", "team_member_rpm_limit"}, + ) + + mock_update_budget.assert_not_awaited() + assert "team_member_budget" not in result + assert "team_member_rpm_limit" not in result diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index c763e9c0e98..2efec3e0b34 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1777,6 +1777,23 @@ class TestHTMLIntegration: assert isinstance(html, str) assert len(html) > 0 + def test_success_page_instructs_manual_close_without_false_countdown(self): + """Browsers refuse window.close() on tabs they did not open via window.open() + (the CLI opens the page with webbrowser.open), so a 'closing in 3...' countdown + is a promise the browser usually can't keep and the page gets stuck on + 'Closing...'. The page must instead always show the manual-close instruction + and never advertise an auto-close that won't happen. + """ + from litellm.proxy.common_utils.html_forms.cli_sso_success import ( + render_cli_sso_success_page, + ) + + html = render_cli_sso_success_page() + + assert "You can now close this window and return to your terminal." in html + assert "Closing..." not in html + assert "This window will close in" not in html + class TestCustomUISSO: """Test the custom UI SSO sign-in handler functionality""" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py index 3c6af3e528a..401ea2ef589 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py @@ -683,36 +683,43 @@ class TestOpenAIPassthroughLoggingHandler: "litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload" ) @patch( - "litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler.OpenAIPassthroughLoggingHandler.get_provider_config" + "litellm.llms.openai.responses.transformation.OpenAIResponsesAPIConfig.transform_response_api_response" ) def test_responses_api_cost_tracking( - self, mock_get_provider_config, mock_get_standard_logging, mock_completion_cost + self, + mock_transform_responses, + mock_get_standard_logging, + mock_completion_cost, ): - """Test cost tracking for responses API route""" + """Test cost tracking for responses API route. + + Mocks the Responses-API transformer (the dedicated one this branch + of the handler dispatches into post-fix) so we can assert the + downstream cost-calculation contract without depending on the + real transformer's full behavior. + """ # Arrange mock_completion_cost.return_value = 0.000050 mock_get_standard_logging.return_value = {"test": "logging_payload"} - # Mock the provider config's transform_response to return a valid ModelResponse - from litellm import ModelResponse + # Mock the Responses transformer's return — a ResponsesAPIResponse + # carrying the usage fields downstream cost-calc expects. + from litellm.types.llms.openai import ResponsesAPIResponse - mock_model_response = ModelResponse( + mock_responses_api_response = ResponsesAPIResponse.model_construct( id="resp_abc123", + object="response", + created_at=1677652288, model="gpt-4o-2024-08-06", - choices=[ - { - "message": { - "role": "assistant", - "content": "Hello! How can I help you today?", - } - } - ], - usage={"prompt_tokens": 20, "completion_tokens": 15, "total_tokens": 35}, + status="completed", + output=[], + usage={ + "input_tokens": 20, + "output_tokens": 15, + "total_tokens": 35, + }, ) - - mock_provider_config = MagicMock() - mock_provider_config.transform_response.return_value = mock_model_response - mock_get_provider_config.return_value = mock_provider_config + mock_transform_responses.return_value = mock_responses_api_response # Mock responses API response mock_responses_response = { @@ -768,6 +775,109 @@ class TestOpenAIPassthroughLoggingHandler: assert mock_logging_obj.model_call_details["model"] == "gpt-4o" assert mock_logging_obj.model_call_details["custom_llm_provider"] == "openai" + @patch("litellm.completion_cost") + @patch( + "litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload" + ) + def test_responses_api_uses_responses_transformer_not_chat_completions( + self, mock_get_standard_logging, mock_completion_cost + ): + """Regression test for the Responses-API cost-tracking dispatch bug. + + BUG: the `elif is_responses:` branch in `openai_passthrough_handler` + was calling `OpenAIConfig.transform_response` (the chat-completions + transformer) on a Responses API payload. Chat-completions + transform_response expects `choices: [...]` in the raw response; + the Responses API uses `output: [...]` and `usage.input_tokens` / + `usage.output_tokens` (not `prompt_tokens` / `completion_tokens`). + The result was a KeyError 'choices' inside + `convert_to_model_response_object`, swallowed by the surrounding + try/except, and the SpendLogs row was written with zero tokens + and zero spend. + + FIX: use the dedicated `OpenAIResponsesAPIConfig.transform_response_api_response` + for the Responses branch. + + This test exercises the REAL transformer (no mocked + `get_provider_config`) so that running it against the un-fixed + handler raises and running it against the fixed handler succeeds. + """ + mock_completion_cost.return_value = 0.000050 + mock_get_standard_logging.return_value = {"test": "logging_payload"} + + # A real-shaped Azure / OpenAI Responses API payload — NO `choices`, + # uses `output` and `usage.input_tokens` / `usage.output_tokens`. + responses_api_body = { + "id": "resp_abc123", + "object": "response", + "created_at": 1677652288, + "model": "gpt-4o-2024-08-06", + "status": "completed", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "Hello!", + } + ], + } + ], + "usage": { + "input_tokens": 20, + "output_tokens": 15, + "total_tokens": 35, + }, + } + + mock_httpx_response = self._create_mock_httpx_response(responses_api_body) + mock_logging_obj = self._create_mock_logging_obj() + passthrough_payload = self._create_passthrough_logging_payload() + + kwargs = { + "passthrough_logging_payload": passthrough_payload, + "model": "gpt-4o", + "custom_llm_provider": "openai", + } + + result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler( + httpx_response=mock_httpx_response, + response_body=responses_api_body, + logging_obj=mock_logging_obj, + url_route="https://api.openai.com/v1/responses", + result="", + start_time=self.start_time, + end_time=self.end_time, + cache_hit=False, + request_body={"model": "gpt-4o", "input": "Tell me about AI"}, + **kwargs, + ) + + # Pre-fix this assertion fails — the handler swallows the + # KeyError raised by the chat-completions transformer and falls + # back to the passthrough_chat_handler which yields a different + # response_cost value. Post-fix, the Responses transformer + # succeeds and we get the mocked 0.000050. + assert result is not None + assert result["kwargs"]["response_cost"] == 0.000050 + assert result["kwargs"]["model"] == "gpt-4o" + + # `completion_cost` must be called with the responses call type + # and a `ResponsesAPIResponse` (not a `ModelResponse`). + mock_completion_cost.assert_called_once() + call_kwargs = mock_completion_cost.call_args[1] + assert call_kwargs["call_type"] == "responses" + + from litellm.types.llms.openai import ResponsesAPIResponse + + assert isinstance(call_kwargs["completion_response"], ResponsesAPIResponse), ( + "completion_response must be a ResponsesAPIResponse; passing a " + "chat-completions ModelResponse means the Responses transformer " + "isn't being used and we're back in the bug." + ) + class TestOpenAIPassthroughIntegration: """Integration tests for OpenAI passthrough cost tracking""" @@ -872,6 +982,126 @@ class TestOpenAIPassthroughIntegration: ) assert self.handler.is_openai_route("") == False + def test_is_supported_openai_endpoint_includes_responses_api(self): + """Regression test for the outer dispatch gate. + + `_is_supported_openai_endpoint` is the gate that decides whether the + OpenAI handler runs for a given URL. Before this gate accepted the + Responses API, calls to `/v1/responses` would fail the gate and the + handler's `elif is_responses:` branch was unreachable in the live + success-handler pipeline — every Responses-API call landed in + `LiteLLM_SpendLogs` with zero tokens / zero spend even though the + handler had a Responses branch internally. + + This test exercises the dispatch decision directly so future + refactors of `_is_supported_openai_endpoint` can't silently + remove Responses from the OR-chain without a test failure. + """ + # Responses must be supported on api.openai.com and openai.azure.com. + assert ( + self.handler._is_supported_openai_endpoint( + "https://api.openai.com/v1/responses" + ) + is True + ) + assert ( + self.handler._is_supported_openai_endpoint( + "https://openai.azure.com/v1/responses" + ) + is True + ) + # The other supported endpoints stay supported (no regression). + assert ( + self.handler._is_supported_openai_endpoint( + "https://api.openai.com/v1/chat/completions" + ) + is True + ) + assert ( + self.handler._is_supported_openai_endpoint( + "https://api.openai.com/v1/images/generations" + ) + is True + ) + assert ( + self.handler._is_supported_openai_endpoint( + "https://api.openai.com/v1/images/edits" + ) + is True + ) + # Unsupported OpenAI endpoints (e.g. /v1/models) still return False. + assert ( + self.handler._is_supported_openai_endpoint( + "https://api.openai.com/v1/models" + ) + is False + ) + + @patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler.OpenAIPassthroughLoggingHandler.openai_passthrough_handler" + ) + @pytest.mark.asyncio + async def test_success_handler_dispatches_responses_api_to_openai_handler( + self, mock_openai_handler + ): + """End-to-end dispatch test for the Responses API path. + + Pre-fix: `_is_supported_openai_endpoint` returned False for + `/v1/responses` URLs, so the OpenAI handler was never called. + This test would fail (mock never invoked) on the un-fixed + success_handler — passes only when the dispatch gate accepts + Responses URLs. + """ + mock_openai_handler.return_value = { + "result": {"id": "resp_abc123"}, + "kwargs": { + "response_cost": 0.0001, + "model": "gpt-4o", + "custom_llm_provider": "openai", + }, + } + + mock_httpx_response = MagicMock(spec=httpx.Response) + mock_httpx_response.text = ( + '{"id": "resp_abc123", "object": "response", ' + '"output": [], "usage": {"input_tokens": 5, "output_tokens": 3}}' + ) + + mock_logging_obj = AsyncMock() + mock_logging_obj.model_call_details = {} + mock_logging_obj.async_success_handler = AsyncMock() + + passthrough_payload = PassthroughStandardLoggingPayload( + url="https://api.openai.com/v1/responses", + request_body={"model": "gpt-4o", "input": "Hello"}, + request_method="POST", + ) + + await self.handler.pass_through_async_success_handler( + httpx_response=mock_httpx_response, + response_body={ + "id": "resp_abc123", + "object": "response", + "output": [], + "usage": {"input_tokens": 5, "output_tokens": 3}, + }, + logging_obj=mock_logging_obj, + url_route="https://api.openai.com/v1/responses", + result="", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={"model": "gpt-4o", "input": "Hello"}, + passthrough_logging_payload=passthrough_payload, + ) + + # The OpenAI handler MUST have been invoked. Pre-fix the dispatch + # gate filtered Responses URLs out and the mock was never called. + mock_openai_handler.assert_called_once() + # And we can verify it was dispatched with the Responses URL. + call_kwargs = mock_openai_handler.call_args.kwargs + assert call_kwargs["url_route"] == "https://api.openai.com/v1/responses" + @patch( "litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler.OpenAIPassthroughLoggingHandler.openai_passthrough_handler" ) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index b75fc27e21d..4eab1a4bf61 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -24,11 +24,13 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( get_vertex_base_url, llm_passthrough_factory_proxy_route, milvus_proxy_route, + mistral_proxy_route, openai_proxy_route, vertex_discovery_proxy_route, vertex_proxy_route, vllm_proxy_route, ) +from litellm.proxy._types import UserAPIKeyAuth from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials @@ -1092,9 +1094,9 @@ class TestVertexAIPassThroughHandler: assert result is not None assert result["result"] is not None - assert result["kwargs"].get("custom_llm_provider") == "gemini", ( - "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" - ) + assert ( + result["kwargs"].get("custom_llm_provider") == "gemini" + ), "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" assert result["kwargs"].get("model") == "gemini-embedding-2-preview" mock_completion_cost.assert_called_once() @@ -1261,6 +1263,78 @@ async def test_is_streaming_request_fn(): assert await is_streaming_request_fn(mock_request) is True +@pytest.mark.asyncio +async def test_mistral_passthrough_accepts_multipart_without_json_parsing(): + boundary = "----litellm-test-boundary" + body = ( + f"--{boundary}\r\n" + 'Content-Disposition: form-data; name="purpose"\r\n\r\n' + "ocr\r\n" + f"--{boundary}\r\n" + 'Content-Disposition: form-data; name="file"; filename="document.pdf"\r\n' + "Content-Type: application/pdf\r\n\r\n" + "%PDF-1.4 test\r\n" + f"--{boundary}--\r\n" + ).encode("utf-8") + + async def receive(): + return { + "type": "http.request", + "body": body, + "more_body": False, + } + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/mistral/v1/files", + "headers": [ + ( + b"content-type", + f"multipart/form-data; boundary={boundary}".encode("utf-8"), + ) + ], + "query_string": b"", + }, + receive=receive, + ) + + captured_kwargs = {} + + async def fake_endpoint(request, fastapi_response, user_api_key_dict): + return {"ok": True} + + def fake_create_pass_through_route(**kwargs): + captured_kwargs.update(kwargs) + return fake_endpoint + + user_api_key_dict = UserAPIKeyAuth(token="test-key") + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="mistral-test-key", + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route", + side_effect=fake_create_pass_through_route, + ), + ): + response = await mistral_proxy_route( + endpoint="v1/files", + request=request, + fastapi_response=Response(), + user_api_key_dict=user_api_key_dict, + ) + + assert response == {"ok": True} + assert captured_kwargs["is_streaming_request"] is False + assert captured_kwargs["custom_headers"] == { + "Authorization": "Bearer mistral-test-key" + } + + class TestBedrockLLMProxyRoute: @pytest.mark.asyncio async def test_bedrock_llm_proxy_route_application_inference_profile(self): diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 164538a2757..677d358428d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -601,6 +601,98 @@ async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch): await pc.load_config(router=None, config_file_path="/no/file.yaml") +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_forwards_callback_specific_params( + tmp_path, monkeypatch +): + """Regression: callback_settings from config must be forwarded to + initialize_callbacks_on_proxy as callback_specific_params. + + Callbacks like DatadogCostManagementLogger read their init params (e.g. + cost_tag_keys) from callback_specific_params[]. If the + argument is dropped at the call site, they silently initialize with empty + params and the configured allowlist never takes effect. + """ + f = tmp_path / "c.yaml" + f.write_text( + "model_list: []\n" + "general_settings: {}\n" + "callback_settings:\n" + " datadog_cost_management:\n" + " cost_tag_keys:\n" + " - capability\n" + " - platform\n" + " - ai_product\n" + "litellm_settings:\n" + ' callbacks: ["datadog_cost_management"]\n' + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + captured = {} + + def _fake_initialize_callbacks_on_proxy(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr( + "litellm.proxy.proxy_server.initialize_callbacks_on_proxy", + _fake_initialize_callbacks_on_proxy, + ) + + pc = ProxyConfig() + await pc.load_config(router=None, config_file_path=str(f)) + + # The callbacks branch must forward the loaded callback_settings. + assert captured.get("callback_specific_params") == { + "datadog_cost_management": { + "cost_tag_keys": ["capability", "platform", "ai_product"] + } + } + + +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash( + tmp_path, monkeypatch +): + """Regression: `callback_settings:` with no body loads as None because + dict.get() only falls back to the default when the key is absent. The None + was forwarded verbatim to initialize_callbacks_on_proxy, where the first + `"" in callback_specific_params` membership test raised + TypeError: argument of type 'NoneType' is not iterable, aborting startup. + Startup must succeed and the callback must initialize with its defaults. + """ + f = tmp_path / "c.yaml" + f.write_text( + "model_list: []\n" + "general_settings: {}\n" + "callback_settings:\n" + "litellm_settings:\n" + ' callbacks: ["compression_interception"]\n' + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, + ) + + original_callbacks = ( + list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] + ) + litellm.callbacks = [] + try: + pc = ProxyConfig() + await pc.load_config(router=None, config_file_path=str(f)) + + assert any( + isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks + ) + finally: + litellm.callbacks = original_callbacks + + # --------------------------------------------------------------------------- # ProxyConfig._init_non_llm_configs # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index ec8b06d9c97..4e5f13fdf88 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -192,8 +192,9 @@ async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set async def test_reconcile_budget_reservation_for_counter_update_failure_invalidates( monkeypatch, ): - """Reservation reconcile raising must invalidate reserved counters but - not propagate the exception.""" + """Reservation reconcile raising must invalidate reserved counters, swallow + the exception, and return an empty set so the caller falls back to the + direct spend-counter increment instead of skipping it.""" import litellm.proxy.spend_tracking.budget_reservation as br monkeypatch.setattr( @@ -213,7 +214,7 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat budget_reservation={"foo": "bar"}, response_cost=1.0 ) - assert result == {"spend:key:abc"} + assert result == set() assert fake_invalidate.called is True diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py new file mode 100644 index 00000000000..c123eeeed36 --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py @@ -0,0 +1,87 @@ +""" +Regression test for enforced-spend underreporting when Redis fails during the +budget-reservation reconcile step of ``increment_spend_counters``. + +Production failure mode: a managed Redis returns an intermittent timeout on the +reconcile increment. Reconcile deletes (invalidates) the shared counter and +gives up, but ``increment_spend_counters`` still treats the counter as +"already reconciled" and skips the direct increment. The actual call cost never +lands in the enforced counter, so budgets stop gating until the next cold +reseed pulls a lagging value from the DB. + +The fix makes the reconcile path fall back to the direct increment when it +fails, so the actual cost is always written to the shared counter. +""" + +import pytest + +from litellm.caching import DualCache +from litellm.proxy import proxy_server + + +class _FlakyRedisCache: + def __init__(self) -> None: + self._store: dict = {} + self._increment_calls = 0 + + async def async_increment(self, key, value, **kwargs): + self._increment_calls += 1 + if self._increment_calls == 1: + raise Exception("Redis timeout") + self._store[key] = float(self._store.get(key, 0.0)) + float(value) + return self._store[key] + + async def async_get_cache(self, key, *args, **kwargs): + return self._store.get(key) + + async def async_delete_cache(self, key, *args, **kwargs): + self._store.pop(key, None) + + async def async_set_cache(self, key, value, *args, **kwargs): + self._store[key] = float(value) + return True + + +@pytest.mark.asyncio +async def test_direct_increment_runs_when_reservation_reconcile_hits_redis_failure( + monkeypatch, +): + hashed_token = "hashed_test_token" + counter_key = f"spend:key:{hashed_token}" + reserved_cost = 0.5 + response_cost = 1.0 + + flaky_redis = _FlakyRedisCache() + flaky_redis._store[counter_key] = reserved_cost + + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + monkeypatch.setattr(proxy_server.spend_counter_cache, "redis_cache", flaky_redis) + proxy_server.spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, value=reserved_cost + ) + + budget_reservation = { + "reserved_cost": reserved_cost, + "finalized": False, + "entries": [ + { + "counter_key": counter_key, + "entity_type": "Key", + "entity_id": hashed_token, + "reserved_cost": reserved_cost, + "applied_adjustment": 0.0, + } + ], + } + + await proxy_server.increment_spend_counters( + token=hashed_token, + team_id=None, + user_id=None, + response_cost=response_cost, + budget_reservation=budget_reservation, + ) + + enforced_spend = await flaky_redis.async_get_cache(key=counter_key) + assert enforced_spend == response_cost diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index aef91ed3c77..2632d8af4f1 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -1314,7 +1314,8 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): assert session_id == "session-123" assert page_size == 1 assert skip == 1 # page=2, page_size=1 - return [mock_spend_logs[1]] + assert 'ORDER BY "startTime" DESC' in sql_query + return [mock_spend_logs[0]] class MockPrismaClient: def __init__(self): @@ -1337,7 +1338,7 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): assert data["page_size"] == 1 assert data["total_pages"] == 2 assert len(data["data"]) == 1 - assert data["data"][0]["request_id"] == "req2" + assert data["data"][0]["request_id"] == "req1" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 0f5a0cbe4b6..b45b31cc67c 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2269,6 +2269,36 @@ class TestHandleLLMApiExceptionDictDetail: assert proxy_exc.message == "Content blocked by guardrail" assert proxy_exc.provider_specific_fields is None + async def test_not_found_error_preserves_404(self): + """NotFoundError with status_code=404 should map to ProxyException code=404.""" + from litellm.exceptions import NotFoundError + + exc = NotFoundError( + message="Model gemini-3.1-flash-lite-preview not found", + model="gemini-3.1-flash-lite-preview", + llm_provider="gemini", + ) + proxy_exc = await self._invoke(exc) + assert proxy_exc.code == "404" + assert "NotFoundError" in proxy_exc.message + + async def test_exception_with_status_code_propagates(self): + """Exception with a statically-set status_code should propagate it.""" + from litellm.llms.vertex_ai.common_utils import VertexAIError + + exc = VertexAIError( + status_code=429, + message="Rate limit exceeded", + ) + proxy_exc = await self._invoke(exc) + assert proxy_exc.code == "429" + + async def test_exception_without_status_code_defaults_to_500(self): + """Exception with no status_code attribute defaults to 500.""" + exc = ValueError("Something broke") + proxy_exc = await self._invoke(exc) + assert proxy_exc.code == "500" + class TestAsyncStreamingDataGeneratorFastPath: """Fast/slow path branching in async_streaming_data_generator.""" diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 4fb725b7ef3..34c88e2fd33 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -795,6 +795,127 @@ class TestProxyInitializationHelpers: assert appended_params["pgbouncer"] == "true" assert appended_params["statement_cache_size"] == 0 + def test_build_db_connection_url_params_disable_prepared_statements(self): + from litellm.proxy.proxy_cli import _build_db_connection_url_params + + params = _build_db_connection_url_params( + connection_limit=10, + pool_timeout=60, + disable_prepared_statements=True, + ) + assert params["pgbouncer"] == "true" + + def test_build_db_connection_url_params_no_pgbouncer_by_default(self): + from litellm.proxy.proxy_cli import _build_db_connection_url_params + + params = _build_db_connection_url_params( + connection_limit=10, + pool_timeout=60, + ) + assert "pgbouncer" not in params + + def test_build_db_connection_url_params_extra_pgbouncer_overrides_flag(self): + from litellm.proxy.proxy_cli import _build_db_connection_url_params + + params = _build_db_connection_url_params( + connection_limit=10, + pool_timeout=60, + disable_prepared_statements=True, + extra_params={"pgbouncer": "false"}, + ) + assert params["pgbouncer"] == "false" + + @pytest.mark.parametrize( + "config_value, expect_pgbouncer", + [ + (True, True), + (False, False), + ("true", True), + ("false", False), + ("not-a-bool", False), + ], + ) + @patch("subprocess.run") + @patch("atexit.register") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch( + "litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False + ) + def test_disable_prepared_statements_forwarded_to_url( + self, + mock_should_update, + mock_setup_db, + mock_atexit_register, + mock_subprocess_run, + config_value, + expect_pgbouncer, + ): + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + runner = CliRunner() + mock_subprocess_run.return_value = MagicMock(returncode=0) + + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock( + return_value={ + "general_settings": { + "database_url": "postgresql://test:test@localhost:5432/test", + "database_disable_prepared_statements": config_value, + } + } + ) + + clean_env = { + k: v + for k, v in os.environ.items() + if k not in ("DATABASE_URL", "DIRECT_URL") + } + + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + { + "proxy_server": mock_proxy_module, + "litellm.proxy.proxy_server": mock_proxy_module, + }, + ), + patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args, + patch( + "litellm.proxy.proxy_cli.append_query_params", + side_effect=lambda url, params: str(url), + ) as mock_append_query_params, + ): + mock_get_args.return_value = { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + + result = runner.invoke( + run_server, + ["--local", "--config", "test-config.yaml", "--skip_server_startup"], + ) + + assert ( + result.exit_code == 0 + ), f"exit_code={result.exit_code}, output={result.output}" + mock_append_query_params.assert_called() + appended_params = mock_append_query_params.call_args.args[1] + if expect_pgbouncer: + assert appended_params["pgbouncer"] == "true" + else: + assert "pgbouncer" not in appended_params + @patch("uvicorn.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 8aa839cdfcb..9f2c5ffd615 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -2319,6 +2319,36 @@ async def test_custom_ui_sso_sign_in_handler_config_loading(): os.unlink(config_file_path) +@pytest.mark.asyncio +async def test_load_config_max_budget_env_var_coerced_to_float(tmp_path, monkeypatch): + """ + max_budget configured as os.environ/MAX_BUDGET resolves to a string; + load_config must coerce it to float so the startup check + `litellm.max_budget > 0` doesn't raise TypeError. + """ + from litellm.proxy.proxy_server import ProxyConfig + + monkeypatch.setenv("MAX_BUDGET", "10") + test_config = { + "model_list": [], + "litellm_settings": {"max_budget": "os.environ/MAX_BUDGET"}, + } + config_file = tmp_path / "config.yaml" + config_file.write_text(yaml.dump(test_config)) + + original_max_budget = litellm.max_budget + try: + proxy_config = ProxyConfig() + await proxy_config.load_config( + router=MagicMock(), config_file_path=str(config_file) + ) + assert isinstance(litellm.max_budget, float) + assert litellm.max_budget == 10.0 + assert litellm.max_budget > 0 + finally: + litellm.max_budget = original_max_budget + + @pytest.mark.asyncio async def test_load_environment_variables_direct_and_os_environ(): """ @@ -4294,6 +4324,111 @@ async def test_init_sso_settings_in_db_empty_settings(): assert uppercased_settings == {} +@pytest.mark.asyncio +async def test_init_sso_settings_in_db_retries_on_transport_error(): + """`_init_sso_settings_in_db` self-heals across one ClientNotConnectedError + via call_with_db_reconnect_retry — mirrors the auth-path behavior so + startup/reload bursts don't spam the log.""" + import prisma + + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + mock_sso_config = MagicMock() + mock_sso_config.sso_settings = {"GOOGLE_CLIENT_ID": "xxx"} + + invocations: list = [] + + async def _flaky_find_unique(**kwargs): + invocations.append(None) + if len(invocations) == 1: + raise prisma.errors.ClientNotConnectedError() + return mock_sso_config + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( + side_effect=_flaky_find_unique + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + + with patch.object( + proxy_config, "_decrypt_and_set_db_env_variables" + ) as mock_decrypt: + await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) + + assert len(invocations) == 2 + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs + assert reconnect_kwargs["reason"] == "init_sso_settings_in_db_lookup_failure" + mock_decrypt.assert_called_once() + + +@pytest.mark.asyncio +async def test_init_sso_settings_in_db_propagates_when_reconnect_fails(): + """When reconnect returns False (cooldown / lock contention), the original + ClientNotConnectedError is caught by the function's `except Exception` and + logged — no retry storm, no crash.""" + import prisma + + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock( + side_effect=prisma.errors.ClientNotConnectedError() + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + + # Should NOT raise — the function's own try/except swallows the propagated error. + await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client) + + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_init_hashicorp_vault_config_override_retries_on_transport_error(): + """`_init_hashicorp_vault_config_override` self-heals across one + ClientNotConnectedError via call_with_db_reconnect_retry.""" + import prisma + + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + proxy_config._last_hashicorp_vault_config = None + + invocations: list = [] + + async def _flaky_find_unique(**kwargs): + invocations.append(None) + if len(invocations) == 1: + raise prisma.errors.ClientNotConnectedError() + return None # No config in DB → function returns early after retry. + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_configoverrides.find_unique = AsyncMock( + side_effect=_flaky_find_unique + ) + mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 + mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 + + await proxy_config._init_hashicorp_vault_config_override( + prisma_client=mock_prisma_client + ) + + assert len(invocations) == 2 + mock_prisma_client.attempt_db_reconnect.assert_awaited_once() + reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs + assert ( + reconnect_kwargs["reason"] + == "init_hashicorp_vault_config_override_lookup_failure" + ) + + def test_update_config_fields_uppercases_env_vars(monkeypatch): """ Ensure environment variables pulled from DB are uppercased when applied so @@ -6474,7 +6609,12 @@ async def test_increment_spend_counters_finalizes_none_cost_reservation(): @pytest.mark.asyncio -async def test_increment_spend_counters_invalidates_bad_reserved_counter_without_failing(): +async def test_increment_spend_counters_falls_back_to_direct_increment_on_bad_reserved_counter(): + """When the reservation reconcile fails, the reserved counters are + invalidated and the actual response cost must still be written via the + direct increment fallback. Leaving the counter at ``None`` lets the next + request reseed a stale value from the DB and silently stops budget gating, + which is the bug this fix addresses.""" from litellm.caching.dual_cache import DualCache from litellm.proxy.proxy_server import increment_spend_counters @@ -6515,7 +6655,7 @@ async def test_increment_spend_counters_invalidates_bad_reserved_counter_without counter_cache.in_memory_cache.get_cache( key="spend:key:key-bad-reserved-counter" ) - is None + == 0.25 ) finally: ps.spend_counter_cache = orig_counter diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py index 7e7e98d1360..437984d9273 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py @@ -193,6 +193,7 @@ async def test_query_first_with_cached_plan_fallback_happy_returns_row( ) -> None: expected = {"token": "abc", "team_spend": 1.0, "team_max_budget": 5.0} prisma_client.db.query_first = AsyncMock(return_value=expected) + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) result = await prisma_client._query_first_with_cached_plan_fallback( "SELECT * FROM x WHERE token = $1", "abc" ) @@ -208,35 +209,110 @@ async def test_query_first_with_cached_plan_fallback_happy_returns_row( "args": ("SELECT * FROM x WHERE token = $1", "abc"), "matches": True, } + prisma_client.attempt_db_reconnect.assert_not_awaited() @pytest.mark.asyncio -async def test_query_first_with_cached_plan_fallback_retries_on_cached_plan_error( +async def test_query_first_with_cached_plan_fallback_reconnects_then_retries_identical_query( prisma_client: PrismaClient, ) -> None: + original_query = 'SELECT * FROM "LiteLLM_VerificationToken" WHERE v.token = $1' expected = {"token": "abc", "team_spend": 1.0, "team_max_budget": 5.0} + manager = MagicMock() + query_first = AsyncMock( + side_effect=[ + RuntimeError("cached plan must not change result type"), + expected, + ] + ) + reconnect = AsyncMock(return_value=True) + manager.attach_mock(query_first, "query_first") + manager.attach_mock(reconnect, "attempt_db_reconnect") + prisma_client.db.query_first = query_first + prisma_client.attempt_db_reconnect = reconnect + + result = await prisma_client._query_first_with_cached_plan_fallback( + original_query, "abc" + ) + + assert result == expected + assert query_first.await_count == 2 + first_call, retry_call = query_first.await_args_list + assert retry_call.args == first_call.args == (original_query, "abc") + reconnect.assert_awaited_once() + assert reconnect.await_args.kwargs.get("force", False) is False + assert [name for name, *_ in manager.mock_calls] == [ + "query_first", + "attempt_db_reconnect", + "query_first", + ] + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_never_deallocates( + prisma_client: PrismaClient, +) -> None: + expected = {"token": "abc"} prisma_client.db.query_first = AsyncMock( side_effect=[ RuntimeError("cached plan must not change result type"), expected, ] ) - result = await prisma_client._query_first_with_cached_plan_fallback( - "SELECT * FROM x WHERE token = $1", "abc" + prisma_client.db.execute_raw = AsyncMock(return_value=0) + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") + + prisma_client.db.execute_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_propagates_when_retry_also_fails( + prisma_client: PrismaClient, +) -> None: + plan_error = RuntimeError("cached plan must not change result type") + prisma_client.db.query_first = AsyncMock(side_effect=[plan_error, plan_error]) + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + with pytest.raises(RuntimeError, match="cached plan must not change result type"): + await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") + + assert prisma_client.db.query_first.await_count == 2 + prisma_client.attempt_db_reconnect.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_retries_when_reconnect_returns_false( + prisma_client: PrismaClient, +) -> None: + expected = {"token": "abc"} + prisma_client.db.query_first = AsyncMock( + side_effect=[ + RuntimeError("cached plan must not change result type"), + expected, + ] ) + prisma_client.attempt_db_reconnect = AsyncMock(return_value=False) + + result = await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") + assert result == expected assert prisma_client.db.query_first.await_count == 2 - second_call_sql = prisma_client.db.query_first.await_args_list[1].args[0] - assert "cache_invalidated_" in second_call_sql @pytest.mark.asyncio async def test_query_first_with_cached_plan_fallback_reraises_non_plan_errors( prisma_client: PrismaClient, ) -> None: - prisma_client.db.query_first = AsyncMock(side_effect=RuntimeError("totally unrelated")) + prisma_client.db.query_first = AsyncMock( + side_effect=RuntimeError("totally unrelated") + ) + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) with pytest.raises(RuntimeError, match="totally unrelated"): await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") + assert prisma_client.db.query_first.await_count == 1 + prisma_client.attempt_db_reconnect.assert_not_awaited() @pytest.mark.asyncio @@ -351,7 +427,9 @@ async def test_get_data_token_find_unique_returns_record( async def test_get_data_token_find_unique_missing_token_raises_401( prisma_client: PrismaClient, ) -> None: - prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=None + ) with pytest.raises(HTTPException) as excinfo: await prisma_client.get_data(token="sk-missing", table_name="key") err = excinfo.value diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py index 07d894d0400..510dcf77afd 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py @@ -1471,3 +1471,191 @@ async def test_affinity_does_not_raise_when_boundary_peer_available(): assert result == [peer] assert request_kwargs.get("_encrypted_content_affinity_pinned") is True + + +@pytest.mark.asyncio +async def test_model_group_affinity_config_enables_encrypted_content_affinity(): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + target_deployment = { + "model_name": model_group, + "litellm_params": {"model": "openai/gpt-5.1-codex"}, + "model_info": {"id": "deployment-b"}, + } + healthy_deployments = [ + { + "model_name": model_group, + "litellm_params": {"model": "openai/gpt-5.1-codex"}, + "model_info": {"id": "deployment-a"}, + }, + target_deployment, + ] + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( + "deployment-b", "rs_test" + ) + request_kwargs = { + "input": [{"type": "reasoning", "id": encoded_id}], + "litellm_metadata": {}, + } + check = EncryptedContentAffinityCheck( + enable_global_affinity=False, + model_group_affinity_config={ + model_group: ["encrypted_content_affinity"], + }, + ) + + filtered = await check.async_filter_deployments( + model=model_group, + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs=request_kwargs, + ) + + assert filtered == [target_deployment] + assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] + assert request_kwargs.get("_encrypted_content_affinity_pinned") is True + + +@pytest.mark.asyncio +async def test_model_group_affinity_config_does_not_disable_global_encrypted_content_affinity(): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + target_deployment = { + "model_name": model_group, + "litellm_params": {"model": "openai/gpt-5.1-codex"}, + "model_info": {"id": "deployment-b"}, + } + healthy_deployments = [ + { + "model_name": model_group, + "litellm_params": {"model": "openai/gpt-5.1-codex"}, + "model_info": {"id": "deployment-a"}, + }, + target_deployment, + ] + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( + "deployment-b", "rs_test" + ) + request_kwargs = { + "input": [{"type": "reasoning", "id": encoded_id}], + "litellm_metadata": {}, + } + check = EncryptedContentAffinityCheck( + enable_global_affinity=True, + model_group_affinity_config={ + model_group: ["deployment_affinity"], + }, + ) + + filtered = await check.async_filter_deployments( + model=model_group, + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs=request_kwargs, + ) + + assert filtered == [target_deployment] + assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] + assert request_kwargs.get("_encrypted_content_affinity_pinned") is True + + +@pytest.mark.asyncio +async def test_model_group_encrypted_content_affinity_overrides_global_deployment_affinity(): + from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( + DeploymentAffinityCheck, + ) + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + user_api_key_hash = "test-user-key" + deployment_a = { + "model_name": model_group, + "litellm_params": { + "model": "openai/gpt-5.1-codex", + "api_key": "mock-api-key-a", + }, + "model_info": {"id": "deployment-a"}, + } + deployment_b = { + "model_name": model_group, + "litellm_params": { + "model": "openai/gpt-5.1-codex", + "api_key": "mock-api-key-b", + }, + "model_info": {"id": "deployment-b"}, + } + router = litellm.Router( + model_list=[deployment_a, deployment_b], + optional_pre_call_checks=["deployment_affinity"], + model_group_affinity_config={ + model_group: ["encrypted_content_affinity"], + }, + num_retries=0, + ) + + try: + callbacks = router.optional_callbacks or [] + deployment_callback = next( + cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) + ) + encrypted_content_callback = next( + cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) + ) + assert callbacks.index(encrypted_content_callback) < callbacks.index( + deployment_callback + ) + assert encrypted_content_callback.enable_global_affinity is False + + cache_key = DeploymentAffinityCheck.get_affinity_cache_key( + model_group=model_group, + user_key=user_api_key_hash, + ) + await deployment_callback.cache.async_set_cache( + key=cache_key, + value={"model_id": "deployment-a"}, + ttl=60, + ) + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( + "deployment-b", "rs_test" + ) + request_kwargs = { + "input": [ + { + "type": "reasoning", + "id": encoded_id, + "encrypted_content": "gAAAAABpnW_yEYmSNEyOG...", + } + ], + "metadata": {"user_api_key_hash": user_api_key_hash}, + "litellm_metadata": {}, + } + + after_deployment_affinity = await deployment_callback.async_filter_deployments( + model=model_group, + healthy_deployments=[deployment_a, deployment_b], + messages=None, + request_kwargs=request_kwargs, + ) + assert after_deployment_affinity == [deployment_a, deployment_b] + + after_encrypted_content_affinity = ( + await encrypted_content_callback.async_filter_deployments( + model=model_group, + healthy_deployments=after_deployment_affinity, + messages=None, + request_kwargs=request_kwargs, + ) + ) + + assert after_encrypted_content_affinity == [deployment_b] + assert request_kwargs.get("_encrypted_content_affinity_pinned") is True + finally: + router.discard() diff --git a/tests/test_litellm/test_anthropic_beta_headers_filtering.py b/tests/test_litellm/test_anthropic_beta_headers_filtering.py index 84867a6e905..9cd27e88c33 100644 --- a/tests/test_litellm/test_anthropic_beta_headers_filtering.py +++ b/tests/test_litellm/test_anthropic_beta_headers_filtering.py @@ -402,6 +402,16 @@ class TestAnthropicBetaHeadersFiltering: test_case["expected"] in filtered ), f"Header '{test_case['input']}' should be mapped to '{test_case['expected']}' for {test_case['provider']}, but got: {filtered}" + def test_filter_and_transform_beta_headers_vertex_ai_keeps_compact(self): + """Vertex AI supports compact context edits, so the compact beta header + must be forwarded instead of stripped (it was previously mapped to null, + which broke compact_20260112 context edits over /v1/messages).""" + filtered = filter_and_transform_beta_headers( + beta_headers=["compact-2026-01-12"], provider="vertex_ai" + ) + + assert filtered == ["compact-2026-01-12"] + def test_null_value_headers_filtered(self): """Test that headers with null values are always filtered out.""" for provider in [ diff --git a/tests/test_litellm/test_claude_fable_5_config.py b/tests/test_litellm/test_claude_fable_5_config.py new file mode 100644 index 00000000000..d8d95fba0da --- /dev/null +++ b/tests/test_litellm/test_claude_fable_5_config.py @@ -0,0 +1,230 @@ +""" +Validate Claude Fable 5 model configuration entries. + +Fable 5 is a new tier above Opus ($10/$50 per MTok) with the same adaptive-only +API surface as Opus 4.7/4.8. The cost-map entries below are what make the model +resolvable across Anthropic, Bedrock, Vertex AI, and Azure AI (Microsoft +Foundry), and the ``supports_adaptive_thinking`` flag is what makes LiteLLM send +``thinking.type='adaptive'`` instead of the legacy ``enabled``/``budget_tokens`` +shape, which Fable 5 rejects with a 400. +""" + +import json +import os + +import pytest + +import litellm +from litellm.constants import BEDROCK_CONVERSE_MODELS +from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap + +REPO_ROOT = os.path.join(os.path.dirname(__file__), "../..") + + +def _load_root_cost_map() -> dict: + json_path = os.path.join(REPO_ROOT, "model_prices_and_context_window.json") + with open(json_path) as f: + return json.load(f) + + +@pytest.fixture +def local_model_cost_map(monkeypatch): + """Force the bundled backup cost map so assertions don't depend on the + network-fetched ``main`` copy (which lags this branch until merge).""" + original_model_cost = litellm.model_cost + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.get_model_info.cache_clear() + try: + yield + finally: + litellm.model_cost = original_model_cost + litellm.get_model_info.cache_clear() + + +def test_fable_5_model_pricing_and_capabilities(): + model_data = _load_root_cost_map() + + expected_models = [ + ("claude-fable-5", "anthropic"), + ("anthropic.claude-fable-5", "bedrock_converse"), + ("vertex_ai/claude-fable-5", "vertex_ai-anthropic_models"), + # Unlike Opus 4.8 (200k on Foundry), Fable 5 has the full 1M context + # window on Microsoft Foundry. + ("azure_ai/claude-fable-5", "azure_ai"), + ] + + for model_name, provider in expected_models: + assert model_name in model_data, f"Missing model entry: {model_name}" + info = model_data[model_name] + + assert info["litellm_provider"] == provider + assert info["mode"] == "chat" + assert info["max_input_tokens"] == 1000000 + assert info["max_output_tokens"] == 128000 + assert info["max_tokens"] == 128000 + + # $10 / $50 per MTok (2x Opus 4.8), with the standard 1.25x 5m + # cache-write, 2x 1h cache-write, and 0.1x cache-read multipliers. + assert info["input_cost_per_token"] == 1e-05 + assert info["output_cost_per_token"] == 5e-05 + assert info["cache_creation_input_token_cost"] == 1.25e-05 + assert info["cache_creation_input_token_cost_above_1hr"] == 2e-05 + assert info["cache_read_input_token_cost"] == 1e-06 + + # Flat-rate across the full 1M context window. + assert "input_cost_per_token_above_200k_tokens" not in info + assert "output_cost_per_token_above_200k_tokens" not in info + + assert info["supports_assistant_prefill"] is False + assert info["supports_function_calling"] is True + assert info["supports_prompt_caching"] is True + assert info["supports_reasoning"] is True + assert info["supports_tool_choice"] is True + assert info["supports_vision"] is True + assert info["supports_xhigh_reasoning_effort"] is True + assert info["supports_max_reasoning_effort"] is True + + +def test_fable_5_bedrock_regional_model_pricing(): + model_data = _load_root_cost_map() + + # Fable 5 launched with us/eu geo inference profiles plus a global profile + # (no au/apac/jp). Global uses base pricing; geo profiles carry the + # standard 10% regional premium. + expected_models = { + "global.anthropic.claude-fable-5": { + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "cache_creation_input_token_cost": 1.25e-05, + "cache_read_input_token_cost": 1e-06, + }, + "us.anthropic.claude-fable-5": { + "input_cost_per_token": 1.1e-05, + "output_cost_per_token": 5.5e-05, + "cache_creation_input_token_cost": 1.375e-05, + "cache_read_input_token_cost": 1.1e-06, + }, + "eu.anthropic.claude-fable-5": { + "input_cost_per_token": 1.1e-05, + "output_cost_per_token": 5.5e-05, + "cache_creation_input_token_cost": 1.375e-05, + "cache_read_input_token_cost": 1.1e-06, + }, + } + + for model_name, expected in expected_models.items(): + assert model_name in model_data, f"Missing model entry: {model_name}" + info = model_data[model_name] + assert info["litellm_provider"] == "bedrock_converse" + assert info["max_input_tokens"] == 1000000 + assert info["max_output_tokens"] == 128000 + assert info["bedrock_output_config_effort_ceiling"] == "xhigh" + for key, value in expected.items(): + assert info[key] == value + + +def test_fable_5_geo_multiplier_without_fast_mode(): + """First-party ``inference_geo='us'`` carries the 1.1x premium, but unlike + the Opus line there is no fast-mode variant for Fable 5; a ``fast`` key + here would silently misprice ``speed='fast'`` requests.""" + model_data = _load_root_cost_map() + entry = model_data["claude-fable-5"]["provider_specific_entry"] + assert entry == {"us": 1.1} + + +def test_fable_5_present_in_bundled_backup(): + """The bundled backup is the runtime fallback (and what tests load with + ``LITELLM_LOCAL_MODEL_COST_MAP=True``) — it must carry the same entries as + the root cost map, otherwise the model resolves on one path but not the + other.""" + backup = GetModelCostMap.load_local_model_cost_map() + root = _load_root_cost_map() + for model_name in ( + "claude-fable-5", + "anthropic.claude-fable-5", + "global.anthropic.claude-fable-5", + "us.anthropic.claude-fable-5", + "eu.anthropic.claude-fable-5", + "vertex_ai/claude-fable-5", + "vertex_ai/claude-fable-5@default", + "azure_ai/claude-fable-5", + ): + assert model_name in backup, f"Missing from backup cost map: {model_name}" + assert backup[model_name] == root[model_name], model_name + + +def test_fable_5_registered_for_bedrock_converse(): + assert "anthropic.claude-fable-5" in BEDROCK_CONVERSE_MODELS + + +def test_fable_5_provider_resolves_via_model_info(local_model_cost_map): + info = litellm.get_model_info(model="claude-fable-5") + assert info["litellm_provider"] == "anthropic" + assert info["max_input_tokens"] == 1000000 + assert info["max_output_tokens"] == 128000 + + +@pytest.mark.parametrize( + "cost_map", + [_load_root_cost_map(), GetModelCostMap.load_local_model_cost_map()], + ids=["root", "bundled_backup"], +) +def test_fable_5_all_variants_carry_adaptive_thinking_flag(cost_map): + """Every Fable 5 entry must advertise ``supports_adaptive_thinking``. + + Adaptive-thinking detection is cost-map driven, so a single variant missing + the flag silently sends the legacy ``thinking.type='enabled'`` shape and the + provider 400s (issue #29188 for the Opus 4.8 equivalent). Fable 5 is even + stricter than Opus 4.8: an explicit ``thinking.type='disabled'`` also 400s, + so adaptive is the only valid thinking shape LiteLLM can emit for it.""" + variants = [k for k in cost_map if "claude-fable-5" in k] + assert variants, "no claude-fable-5 entries found in cost map" + missing = [ + k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True + ] + assert not missing, f"missing supports_adaptive_thinking: {missing}" + + +@pytest.mark.parametrize( + "model", + [ + "claude-fable-5", + "anthropic/claude-fable-5", + "anthropic.claude-fable-5", + "bedrock/us.anthropic.claude-fable-5", + "bedrock/invoke/eu.anthropic.claude-fable-5", + "bedrock/global.anthropic.claude-fable-5", + "vertex_ai/claude-fable-5", + "azure_ai/claude-fable-5", + ], +) +def test_adaptive_thinking_detected_for_fable_5(local_model_cost_map, model): + """Provider-routed ids must resolve to a flagged entry so ``reasoning_effort`` + maps to ``thinking.type='adaptive'`` + ``output_config.effort``.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + assert AnthropicModelInfo._is_adaptive_thinking_model(model) is True + + +@pytest.mark.parametrize( + "cost_map", + [_load_root_cost_map(), GetModelCostMap.load_local_model_cost_map()], + ids=["root", "bundled_backup"], +) +def test_sampling_params_flag_on_all_models_that_removed_them(cost_map): + """Fable 5 and Opus 4.7/4.8 reject ``top_p``/``top_k``/``temperature != 1``; + the drop/raise gating is cost-map driven, so every variant must carry an + explicit ``supports_sampling_params: false``. The perplexity route is + exempt: it is OpenAI-compatible and maps sampling params upstream.""" + variants = [ + k + for k in cost_map + if any(v in k for v in ("claude-fable-5", "claude-opus-4-7", "claude-opus-4-8")) + and not k.startswith("perplexity/") + ] + assert variants, "no matching entries found in cost map" + missing = [ + k for k in variants if cost_map[k].get("supports_sampling_params") is not False + ] + assert not missing, f"missing supports_sampling_params=false: {missing}" diff --git a/tests/test_litellm/test_register_model_custom_pricing.py b/tests/test_litellm/test_register_model_custom_pricing.py index 719cb8eecd2..e384d3e1161 100644 --- a/tests/test_litellm/test_register_model_custom_pricing.py +++ b/tests/test_litellm/test_register_model_custom_pricing.py @@ -301,6 +301,126 @@ def test_register_model_strips_none_litellm_provider_from_get_model_info(monkeyp litellm.model_cost.pop(model_key, None) +def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key(): + """Registering a custom override under a key shape that + ``get_model_info`` cannot resolve (e.g. a double provider prefix like + ``bedrock/bedrock/us.anthropic.claude-sonnet-4-6``) must still inherit + the built-in cache pricing for the underlying model. + + Before the fix ``register_model`` fell back to an empty ``existing_model`` + so the merged entry only carried the fields the user set explicitly + (input/output cost). ``cache_creation_input_token_cost`` and + ``cache_read_input_token_cost`` were absent, and the cost calculator + silently charged 0 for every cache token, dropping the bulk of the bill + for cache-heavy Anthropic traffic. + + Regression for the cache-pricing dropout under partial overrides. + """ + from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token + from litellm.types.utils import PromptTokensDetailsWrapper, Usage + + original_model_cost = litellm.model_cost + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + builtin_key = "us.anthropic.claude-sonnet-4-6" + registered_key = f"bedrock/bedrock/{builtin_key}" + builtin = litellm.model_cost[builtin_key] + + assert builtin["cache_creation_input_token_cost"] > 0 + assert builtin["cache_read_input_token_cost"] > 0 + + try: + litellm.register_model( + { + registered_key: { + "input_cost_per_token": builtin["input_cost_per_token"], + "output_cost_per_token": builtin["output_cost_per_token"], + "litellm_provider": "bedrock", + } + } + ) + + registered = litellm.model_cost[registered_key] + assert ( + registered.get("cache_creation_input_token_cost") + == builtin["cache_creation_input_token_cost"] + ) + assert ( + registered.get("cache_read_input_token_cost") + == builtin["cache_read_input_token_cost"] + ) + assert registered["litellm_provider"] == "bedrock" + + usage = Usage( + prompt_tokens=1100, + completion_tokens=100, + total_tokens=1200, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=800, + text_tokens=100, + ), + cache_creation_input_tokens=200, + ) + + input_cost, output_cost = generic_cost_per_token( + model=registered_key, + usage=usage, + custom_llm_provider="bedrock", + ) + + text_only_cost = builtin["input_cost_per_token"] * 100 + expected_input_cost = ( + text_only_cost + + builtin["cache_read_input_token_cost"] * 800 + + builtin["cache_creation_input_token_cost"] * 200 + ) + assert abs(input_cost - expected_input_cost) < 1e-12 + assert abs(output_cost - builtin["output_cost_per_token"] * 100) < 1e-12 + assert input_cost > text_only_cost + 1e-12 + finally: + litellm.model_cost.pop(registered_key, None) + litellm.model_cost = original_model_cost + os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None) + from litellm.utils import _invalidate_model_cost_lowercase_map + + _invalidate_model_cost_lowercase_map() + + +def test_register_model_warns_when_no_builtin_match_for_cache_pricing(caplog): + """When a custom override is registered under a key that neither + ``get_model_info`` nor any prefix/region variant can resolve to a + built-in entry, ``register_model`` must warn that cache cost fields will + default to 0 instead of silently producing an under-billed entry. + """ + import logging + + from litellm._logging import verbose_logger + + registered_key = "bedrock/totally-made-up-model-alias-xyz" + litellm.model_cost.pop(registered_key, None) + + try: + with caplog.at_level(logging.WARNING, logger=verbose_logger.name): + litellm.register_model( + { + registered_key: { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "litellm_provider": "bedrock", + } + } + ) + + assert any( + registered_key in record.message + and "cache_creation_input_token_cost" in record.message + for record in caplog.records + ), "expected a warning naming the unmapped key and the cache cost fields" + finally: + litellm.model_cost.pop(registered_key, None) + + def test_register_model_router_add_deployment_custom_pricing_applies(): """End-to-end regression for https://github.com/BerriAI/litellm/issues/28336. @@ -344,9 +464,9 @@ def test_register_model_router_add_deployment_custom_pricing_applies(): f"{model_key} / {deployment_model}" ) for k in registered_keys: - assert _check_provider_match(litellm.model_cost[k], "openai") is True, ( - f"custom pricing for {k} was dropped by _check_provider_match" - ) + assert ( + _check_provider_match(litellm.model_cost[k], "openai") is True + ), f"custom pricing for {k} was dropped by _check_provider_match" finally: litellm.model_cost.pop(model_key, None) litellm.model_cost.pop(deployment_model, None) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index cd235d8de67..e681247959f 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -80,6 +80,256 @@ def test_router_with_model_info_and_model_group(): ) +def test_router_model_group_encrypted_content_affinity_callback_registration(): + from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( + DeploymentAffinityCheck, + ) + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + model_group_affinity_config = { + model_group: ["encrypted_content_affinity"], + } + original_callbacks = list(litellm.callbacks) + litellm.callbacks = [] + router = None + + try: + router = litellm.Router( + model_list=[ + { + "model_name": model_group, + "litellm_params": { + "model": "openai/gpt-5.1-codex", + "api_key": "mock-api-key", + }, + } + ], + model_group_affinity_config=model_group_affinity_config, + num_retries=0, + ) + callbacks = router.optional_callbacks or [] + encrypted_content_callbacks = [ + cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) + ] + deployment_callback = next( + cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) + ) + assert len(encrypted_content_callbacks) == 1 + assert encrypted_content_callbacks[0].enable_global_affinity is False + assert ( + encrypted_content_callbacks[0].model_group_affinity_config + == model_group_affinity_config + ) + assert callbacks.index(encrypted_content_callbacks[0]) < callbacks.index( + deployment_callback + ) + assert litellm.callbacks.index(encrypted_content_callbacks[0]) < ( + litellm.callbacks.index(deployment_callback) + ) + + router._add_encrypted_content_affinity_check(enable_global_affinity=True) + + callbacks = router.optional_callbacks or [] + encrypted_content_callbacks = [ + cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) + ] + assert len(encrypted_content_callbacks) == 1 + assert encrypted_content_callbacks[0].enable_global_affinity is True + assert encrypted_content_callbacks[0].router is router + finally: + if router is not None: + router.discard() + litellm.callbacks = original_callbacks + + +@pytest.mark.asyncio +async def test_encrypted_content_affinity_model_group_config_is_additive(): + from litellm.responses.utils import ResponsesAPIRequestUtils + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + target_deployment = { + "model_name": model_group, + "litellm_params": {"model": "openai/gpt-5.1-codex"}, + "model_info": {"id": "deployment-b"}, + } + healthy_deployments = [ + { + "model_name": model_group, + "litellm_params": {"model": "openai/gpt-5.1-codex"}, + "model_info": {"id": "deployment-a"}, + }, + target_deployment, + ] + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( + "deployment-b", "rs_test" + ) + + assert EncryptedContentAffinityCheck.has_model_group_affinity_enabled( + {model_group: ["encrypted_content_affinity"]} + ) + assert not EncryptedContentAffinityCheck.has_model_group_affinity_enabled(None) + + per_group_check = EncryptedContentAffinityCheck( + enable_global_affinity=False, + model_group_affinity_config={ + model_group: ["encrypted_content_affinity"], + }, + ) + request_kwargs = { + "input": [{"type": "reasoning", "id": encoded_id}], + "litellm_metadata": {}, + } + filtered = await per_group_check.async_filter_deployments( + model=model_group, + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs=request_kwargs, + ) + + assert filtered == [target_deployment] + assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] + + disabled_check = EncryptedContentAffinityCheck( + enable_global_affinity=False, + model_group_affinity_config={ + "other-model-group": ["encrypted_content_affinity"], + }, + ) + disabled_request_kwargs = { + "input": [{"type": "reasoning", "id": encoded_id}], + "litellm_metadata": {}, + } + unfiltered = await disabled_check.async_filter_deployments( + model=model_group, + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs=disabled_request_kwargs, + ) + + assert unfiltered == healthy_deployments + assert "encrypted_content_affinity_enabled" not in disabled_request_kwargs[ + "litellm_metadata" + ] + + global_check = EncryptedContentAffinityCheck( + enable_global_affinity=True, + model_group_affinity_config={ + model_group: ["deployment_affinity"], + }, + ) + global_request_kwargs = { + "input": [{"type": "reasoning", "id": encoded_id}], + "litellm_metadata": {}, + } + globally_filtered = await global_check.async_filter_deployments( + model=model_group, + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs=global_request_kwargs, + ) + + assert globally_filtered == [target_deployment] + assert global_request_kwargs["litellm_metadata"][ + "encrypted_content_affinity_enabled" + ] + + +@pytest.mark.asyncio +async def test_encrypted_content_affinity_takes_priority_over_user_key_affinity(): + from litellm.responses.utils import ResponsesAPIRequestUtils + from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( + DeploymentAffinityCheck, + ) + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + model_group = "openai.gpt-5.1-codex" + user_api_key_hash = "test-user-key" + deployment_a = { + "model_name": model_group, + "litellm_params": { + "model": "openai/gpt-5.1-codex", + "api_key": "mock-api-key-a", + }, + "model_info": {"id": "deployment-a"}, + } + deployment_b = { + "model_name": model_group, + "litellm_params": { + "model": "openai/gpt-5.1-codex", + "api_key": "mock-api-key-b", + }, + "model_info": {"id": "deployment-b"}, + } + original_callbacks = list(litellm.callbacks) + litellm.callbacks = [] + router = None + + try: + router = litellm.Router( + model_list=[deployment_a, deployment_b], + model_group_affinity_config={ + model_group: [ + "deployment_affinity", + "encrypted_content_affinity", + ], + }, + num_retries=0, + ) + callbacks = router.optional_callbacks or [] + deployment_callback = next( + cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck) + ) + encrypted_content_callback = next( + cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck) + ) + assert callbacks.index(encrypted_content_callback) < callbacks.index( + deployment_callback + ) + assert litellm.callbacks.index(encrypted_content_callback) < ( + litellm.callbacks.index(deployment_callback) + ) + + cache_key = DeploymentAffinityCheck.get_affinity_cache_key( + model_group=model_group, + user_key=user_api_key_hash, + ) + await deployment_callback.cache.async_set_cache( + key=cache_key, + value={"model_id": "deployment-a"}, + ttl=60, + ) + encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id( + "deployment-b", "rs_test" + ) + request_kwargs = { + "input": [{"type": "reasoning", "id": encoded_id}], + "litellm_metadata": {"user_api_key_hash": user_api_key_hash}, + } + + filtered = await router.async_callback_filter_deployments( + model=model_group, + healthy_deployments=[deployment_a, deployment_b], + messages=None, + parent_otel_span=None, + request_kwargs=request_kwargs, + ) + + assert filtered == [deployment_b] + assert request_kwargs.get("_encrypted_content_affinity_pinned") is True + finally: + if router is not None: + router.discard() + litellm.callbacks = original_callbacks + + @pytest.mark.asyncio async def test_arouter_with_tags_and_fallbacks(): """ @@ -4311,6 +4561,48 @@ def test_get_available_deployment_for_pass_through_raises_when_dict_blocked(): ) +def test_initialize_deployment_for_pass_through_keeps_bedrock_iam_deployment(): + """ + Bedrock deployments using IAM/OIDC auth have no api_key; pass-through + init must not raise and drop them from routing (#27728). + """ + import litellm + + router = litellm.Router( + model_list=[ + { + "model_name": "bedrock-claude", + "litellm_params": { + "model": "bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0", + "aws_role_name": "arn:aws:iam::123456789012:role/my-role", + "aws_session_name": "my-session", + "use_in_pass_through": True, + }, + "model_info": {"id": "bedrock-iam-pt"}, + } + ] + ) + assert [m["model_info"]["id"] for m in router.get_model_list()] == [ + "bedrock-iam-pt" + ] + + +def test_initialize_deployment_for_pass_through_sets_credentials_with_api_key(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + passthrough_endpoint_router, + ) + + passthrough_endpoint_router.credentials.clear() + router = _router_with_two_pass_through_deployments([False, False]) + assert len(router.get_model_list()) == 2 + assert ( + passthrough_endpoint_router.get_credentials( + custom_llm_provider="openai", region_name=None + ) + == "sk-fake-for-tests" + ) + + def test_get_deployment_credentials_returns_none_for_blocked_deployment(): router = _router_with_two_deployments([True, False]) assert router.get_deployment_credentials(model_id="dep-0") is None diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index 9454e03e918..ee64f44d32c 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -402,3 +402,138 @@ def test_should_not_downgrade_chatgpt_shared_key_mode_with_alias_override(): assert bridge_model_info["mode"] == "responses" finally: _restore_model_cost_entries(model_keys) + + +def test_partial_custom_pricing_inherits_builtin_cache_pricing(): + """A deployment that overrides only input/output cost on a cache-supporting + model must still bill cache_read and cache_creation tokens. Before the + fix the deploy-id entry was registered with the user's two fields and + nothing else, so the cost calculator silently billed cache tokens at 0. + Regression for the prompt-caching cost dropout reported by the customer. + """ + backend_model = "anthropic/claude-sonnet-4-5-20250929" + deploy_id = "claude-deploy-partial-pricing" + + builtin_info = litellm.get_model_info(model=backend_model) + builtin_cache_create = builtin_info["cache_creation_input_token_cost"] + builtin_cache_read = builtin_info["cache_read_input_token_cost"] + assert builtin_cache_create is not None and builtin_cache_create > 0 + assert builtin_cache_read is not None and builtin_cache_read > 0 + + model_keys = { + deploy_id: litellm.model_cost.get(deploy_id), + backend_model: copy.deepcopy(litellm.model_cost.get(backend_model)), + } + try: + Router( + model_list=[ + { + "model_name": "claude-custom", + "litellm_params": { + "model": backend_model, + "api_key": "fake-key", + }, + "model_info": { + "id": deploy_id, + "input_cost_per_token": 0.000003, + "output_cost_per_token": 0.000015, + }, + } + ], + ) + + entry = litellm.model_cost[deploy_id] + assert entry["input_cost_per_token"] == 0.000003 + assert entry["output_cost_per_token"] == 0.000015 + assert entry.get("cache_creation_input_token_cost") == builtin_cache_create + assert entry.get("cache_read_input_token_cost") == builtin_cache_read + finally: + _restore_model_cost_entries(model_keys) + + +def test_partial_pricing_does_not_overwrite_explicit_cache_fields(): + """When the user explicitly sets cache_*_input_token_cost on a deployment, + those values must not be replaced by the built-in fallback. + """ + backend_model = "anthropic/claude-sonnet-4-5-20250929" + deploy_id = "claude-deploy-explicit-cache" + + explicit_cache_create = 0.00001 + explicit_cache_read = 0.0000005 + builtin_info = litellm.get_model_info(model=backend_model) + assert builtin_info["cache_creation_input_token_cost"] != explicit_cache_create + assert builtin_info["cache_read_input_token_cost"] != explicit_cache_read + + model_keys = { + deploy_id: litellm.model_cost.get(deploy_id), + backend_model: copy.deepcopy(litellm.model_cost.get(backend_model)), + } + try: + Router( + model_list=[ + { + "model_name": "claude-custom-explicit", + "litellm_params": { + "model": backend_model, + "api_key": "fake-key", + }, + "model_info": { + "id": deploy_id, + "input_cost_per_token": 0.000003, + "output_cost_per_token": 0.000015, + "cache_creation_input_token_cost": explicit_cache_create, + "cache_read_input_token_cost": explicit_cache_read, + }, + } + ], + ) + + entry = litellm.model_cost[deploy_id] + assert entry.get("cache_creation_input_token_cost") == explicit_cache_create + assert entry.get("cache_read_input_token_cost") == explicit_cache_read + finally: + _restore_model_cost_entries(model_keys) + + +def test_inherit_builtin_cache_pricing_fills_only_missing_fields(): + """Direct unit test of the helper: missing cache fields are filled from the + backend model's built-in entry, while an explicitly set cache field and the + user's input/output pricing are left untouched. + """ + backend_model = "anthropic/claude-sonnet-4-5-20250929" + builtin_info = litellm.get_model_info(model=backend_model) + builtin_cache_create = builtin_info["cache_creation_input_token_cost"] + builtin_cache_read = builtin_info["cache_read_input_token_cost"] + assert builtin_cache_create is not None and builtin_cache_create > 0 + assert builtin_cache_read is not None and builtin_cache_read > 0 + + explicit_cache_read = builtin_cache_read + 1 + model_info = { + "input_cost_per_token": 0.000003, + "cache_read_input_token_cost": explicit_cache_read, + } + + Router._inherit_builtin_cache_pricing( + model_info=model_info, + backend_model=backend_model, + custom_llm_provider="anthropic", + ) + + assert model_info["input_cost_per_token"] == 0.000003 + assert model_info["cache_read_input_token_cost"] == explicit_cache_read + assert model_info["cache_creation_input_token_cost"] == builtin_cache_create + + +def test_inherit_builtin_cache_pricing_noop_for_unknown_backend(): + """No canonical entry for the backend model means the helper leaves the + passed-in dict unchanged rather than raising. + """ + model_info = {"input_cost_per_token": 0.000003} + + Router._inherit_builtin_cache_pricing( + model_info=model_info, + backend_model="this-backend-model-does-not-exist-x9y8z7", + custom_llm_provider=None, + ) + + assert model_info == {"input_cost_per_token": 0.000003} diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index f179e9c8f93..62cb8154b6d 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -700,6 +700,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_read_input_token_cost": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens": {"type": "number"}, + "cache_read_input_token_cost_above_512k_tokens": {"type": "number"}, "cache_read_input_token_cost_batches": {"type": "number"}, "cache_creation_input_token_cost_above_1hr_above_200k_tokens": { "type": "number" @@ -721,6 +722,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_token_above_200k_tokens": {"type": "number"}, "input_cost_per_token_above_256k_tokens": {"type": "number"}, "input_cost_per_token_above_272k_tokens": {"type": "number"}, + "input_cost_per_token_above_512k_tokens": {"type": "number"}, "cache_read_input_token_cost_flex": {"type": "number"}, "cache_read_input_token_cost_priority": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens_priority": { @@ -811,6 +813,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "output_cost_per_token_above_200k_tokens": {"type": "number"}, "output_cost_per_token_above_256k_tokens": {"type": "number"}, "output_cost_per_token_above_272k_tokens": {"type": "number"}, + "output_cost_per_token_above_512k_tokens": {"type": "number"}, "output_cost_per_image_above_1024_and_1024_pixels": {"type": "number"}, "output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": { "type": "number" @@ -858,6 +861,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_xhigh_reasoning_effort": {"type": "boolean"}, "supports_max_reasoning_effort": {"type": "boolean"}, "supports_adaptive_thinking": {"type": "boolean"}, + "supports_sampling_params": {"type": "boolean"}, "supports_service_tier": {"type": "boolean"}, "supports_preset": {"type": "boolean"}, "supports_output_config": {"type": "boolean"}, @@ -931,6 +935,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_native_streaming": {"type": "boolean"}, "supports_image_size": {"type": "boolean"}, "supports_native_structured_output": {"type": "boolean"}, + "use_openai_responses_path": {"type": "boolean"}, "tiered_pricing": { "type": "array", "items": { diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index e2a7a9a4c49..7946202bffc 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -164,11 +164,6 @@ "count": 2 } }, - "src/app/(dashboard)/layout.tsx": { - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, "src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx": { "no-restricted-imports": { "count": 1 @@ -228,11 +223,6 @@ "count": 1 } }, - "src/app/page.tsx": { - "unused-imports/no-unused-imports": { - "count": 2 - } - }, "src/components/AIHub/AgentHubTableColumns.test.tsx": { "unused-imports/no-unused-imports": { "count": 1 @@ -1452,11 +1442,6 @@ "count": 1 } }, - "src/components/navbar.tsx": { - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, "src/components/networking.tsx": { "max-params": { "count": 23 diff --git a/ui/litellm-dashboard/knip.json b/ui/litellm-dashboard/knip.json index e95c0acef3f..6f129398981 100644 --- a/ui/litellm-dashboard/knip.json +++ b/ui/litellm-dashboard/knip.json @@ -1,8 +1,9 @@ { "$schema": "https://unpkg.com/knip@5/schema.json", - "entry": ["scripts/**/*.ts"], - "project": ["src/**/*.{ts,tsx}", "tests/**/*.{ts,tsx}", "scripts/**/*.ts", "e2e_tests/**/*.ts"], + "entry": ["scripts/**/*.{ts,mjs}"], + "project": ["src/**/*.{ts,tsx}", "tests/**/*.{ts,tsx}", "scripts/**/*.{ts,mjs}", "e2e_tests/**/*.ts"], "ignore": ["src/lib/http/schema.d.ts"], + "ignoreDependencies": ["openapi-typescript"], "playwright": { "config": "e2e_tests/playwright.config.ts", "entry": ["e2e_tests/**/*.spec.ts", "e2e_tests/**/*.setup.ts", "e2e_tests/globalSetup.ts"] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-reference/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-reference/page.tsx index 02bed1adbe5..a4a4d3d0f43 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-reference/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-reference/page.tsx @@ -1,10 +1,12 @@ "use client"; import APIReferenceView from "@/app/(dashboard)/api-reference/APIReferenceView"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings"; const APIReferencePage = () => { - const proxySettings = useProxySettings(); + const { accessToken } = useAuthorized(); + const proxySettings = useProxySettings(accessToken); return ; }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/proxySettings/useProxySettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/proxySettings/useProxySettings.ts index d4fb3073856..82cefd800f4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/proxySettings/useProxySettings.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/proxySettings/useProxySettings.ts @@ -1,21 +1,26 @@ -import { useState, useEffect } from "react"; import { fetchProxySettings } from "@/utils/proxyUtils"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useQuery } from "@tanstack/react-query"; +import { createQueryKeys } from "../common/queryKeysFactory"; -export default function useProxySettings() { - const { accessToken } = useAuthorized(); - const [proxySettings, setProxySettings] = useState({ - PROXY_BASE_URL: "", - PROXY_LOGOUT_URL: "", - LITELLM_UI_API_DOC_BASE_URL: null as string | null, - }); +export const proxySettingsKeys = createQueryKeys("proxySettings"); - useEffect(() => { - if (!accessToken) return; - fetchProxySettings(accessToken).then((settings) => { - if (settings) setProxySettings(settings); - }); - }, [accessToken]); - - return proxySettings; +export interface ProxySettings { + PROXY_BASE_URL: string; + PROXY_LOGOUT_URL: string; + LITELLM_UI_API_DOC_BASE_URL?: string | null; +} + +const EMPTY_PROXY_SETTINGS: ProxySettings = { + PROXY_BASE_URL: "", + PROXY_LOGOUT_URL: "", + LITELLM_UI_API_DOC_BASE_URL: null, +}; + +export default function useProxySettings(accessToken: string | null): ProxySettings { + const { data } = useQuery({ + queryKey: [...proxySettingsKeys.all, accessToken], + queryFn: () => fetchProxySettings(accessToken), + enabled: Boolean(accessToken), + }); + return data ?? EMPTY_PROXY_SETTINGS; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index bec1ad14297..df5b2ab4511 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -1,65 +1,63 @@ "use client"; -import React, { Suspense, useEffect, useState } from "react"; +import React, { Suspense, useState } from "react"; import Navbar from "@/components/navbar"; +import LoadingScreen from "@/components/common_components/LoadingScreen"; import { ThemeProvider } from "@/contexts/ThemeContext"; +import { useAuth } from "@/contexts/AuthContext"; import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useRouter, useSearchParams, usePathname } from "next/navigation"; import { DebugWarningBanner } from "@/components/DebugWarningBanner"; import { MIGRATED_PAGES, migratedHref, legacyPageHref, legacyKeyForPathname } from "@/utils/migratedPages"; -import i18n from "@/lib/i18n"; -function LayoutContent({ children }: { children: React.ReactNode }) { +function DashboardShell({ children }: { children: React.ReactNode }) { const router = useRouter(); const searchParams = useSearchParams(); const pathname = usePathname(); - const { accessToken } = useAuthorized(); - const [sidebarCollapsed, setSidebarCollapsed] = React.useState(false); - const [page, setPage] = useState(() => { - return legacyKeyForPathname(pathname) || searchParams.get("page") || "api-keys"; - }); + const { accessToken } = useAuth(); + const [sidebarCollapsed, setSidebarCollapsed] = useState(false); - const handleSetPage = (newPage: string) => { + const page = legacyKeyForPathname(pathname) || searchParams.get("page") || "api-keys"; + + const navigateToPage = (newPage: string) => { const migratedRoute = MIGRATED_PAGES[newPage]; router.push(migratedRoute ? migratedHref(migratedRoute) : legacyPageHref(newPage)); - setPage(newPage); }; - useEffect(() => { - setPage(legacyKeyForPathname(pathname) || searchParams.get("page") || "api-keys"); - }, [pathname, searchParams]); + return ( +
+ setSidebarCollapsed((v) => !v)} + /> + +
+
+ +
+
{children}
+
+
+ ); +} - const toggleSidebar = () => setSidebarCollapsed((v) => !v); +function LayoutContent({ children }: { children: React.ReactNode }) { + const searchParams = useSearchParams(); + const { accessToken } = useAuth(); + const isInvitationFlow = Boolean(searchParams.get("invitation_id")); return ( - -
- {}} - accessToken={accessToken} - /> - -
-
- -
-
{children}
-
-
+ + {isInvitationFlow ? children : {children}} ); } export default function Layout({ children }: { children: React.ReactNode }) { return ( - {i18n.t("common.loading")}} - > + }> {children} ); diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/page.tsx similarity index 55% rename from ui/litellm-dashboard/src/app/page.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/page.tsx index 81cce930998..0854b085fae 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/page.tsx @@ -1,7 +1,6 @@ "use client"; -import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider"; -import OldModelDashboard from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView"; +import ModelsAndEndpointsView from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView"; import PlaygroundPage from "@/app/(dashboard)/playground/page"; import AdminPanel from "@/components/AdminPanel"; import AgentsPanel from "@/components/agents"; @@ -10,6 +9,7 @@ import CacheDashboard from "@/components/cache_dashboard"; import ClaudeCodePluginsPanel from "@/components/claude_code_plugins"; import { teamListCall as v2TeamListCall } from "@/app/(dashboard)/hooks/teams/useTeams"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; +import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings"; import LoadingScreen from "@/components/common_components/LoadingScreen"; import { CostTrackingSettings } from "@/components/CostTrackingSettings"; import GeneralSettings from "@/components/general_settings"; @@ -19,7 +19,6 @@ import PoliciesPanel from "@/components/policies"; import { Team } from "@/components/key_team_helpers/key_list"; import { MCPServers } from "@/components/mcp_tools"; import ModelHubTable from "@/components/AIHub/ModelHubTable"; -import Navbar from "@/components/navbar"; import { Organization, proxyBaseUrl, getInProductNudgesCall } from "@/components/networking"; import NewUsagePage from "@/components/UsagePage/components/UsagePageView"; import OldTeams from "@/components/OldTeams"; @@ -44,7 +43,6 @@ import { MemoryView } from "@/components/MemoryView"; import WorkflowRuns from "@/components/workflow_runs"; import SpendLogsTable from "@/components/view_logs"; import ViewUserDashboard from "@/components/view_users"; -import { ThemeProvider } from "@/contexts/ThemeContext"; import { useAuth } from "@/contexts/AuthContext"; import { buildLoginUrlWithReturn, @@ -55,16 +53,8 @@ import { } from "@/utils/returnUrlUtils"; import { isAdminRole } from "@/utils/roles"; import { MIGRATED_PAGES, migratedHref } from "@/utils/migratedPages"; -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { useRouter, useSearchParams } from "next/navigation"; import { Suspense, useEffect, useMemo, useRef, useState } from "react"; -import { ConfigProvider, theme } from "antd"; - -interface ProxySettings { - PROXY_BASE_URL: string; - PROXY_LOGOUT_URL: string; - LITELLM_UI_API_DOC_BASE_URL?: string | null; -} function CreateKeyPageContent() { const { authLoading, token, userID, userRole, userEmail, accessToken, premiumUser, setUserRole, setUserEmail } = @@ -74,10 +64,7 @@ function CreateKeyPageContent() { const [keys, setKeys] = useState([]); const [organizations, setOrganizations] = useState([]); const [userModels, setUserModels] = useState([]); - const [proxySettings, setProxySettings] = useState({ - PROXY_BASE_URL: "", - PROXY_LOGOUT_URL: "", - }); + const proxySettings = useProxySettings(accessToken); const router = useRouter(); const searchParams = useSearchParams()!; @@ -96,12 +83,6 @@ function CreateKeyPageContent() { const [showClaudeCodePrompt, setShowClaudeCodePrompt] = useState(false); const [showClaudeCodeModal, setShowClaudeCodeModal] = useState(false); - // Dark mode state - const [isDarkMode, setIsDarkMode] = useState(false); - const toggleDarkMode = () => { - setIsDarkMode(!isDarkMode); - }; - const invitation_id = searchParams.get("invitation_id"); // Parse URL query parameters for pre-filling the create key form @@ -154,33 +135,11 @@ function CreateKeyPageContent() { }; }, [searchParams, autoOpenCreate]); - // Get page from URL, default to 'api-keys' if not present - const [page, setPage] = useState(() => { - return searchParams.get("page") || "api-keys"; - }); - - const updatePage = (newPage: string) => { - const migratedRoute = MIGRATED_PAGES[newPage]; - if (migratedRoute) { - router.push(migratedHref(migratedRoute)); - setPage(newPage); - return; - } - const newSearchParams = new URLSearchParams(searchParams); - newSearchParams.set("page", newPage); - window.history.pushState(null, "", `?${newSearchParams.toString()}`); - setPage(newPage); - }; - - const [sidebarCollapsed, setSidebarCollapsed] = useState(false); + const page = searchParams.get("page") || "api-keys"; // Track if we've already attempted a return URL redirect to prevent race conditions const hasAttemptedReturnRedirectRef = useRef(false); - const toggleSidebar = () => { - setSidebarCollapsed(!sidebarCollapsed); - }; - const addKey = (data: any) => { setKeys((prevData) => (prevData ? [...prevData, data] : [data])); setCreateClicked(() => !createClicked); @@ -349,14 +308,26 @@ function CreateKeyPageContent() { } return ( - }> - - - {invitation_id ? ( + <> + {invitation_id ? ( + + ) : ( + <> + {page == "api-keys" ? ( - ) : ( -
- + ) : page == "llm-playground" ? ( + + ) : page == "users" ? ( + + ) : page == "teams" ? ( + + ) : page == "organizations" ? ( + + ) : page == "admin-panel" ? ( + + ) : page == "logging-and-alerts" ? ( + + ) : page == "budgets" ? ( + + ) : page == "guardrails" ? ( + + ) : page == "policies" ? ( + + ) : page == "agents" ? ( + + ) : page == "prompts" ? ( + + ) : page == "transform-request" ? ( + + ) : page == "router-settings" ? ( + + ) : page == "ui-theme" ? ( + + ) : page == "cost-tracking" ? ( + + ) : page == "model-hub-table" ? ( + isAdminRole(userRole) ? ( + -
-
- -
- {page == "api-keys" ? ( - - ) : page == "models" ? ( - - ) : page == "llm-playground" ? ( - - ) : page == "users" ? ( - - ) : page == "teams" ? ( - - ) : page == "organizations" ? ( - - ) : page == "admin-panel" ? ( - - ) : page == "logging-and-alerts" ? ( - - ) : page == "budgets" ? ( - - ) : page == "guardrails" ? ( - - ) : page == "policies" ? ( - - ) : page == "agents" ? ( - - ) : page == "prompts" ? ( - - ) : page == "transform-request" ? ( - - ) : page == "router-settings" ? ( - - ) : page == "ui-theme" ? ( - - ) : page == "cost-tracking" ? ( - - ) : page == "model-hub-table" ? ( - isAdminRole(userRole) ? ( - - ) : ( - - ) - ) : page == "caching" ? ( - - ) : page == "pass-through-settings" ? ( - - ) : page == "logs" ? ( - - ) : page == "mcp-servers" ? ( - - ) : page == "search-tools" ? ( - - ) : page == "tag-management" ? ( - - ) : page == "skills" || page == "claude-code-plugins" ? ( - - ) : page == "access-groups" ? ( - - ) : page == "projects" ? ( - - ) : page == "vector-stores" ? ( - - ) : page == "tool-policies" ? ( - - ) : page == "workflows" ? ( - - ) : page == "memory" ? ( - - ) : page == "guardrails-monitor" ? ( - - ) : page == "new_usage" ? ( - - ) : ( - - )} -
- - {/* Survey Components */} - - - - {/* Claude Code Components */} - - -
+ ) : ( + + ) + ) : page == "caching" ? ( + + ) : page == "pass-through-settings" ? ( + + ) : page == "logs" ? ( + + ) : page == "mcp-servers" ? ( + + ) : page == "search-tools" ? ( + + ) : page == "tag-management" ? ( + + ) : page == "skills" || page == "claude-code-plugins" ? ( + + ) : page == "access-groups" ? ( + + ) : page == "projects" ? ( + + ) : page == "vector-stores" ? ( + + ) : page == "tool-policies" ? ( + + ) : page == "workflows" ? ( + + ) : page == "memory" ? ( + + ) : page == "guardrails-monitor" ? ( + + ) : page == "new_usage" ? ( + + ) : ( + )} -
-
-
+ + {/* Survey Components */} + + + + {/* Claude Code Components */} + + + + )} + ); } diff --git a/ui/litellm-dashboard/src/components/navbar.test.tsx b/ui/litellm-dashboard/src/components/navbar.test.tsx index 274e81db527..bdd0681dfa8 100644 --- a/ui/litellm-dashboard/src/components/navbar.test.tsx +++ b/ui/litellm-dashboard/src/components/navbar.test.tsx @@ -72,7 +72,10 @@ vi.mock("./Navbar/UserDropdown/UserDropdown", async (importOriginal) => { }); vi.mock("@/utils/proxyUtils", () => ({ - fetchProxySettings: vi.fn(), + fetchProxySettings: vi.fn().mockResolvedValue({ + PROXY_BASE_URL: "", + PROXY_LOGOUT_URL: "https://example.com/logout", + }), })); // Mock CommunityEngagementButtons component @@ -137,8 +140,6 @@ Object.defineProperty(window, "location", { describe("Navbar", () => { const defaultProps = { - proxySettings: {}, - setProxySettings: vi.fn(), accessToken: "test-token", isPublicPage: false, }; @@ -298,7 +299,9 @@ describe("Navbar", () => { const cookieUtils = vi.mocked(await import("@/utils/cookieUtils")); expect(cookieUtils.clearTokenCookies).toHaveBeenCalled(); - expect(window.location.href).toBe(""); + await waitFor(() => { + expect(window.location.href).toBe("https://example.com/logout"); + }); }); it("should not render dark mode toggle slider", () => { diff --git a/ui/litellm-dashboard/src/components/navbar.tsx b/ui/litellm-dashboard/src/components/navbar.tsx index e8895c95cc6..71b99d61fcf 100644 --- a/ui/litellm-dashboard/src/components/navbar.tsx +++ b/ui/litellm-dashboard/src/components/navbar.tsx @@ -6,11 +6,11 @@ import { getProxyBaseUrl } from "@/components/networking"; import { useTheme } from "@/contexts/ThemeContext"; import { clearTokenCookies } from "@/utils/cookieUtils"; import { clearStoredReturnUrl } from "@/utils/returnUrlUtils"; -import { fetchProxySettings } from "@/utils/proxyUtils"; +import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings"; import { DownOutlined, MenuFoldOutlined, MenuUnfoldOutlined } from "@ant-design/icons"; import { Tag } from "antd"; import Link from "next/link"; -import React, { useEffect, useState } from "react"; +import React from "react"; import { useTranslation } from "react-i18next"; import { BlogDropdown } from "./Navbar/BlogDropdown/BlogDropdown"; import { CommunityEngagementButtons } from "./Navbar/CommunityEngagementButtons/CommunityEngagementButtons"; @@ -21,8 +21,6 @@ import UserDropdown from "./Navbar/UserDropdown/UserDropdown"; import WorkerDropdown from "./Navbar/WorkerDropdown/WorkerDropdown"; interface NavbarProps { - proxySettings: any; - setProxySettings: React.Dispatch>; accessToken: string | null; isPublicPage: boolean; sidebarCollapsed?: boolean; @@ -30,8 +28,6 @@ interface NavbarProps { } const Navbar: React.FC = ({ - proxySettings, - setProxySettings, accessToken, isPublicPage = false, sidebarCollapsed = false, @@ -39,7 +35,7 @@ const Navbar: React.FC = ({ }) => { const { t } = useTranslation(); const baseUrl = getProxyBaseUrl(); - const [logoutUrl, setLogoutUrl] = useState(""); + const proxySettings = useProxySettings(accessToken); const { logoUrl } = useTheme(); const { data: healthData } = useHealthReadinessDetails(accessToken); const version = healthData?.litellm_version; @@ -50,29 +46,11 @@ const Navbar: React.FC = ({ const imageUrl = logoUrl || `${baseUrl}/get_image`; - useEffect(() => { - const initializeProxySettings = async () => { - if (accessToken) { - const settings = await fetchProxySettings(accessToken); - console.log("response from fetchProxySettings", settings); - if (settings) { - setProxySettings(settings); - } - } - }; - - initializeProxySettings(); - }, [accessToken]); - - useEffect(() => { - setLogoutUrl(proxySettings?.PROXY_LOGOUT_URL || ""); - }, [proxySettings]); - const handleLogout = () => { clearTokenCookies(); localStorage.removeItem("litellm_selected_worker_id"); localStorage.removeItem("litellm_worker_url"); - window.location.href = logoutUrl; + window.location.href = proxySettings.PROXY_LOGOUT_URL || ""; }; const handleWorkerSwitch = (workerId: string) => { diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index b33c2b741a1..43f85cc8674 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -467,3 +467,50 @@ describe("teamInfoCall", () => { expect(parsed.searchParams.has("team_id")).toBe(false); }); }); + +describe("sessionSpendLogsCall", () => { + const originalFetch = global.fetch; + + beforeEach(() => { + vi.clearAllMocks(); + }); + + afterEach(() => { + global.fetch = originalFetch; + }); + + it("should request the first page with defaults so the caller can page through the session", async () => { + const mockFetch = vi.fn().mockResolvedValue({ + ok: true, + json: vi.fn().mockResolvedValue({ data: [], total: 0, page: 1, page_size: 100, total_pages: 1 }), + } as any); + global.fetch = mockFetch as any; + + await Networking.sessionSpendLogsCall("token", "session-123"); + + expect(mockFetch).toHaveBeenCalledOnce(); + const [url] = mockFetch.mock.calls[0]; + const urlStr = typeof url === "string" ? url : (url as Request).url; + const parsed = typeof url === "string" ? new URL(url, "http://example.com") : new URL((url as Request).url); + + expect(urlStr).toContain("/spend/logs/session/ui"); + expect(parsed.searchParams.get("session_id")).toBe("session-123"); + expect(parsed.searchParams.get("page")).toBe("1"); + expect(parsed.searchParams.get("page_size")).toBe("100"); + }); + + it("should pass explicit page and page_size query params for later pages", async () => { + const mockFetch = vi.fn().mockResolvedValue({ + ok: true, + json: vi.fn().mockResolvedValue({ data: [], total: 250, page: 3, page_size: 100, total_pages: 3 }), + } as any); + global.fetch = mockFetch as any; + + await Networking.sessionSpendLogsCall("token", "session-123", 3, 100); + + const [url] = mockFetch.mock.calls[0]; + const parsed = typeof url === "string" ? new URL(url, "http://example.com") : new URL((url as Request).url); + expect(parsed.searchParams.get("page")).toBe("3"); + expect(parsed.searchParams.get("page_size")).toBe("100"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 7568f9ed972..268c937cc82 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5617,13 +5617,27 @@ export const teamPermissionsUpdateCall = async (accessToken: string, teamId: str }; /** - * Get all spend logs for a particular session + * Get a page of spend logs for a particular session. + * + * The backend paginates this endpoint (page / page_size, returning + * { data, total, page, page_size, total_pages }). Callers that need the whole + * session should page through total_pages and accumulate the results. */ -export const sessionSpendLogsCall = async (accessToken: string, session_id: string) => { +export const sessionSpendLogsCall = async ( + accessToken: string, + session_id: string, + page: number = 1, + page_size: number = 100, +) => { try { + const params = new URLSearchParams({ + session_id, + page: String(page), + page_size: String(page_size), + }); let url = proxyBaseUrl - ? `${proxyBaseUrl}/spend/logs/session/ui?session_id=${encodeURIComponent(session_id)}` - : `/spend/logs/session/ui?session_id=${encodeURIComponent(session_id)}`; + ? `${proxyBaseUrl}/spend/logs/session/ui?${params.toString()}` + : `/spend/logs/session/ui?${params.toString()}`; const response = await fetch(url, { method: "GET", diff --git a/ui/litellm-dashboard/src/components/public_model_hub.tsx b/ui/litellm-dashboard/src/components/public_model_hub.tsx index fd6ada836a3..4705731e26f 100644 --- a/ui/litellm-dashboard/src/components/public_model_hub.tsx +++ b/ui/litellm-dashboard/src/components/public_model_hub.tsx @@ -123,7 +123,6 @@ const PublicModelHub: React.FC = ({ accessToken, isEmbedded const [selectedModel, setSelectedModel] = useState(null); const [selectedAgent, setSelectedAgent] = useState(null); const [selectedMcpServer, setSelectedMcpServer] = useState(null); - const [proxySettings, setProxySettings] = useState({}); const [activeTab, setActiveTab] = useState("models"); const [skillHubData, setSkillHubData] = useState([]); const [skillLoading, setSkillLoading] = useState(false); @@ -983,14 +982,7 @@ const PublicModelHub: React.FC = ({ accessToken, isEmbedded
{/* Navigation - only show when not embedded */} - {!isEmbedded && ( - - )} + {!isEmbedded && }
{/* Embedded Explainer - only shown when embedded in dashboard */} diff --git a/ui/litellm-dashboard/src/components/routing_groups/index.tsx b/ui/litellm-dashboard/src/components/routing_groups/index.tsx index d33ec0dffad..6cd4d563865 100644 --- a/ui/litellm-dashboard/src/components/routing_groups/index.tsx +++ b/ui/litellm-dashboard/src/components/routing_groups/index.tsx @@ -7,6 +7,7 @@ import { Trans, useTranslation } from "react-i18next"; import { useRoutingGroups, useSaveRoutingGroups } from "@/app/(dashboard)/hooks/routingGroups/useRoutingGroups"; import { useRouterFields } from "@/app/(dashboard)/hooks/router/useRouterFields"; import { useModelHub } from "@/app/(dashboard)/hooks/models/useModels"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings"; import RoutingGroupsTable from "./RoutingGroupsTable"; import RoutingGroupModal from "./RoutingGroupModal"; @@ -20,7 +21,8 @@ const RoutingGroups: React.FC = () => { const { data, isLoading, refetch, isFetching } = useRoutingGroups(); const { data: routerFields } = useRouterFields(); const { data: modelHub } = useModelHub(); - const proxySettings = useProxySettings(); + const { accessToken } = useAuthorized(); + const proxySettings = useProxySettings(accessToken); const saveMutation = useSaveRoutingGroups(); const [searchQuery, setSearchQuery] = useState(""); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx index bc748cdf93e..b8d2c99e57c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx @@ -29,6 +29,14 @@ export interface LogDetailsDrawerProps { const SIDEBAR_WIDTH_PX = 224; +// Session logs are fetched page-by-page from the paginated backend and +// accumulated so the drawer can show the whole session. page_size is the +// backend maximum (le=100); the page cap bounds the fetch and the +// (un-virtualized) sidebar list for pathological sessions, keeping the most +// recent logs since the endpoint returns newest-first. +const SESSION_PAGE_SIZE = 100; +const MAX_SESSION_PAGES = 50; + /* ------------------------------------------------------------------ */ /* TraceEventRow — compact event row used in both session & non- */ /* session sidebar lists. Extracted to avoid JSX duplication. */ @@ -113,13 +121,39 @@ export function LogDetailsDrawer({ const [isSidebarCollapsed, setIsSidebarCollapsed] = useState(false); const [copiedLeftPanelId, setCopiedLeftPanelId] = useState(false); - const { data: sessionLogs = [] } = useQuery({ + const { data: sessionData } = useQuery({ queryKey: ["sessionLogs", sessionId], queryFn: async () => { - if (!sessionId || !accessToken) return []; - const response = await sessionSpendLogsCall(accessToken, sessionId); - const allSessionLogs: LogEntry[] = response.data || response || []; - return allSessionLogs + if (!sessionId || !accessToken) return { logs: [] as LogEntry[], total: 0 }; + + // Fetch the first page, then page through the rest so sessions with more + // than one page of logs are shown in full (capped for safety). + const firstPage = await sessionSpendLogsCall(accessToken, sessionId, 1, SESSION_PAGE_SIZE); + let rows: LogEntry[] = firstPage.data || firstPage || []; + const pagesToFetch = Math.min(firstPage.total_pages ?? 1, MAX_SESSION_PAGES); + + if (pagesToFetch > 1) { + const BATCH = 5; + const remaining: Awaited>[] = []; + for (let start = 2; start <= pagesToFetch; start += BATCH) { + const end = Math.min(start + BATCH - 1, pagesToFetch); + const batch = await Promise.all( + Array.from({ length: end - start + 1 }, (_, i) => + sessionSpendLogsCall(accessToken, sessionId, start + i, SESSION_PAGE_SIZE), + ), + ); + remaining.push(...batch); + } + for (const page of remaining) { + rows = rows.concat(page.data || []); + } + } + + // Fall back to the accumulated row count (not just the first page) when the + // backend omits total, so the truncation note reflects what was fetched. + const total: number = firstPage.total ?? rows.length; + + const logs = rows .map((row) => ({ ...row, request_duration_ms: row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime), @@ -128,24 +162,49 @@ export function LogDetailsDrawer({ const aIsMcp = MCP_CALL_TYPES.includes(a.call_type) ? 1 : 0; const bIsMcp = MCP_CALL_TYPES.includes(b.call_type) ? 1 : 0; if (aIsMcp !== bIsMcp) return aIsMcp - bIsMcp; - return new Date(a.startTime).getTime() - new Date(b.startTime).getTime(); + // Newest first, matching the all-sessions logs overview. MCP calls + // stay grouped last (above), newest-first within that group too. + return new Date(b.startTime).getTime() - new Date(a.startTime).getTime(); }); + + return { logs, total }; }, enabled: Boolean(open && isSessionMode && sessionId && accessToken), }); + const sessionLogs: LogEntry[] = sessionData?.logs ?? []; + // total reported by the backend; when the page cap truncates the fetch this + // exceeds sessionLogs.length, which drives the "showing most recent" note. + const sessionTotalCount = sessionData?.total ?? sessionLogs.length; + const sessionTruncated = sessionTotalCount > sessionLogs.length; + + // Default selection for a freshly opened session: the most recent log (latest + // startTime). The list is sorted newest-first, but MCP calls are grouped last, + // so the latest log by time is not necessarily sessionLogs[0]; compute it + // explicitly. A clicked/remembered log still wins over this default. + const mostRecentLog = useMemo( + () => + sessionLogs.reduce( + (latest, row) => + !latest || new Date(row.startTime).getTime() > new Date(latest.startTime).getTime() ? row : latest, + null, + ), + [sessionLogs], + ); + const currentLog = useMemo(() => { if (!isSessionMode) return logEntry; if (!sessionLogs.length) return null; + const fallbackLog = mostRecentLog ?? sessionLogs[0]; if (selectedSessionRequestId) { - return sessionLogs.find((row) => row.request_id === selectedSessionRequestId) || sessionLogs[0]; + return sessionLogs.find((row) => row.request_id === selectedSessionRequestId) || fallbackLog; } if (logEntry?.request_id) { const clickedLog = sessionLogs.find((row) => row.request_id === logEntry.request_id); - return clickedLog || sessionLogs[0]; + return clickedLog || fallbackLog; } - return sessionLogs[0]; - }, [isSessionMode, logEntry, selectedSessionRequestId, sessionLogs]); + return fallbackLog; + }, [isSessionMode, logEntry, selectedSessionRequestId, sessionLogs, mostRecentLog]); useEffect(() => { if (!isSessionMode || !sessionLogs.length) return; @@ -153,10 +212,10 @@ export function LogDetailsDrawer({ const fallbackRequestId = logEntry?.request_id && sessionLogs.some((row) => row.request_id === logEntry.request_id) ? logEntry.request_id - : sessionLogs[0].request_id; + : (mostRecentLog ?? sessionLogs[0]).request_id; setSelectedSessionRequestId(fallbackRequestId); } - }, [isSessionMode, logEntry, selectedSessionRequestId, sessionLogs]); + }, [isSessionMode, logEntry, selectedSessionRequestId, sessionLogs, mostRecentLog]); // Reset transient UI state when the drawer opens or closes. useEffect(() => { @@ -336,6 +395,11 @@ export function LogDetailsDrawer({ )}
+ {isSessionMode && sessionTruncated && ( +
+ Showing most recent {logsForList.length} of {sessionTotalCount} +
+ )}
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8379d0536a6..8e470819557 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -22012,6 +22012,11 @@ export interface components { * @default 60 */ database_connection_timeout: number | null; + /** + * Database Disable Prepared Statements + * @description Disable server-side prepared statements by setting Prisma's `pgbouncer=true` URL param. Use this for pgbouncer transaction-pooling deployments, or to prevent the 'cached plan must not change result type' error that pooled connections hit during rolling schema migrations. An explicit `pgbouncer` in `database_extra_connection_params` takes precedence. + */ + database_disable_prepared_statements?: boolean | null; /** * Database Extra Connection Params * @description Escape hatch: extra key/value pairs appended verbatim to the Prisma DATABASE_URL / DIRECT_URL query string (e.g. `sslmode`, `pgbouncer`, `statement_cache_size`). Keys here override any default LiteLLM sets. @@ -25081,6 +25086,12 @@ export interface components { * @default false */ use_litellm_proxy: boolean | null; + /** + * Use Xai Oauth + * @description Use stored xAI OAuth credentials when no xAI API key is configured. + * @default false + */ + use_xai_oauth: boolean | null; /** Vector Store Id */ vector_store_id?: string | null; /** Vertex Credentials */ @@ -32679,6 +32690,12 @@ export interface components { * @default false */ use_litellm_proxy: boolean | null; + /** + * Use Xai Oauth + * @description Use stored xAI OAuth credentials when no xAI API key is configured. + * @default false + */ + use_xai_oauth: boolean | null; /** Vector Store Id */ vector_store_id?: string | null; /** Vertex Credentials */ diff --git a/ui/litellm-dashboard/src/utils/migratedPages.test.ts b/ui/litellm-dashboard/src/utils/migratedPages.test.ts index e9aceea8148..bd4ad7af5b8 100644 --- a/ui/litellm-dashboard/src/utils/migratedPages.test.ts +++ b/ui/litellm-dashboard/src/utils/migratedPages.test.ts @@ -1,8 +1,13 @@ -import { describe, it, expect, vi, beforeEach } from "vitest"; +import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; describe("migratedHref / legacyPageHref", () => { beforeEach(() => { vi.resetModules(); + vi.stubEnv("NODE_ENV", "test"); + }); + + afterEach(() => { + vi.unstubAllEnvs(); }); it("builds a /ui-rooted path when serverRootPath is /", async () => { @@ -37,9 +42,48 @@ describe("migratedHref / legacyPageHref", () => { }); }); +describe("dev server (NODE_ENV=development)", () => { + beforeEach(() => { + vi.resetModules(); + vi.stubEnv("NODE_ENV", "development"); + }); + + afterEach(() => { + vi.unstubAllEnvs(); + }); + + it("builds root-relative hrefs because next dev serves the app at /, not /ui", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { migratedHref, legacyPageHref } = await import("./migratedPages"); + + expect(migratedHref("api-reference")).toBe("/api-reference"); + expect(legacyPageHref("models")).toBe("/?page=models"); + }); + + it("ignores serverRootPath, which only applies to proxy-mounted deployments", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/team-x/" })); + const { migratedHref } = await import("./migratedPages"); + + expect(migratedHref("api-reference")).toBe("/api-reference"); + }); + + it("maps a bare migrated path back to its legacy sidebar key", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { legacyKeyForPathname } = await import("./migratedPages"); + + expect(legacyKeyForPathname("/api-reference/")).toBe("api_ref"); + expect(legacyKeyForPathname("/")).toBeNull(); + }); +}); + describe("legacyKeyForPathname", () => { beforeEach(() => { vi.resetModules(); + vi.stubEnv("NODE_ENV", "test"); + }); + + afterEach(() => { + vi.unstubAllEnvs(); }); it("maps a migrated path back to its legacy sidebar key (including trailing slash)", async () => { diff --git a/ui/litellm-dashboard/src/utils/migratedPages.ts b/ui/litellm-dashboard/src/utils/migratedPages.ts index 2c27e4fee64..d6f1e7d6f1d 100644 --- a/ui/litellm-dashboard/src/utils/migratedPages.ts +++ b/ui/litellm-dashboard/src/utils/migratedPages.ts @@ -15,6 +15,11 @@ export const MIGRATED_PAGES: Record = { }; function uiBase(): string { + // next dev serves the app at the root; only the proxy mounts the static export under /ui + // (and optionally under server_root_path). Inlined at build time, so production is unaffected. + if (process.env.NODE_ENV === "development") { + return ""; + } const root = serverRootPath && serverRootPath !== "/" ? `/${serverRootPath.replace(/^\/+|\/+$/g, "")}` : ""; return `${root}/ui`; } diff --git a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx index 2ebde4f295d..60475380b14 100644 --- a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx +++ b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx @@ -149,7 +149,7 @@ vi.mock("@/lib/cva.config", () => ({ })); import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import CreateKeyPage from "@/app/page"; +import CreateKeyPage from "@/app/(dashboard)/page"; import { AuthProvider } from "@/contexts/AuthContext"; // The page consumes auth state via useAuth(). Wrap it so the hook resolves @@ -242,7 +242,7 @@ describe("CreateKeyPage auth behavior", () => { expect(wroteDeletion).toBe(true); }); - it("does NOT redirect when token is valid and renders the app chrome", async () => { + it("does NOT redirect when token is valid and renders the page content", async () => { // Arrange: valid token in cookie setCookie("token=validtoken"); @@ -269,9 +269,9 @@ describe("CreateKeyPage auth behavior", () => { expect(window.location.replace).not.toHaveBeenCalled(); }); - // And some top-level UI appears (Navbar stub) + // And the default page content appears (UserDashboard stub; chrome now lives in the layout) await waitFor(() => { - expect(screen.getByTestId("navbar")).toBeInTheDocument(); + expect(screen.getByTestId("user-dashboard")).toBeInTheDocument(); }); }); diff --git a/uv.lock b/uv.lock index 2403a7fbf03..1100db783d3 100644 --- a/uv.lock +++ b/uv.lock @@ -3297,6 +3297,11 @@ dependencies = [ caching = [ { name = "diskcache" }, ] +cli = [ + { name = "pyyaml" }, + { name = "requests" }, + { name = "rich" }, +] extra-proxy = [ { name = "a2a-sdk" }, { name = "azure-identity" }, @@ -3528,10 +3533,13 @@ requires-dist = [ { name = "pyroscope-io", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.8.16,<1.0" }, { name = "python-dotenv", specifier = ">=1.0.0,<2.0" }, { name = "python-multipart", marker = "extra == 'proxy'", specifier = ">=0.0.27,<1.0" }, + { name = "pyyaml", marker = "extra == 'cli'", specifier = ">=6.0.3,<7.0" }, { name = "pyyaml", marker = "extra == 'proxy'", specifier = ">=6.0.3,<7.0" }, { name = "redisvl", marker = "python_full_version < '3.14' and extra == 'extra-proxy'", specifier = ">=0.4.1,<1.0" }, + { name = "requests", marker = "extra == 'cli'", specifier = ">=2.32.0,<3.0" }, { name = "resend", marker = "extra == 'extra-proxy'", specifier = ">=2.23.0,<3.0" }, { name = "restrictedpython", marker = "extra == 'proxy'", specifier = ">=8.1,<9.0" }, + { name = "rich", marker = "extra == 'cli'", specifier = ">=13.9.4,<14.0" }, { name = "rich", marker = "extra == 'proxy'", specifier = ">=13.9.4,<14.0" }, { name = "rq", marker = "extra == 'proxy'", specifier = ">=2.7.0,<3.0" }, { name = "semantic-router", marker = "python_full_version < '3.14' and extra == 'semantic-router'", specifier = ">=0.1.15,<1.0" }, @@ -3545,7 +3553,7 @@ requires-dist = [ { name = "uvloop", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.21.0,<1.0" }, { name = "websockets", marker = "extra == 'proxy'", specifier = ">=15.0.1,<16.0" }, ] -provides-extras = ["proxy", "extra-proxy", "utils", "caching", "semantic-router", "mlflow", "grpc", "stt-nvidia-riva", "google", "proxy-runtime"] +provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "semantic-router", "mlflow", "grpc", "stt-nvidia-riva", "google", "proxy-runtime"] [package.metadata.requires-dev] ci = [