Merge branch 'main' into docs/prompt-caching-gemini-support

This commit is contained in:
Krish Dholakia 2026-03-21 10:28:39 -07:00 • committed by GitHub
commit c8a7d5d237
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
128 changed files with 7355 additions and 1222 deletions

View file

@ -54,6 +54,7 @@ response = completion(
- stream
- tools
- tool_choice
- include_server_side_tool_invocations
- functions
- response_format
- n
@ -856,7 +857,112 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
</TabItem>
</Tabs>
### URL Context
### Context Circulation (Server-Side Tool Combination)
Context circulation allows Gemini 3+ models to combine **built-in tools** (like Google Search) with **your custom functions** in the same request. Without it, Gemini returns an error if you try to use both.
When enabled, Gemini can execute Google Search server-side, use those results to decide whether to call your custom functions, and return the full chain of reasoning.
**How it works:**
1. You pass `include_server_side_tool_invocations=True` along with both Google Search and your function tools
2. Gemini executes server-side tools internally and returns `toolCall`/`toolResponse` parts alongside any `functionCall` parts
3. LiteLLM extracts the server-side invocations into `provider_specific_fields["server_side_tool_invocations"]`
4. On subsequent turns, include the full assistant message in your conversation history — LiteLLM re-injects the server-side parts automatically
<Tabs>
<TabItem value="sdk" label="SDK">
```python
from litellm import completion
response = completion(
model="gemini/gemini-3-flash-preview",
messages=[{"role": "user", "content": "What's the weather in Buenos Aires? If it's raining, schedule a meeting."}],
tools=[
{"type": "web_search_preview"}, # Google Search (server-side)
{
"type": "function",
"function": {
"name": "schedule_meeting",
"description": "Schedule a meeting",
"parameters": {
"type": "object",
"properties": {"reason": {"type": "string"}},
"required": ["reason"],
},
},
},
],
include_server_side_tool_invocations=True,
)
msg = response.choices[0].message
# Server-side tool results are in provider_specific_fields
psf = msg.provider_specific_fields or {}
for invocation in psf.get("server_side_tool_invocations", []):
print(invocation["tool_type"]) # e.g. "GOOGLE_SEARCH_WEB"
print(invocation["id"])
print(invocation["args"]) # e.g. {"queries": ["weather Buenos Aires"]}
print(invocation["response"]) # Search results from Google
# For multi-turn: just append the full message to history
messages.append(msg)
messages.append({"role": "user", "content": "Thanks!"})
# LiteLLM automatically re-injects the server-side parts + thought signatures
response2 = completion(
model="gemini/gemini-3-flash-preview",
messages=messages,
tools=tools,
include_server_side_tool_invocations=True,
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
1. Setup config.yaml
```yaml
model_list:
- model_name: gemini-3-flash
litellm_params:
model: gemini/gemini-3-flash-preview
api_key: os.environ/GEMINI_API_KEY
```
2. Start Proxy
```bash
$ litellm --config /path/to/config.yaml
```
3. Make Request
```bash
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-d '{
"model": "gemini-3-flash",
"messages": [{"role": "user", "content": "What is the weather in Buenos Aires?"}],
"tools": [
{"type": "web_search_preview"},
{"type": "function", "function": {"name": "schedule_meeting", "description": "Schedule a meeting", "parameters": {"type": "object", "properties": {"reason": {"type": "string"}}}}}
],
"include_server_side_tool_invocations": true
}'
```
</TabItem>
</Tabs>
:::info
- Context circulation requires **Gemini 3+** models
- Server-side tool invocations (`toolCall`/`toolResponse`) are **not** included in `tool_calls` — they are in `provider_specific_fields["server_side_tool_invocations"]` because they were already executed by Google, not by your code
- `thought_signatures` are automatically preserved alongside server-side invocations for multi-turn coherence
:::
### URL Context
<Tabs>
<TabItem value="sdk" label="SDK">

View file

@ -117,6 +117,14 @@ guardrails:
:::
:::note Streaming and post_call guardrails
For **streaming responses**, `post_call` guardrails run on the fully assembled response **after** all chunks have been delivered to the client. This means `post_call` guardrails on streaming are **audit-only** — they can inspect and log the complete response, but cannot block content delivery. Guardrail results are recorded in `guardrail_information` within the logging payload for compliance and auditing.
To filter or block streaming content in real-time, use `async_post_call_streaming_iterator_hook` instead, which processes chunks as they arrive.
:::
<details>
<summary>Advanced: Multiple modes with individual event hooks</summary>
@ -655,8 +663,8 @@ class myCustomGuardrail(CustomGuardrail):
| `apply_guardrail` | Simple method to check and optionally modify text | ✅ | INPUT or OUTPUT | ✅ | ✅ | ✅ |
| `async_pre_call_hook` | A hook that runs before the LLM API call | ✅ | INPUT | ✅ | ❌ | ✅ |
| `async_moderation_hook` | A hook that runs during the LLM API call| ✅ | INPUT | ❌ | ❌ | ✅ |
| `async_post_call_success_hook` | A hook that runs after a successful LLM API call| ✅ | INPUT, OUTPUT | ❌ | ✅ | ✅ |
| `async_post_call_streaming_iterator_hook` | A hook that processes streaming responses | ✅ | OUTPUT | ❌ | ✅ | ✅ |
| `async_post_call_success_hook` | A hook that runs after a successful LLM API call. For streaming, runs on the assembled response after delivery (audit-only, cannot block). | ✅ | INPUT, OUTPUT | ❌ | ✅ | ✅ (non-streaming only) |
| `async_post_call_streaming_iterator_hook` | A hook that processes streaming responses in real-time (can filter/block chunks) | ✅ | OUTPUT | ❌ | ✅ | ✅ |
## Frequently Asked Questions

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-enterprise"
version = "0.1.34"
version = "0.1.35"
description = "Package for LiteLLM Enterprise features"
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.1.33"
version = "0.1.35"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-enterprise==",

View file

@ -1,9 +1,9 @@
-- CreateIndex
CREATE INDEX "LiteLLM_TeamTable_organization_id_idx" ON "LiteLLM_TeamTable"("organization_id");
CREATE INDEX IF NOT EXISTS "LiteLLM_TeamTable_organization_id_idx" ON "LiteLLM_TeamTable"("organization_id");
-- CreateIndex
CREATE INDEX "LiteLLM_TeamTable_team_alias_idx" ON "LiteLLM_TeamTable"("team_alias");
CREATE INDEX IF NOT EXISTS "LiteLLM_TeamTable_team_alias_idx" ON "LiteLLM_TeamTable"("team_alias");
-- CreateIndex
CREATE INDEX "LiteLLM_TeamTable_created_at_idx" ON "LiteLLM_TeamTable"("created_at");
CREATE INDEX IF NOT EXISTS "LiteLLM_TeamTable_created_at_idx" ON "LiteLLM_TeamTable"("created_at");

View file

@ -60,13 +60,20 @@ class AnthropicCacheControlHook(CustomPromptManagement):
# Create a deep copy of messages to avoid modifying the original list
processed_messages = copy.deepcopy(messages)
# Process message-level cache controls
# Separate message-level and non-message-level injection points
remaining_points = []
for point in injection_points:
if point.get("location") == "message":
point = cast(CacheControlMessageInjectionPoint, point)
processed_messages = self._process_message_injection(
point=point, messages=processed_messages
)
else:
remaining_points.append(point)
# Pass through non-message injection points for provider-specific handling
if remaining_points:
non_default_params["cache_control_injection_points"] = remaining_points
return model, processed_messages, non_default_params

View file

@ -8,6 +8,7 @@ server-side using litellm router's search tools.
import asyncio
import math
import uuid
from typing import Any, Dict, List, Optional, Tuple, Union, cast
import litellm
@ -27,7 +28,9 @@ from litellm.integrations.websearch_interception.transformation import (
from litellm.types.integrations.websearch_interception import (
WebSearchInterceptionConfig,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
class WebSearchInterceptionLogger(CustomLogger):
@ -67,6 +70,111 @@ class WebSearchInterceptionLogger(CustomLogger):
self.search_tool_name = search_tool_name
self._request_has_websearch = False # Track if current request has web search
async def try_short_circuit_search(
self,
model: str,
messages: List[Dict],
tools: Optional[List[Dict]],
custom_llm_provider: Optional[str],
) -> Optional[Dict[str, Any]]:
"""
Short-circuit web-search-only requests by executing the search directly.
Claude Code sends web search as a separate, standalone /v1/messages
request with a simple prompt and only web_search tool(s). For providers
that don't natively support web search (e.g. github_copilot), there is
no need to route this through the backend LLM — we can detect the
pattern, execute the search via Tavily/Perplexity, and return a
synthetic Anthropic response immediately.
Args:
model: Model name from the request
messages: Messages list from the request
tools: Tools list from the request
custom_llm_provider: Provider name
Returns:
An AnthropicMessagesResponse dict if short-circuited, or None to
continue normal processing.
"""
if not tools:
return None
# Check if provider is in enabled list
provider_str = custom_llm_provider or ""
if (
self.enabled_providers is not None
and provider_str not in self.enabled_providers
):
return None
# Only short-circuit for providers without native Anthropic Messages
# support. Providers that have a BaseAnthropicMessagesConfig (bedrock,
# vertex_ai, azure_ai, anthropic) already use the agentic loop, which
# includes a follow-up LLM call to synthesize the answer from search
# results. Short-circuiting those would skip that synthesis step and
# return raw search text — a regression for existing users.
try:
provider_enum = LlmProviders(provider_str)
anthropic_config = (
ProviderConfigManager.get_provider_anthropic_messages_config(
model=model, provider=provider_enum
)
)
if anthropic_config is not None:
verbose_logger.debug(
f"WebSearchInterception: Skipping short-circuit for {provider_str} "
"(provider has native Anthropic Messages support, using agentic loop)"
)
return None
except (ValueError, Exception):
pass # unknown provider enum → safe to short-circuit
# All tools must be web search tools
if not all(is_web_search_tool(t) for t in tools):
return None
# Extract search query from the last user message
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_last_user_message,
)
query = get_last_user_message(cast(List[AllMessageValues], messages))
if not query:
return None
verbose_logger.debug(
"WebSearchInterception: Short-circuit search detected "
f"(provider={provider_str}, query='{query}')"
)
# Execute search
try:
search_result_text = await self._execute_search(query)
except Exception as e:
verbose_logger.error(
f"WebSearchInterception: Short-circuit search failed: {e}"
)
search_result_text = f"Search failed: {e}"
# Build synthetic Anthropic response
response: Dict[str, Any] = {
"id": f"msg_{str(uuid.uuid4())}",
"type": "message",
"role": "assistant",
"model": model,
"content": [{"type": "text", "text": search_result_text}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 0, "output_tokens": 0},
}
verbose_logger.debug(
"WebSearchInterception: Short-circuit search completed, "
f"returning synthetic response ({len(search_result_text)} chars)"
)
return response
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[Any]
) -> Optional[dict]:

View file

@ -1686,6 +1686,28 @@ class Logging(LiteLLMLoggingBaseClass):
)
return logging_result
def _merge_hidden_params_from_response_into_metadata(self, logging_result: Any) -> None:
"""
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
Non-streaming success uses _process_hidden_params_and_response_cost (skipped when
stream=True). Streaming assembles the full response later; without this merge,
OTEL/callbacks that read metadata.hidden_params miss cost-related fields.
"""
if logging_result is None:
return
hidden_params = getattr(logging_result, "_hidden_params", None)
if not hidden_params:
return
if self.model_call_details.get("litellm_params") is None:
return
self.model_call_details["litellm_params"].setdefault("metadata", {})
if self.model_call_details["litellm_params"]["metadata"] is None:
self.model_call_details["litellm_params"]["metadata"] = {}
self.model_call_details["litellm_params"]["metadata"]["hidden_params"] = getattr(
logging_result, "_hidden_params", {}
)
def _process_hidden_params_and_response_cost(
self,
logging_result,
@ -2010,6 +2032,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(result=complete_streaming_response)
self._merge_hidden_params_from_response_into_metadata(
complete_streaming_response
)
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"
@ -2545,6 +2570,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
self.model_call_details["response_cost"] = None
self._merge_hidden_params_from_response_into_metadata(
complete_streaming_response
)
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"

View file

@ -2150,22 +2150,36 @@ class CustomStreamWrapper:
self.sent_stream_usage = True
return response
asyncio.create_task(
self.logging_obj.async_success_handler(
_deferred_cb = getattr(
self.logging_obj,
"_on_deferred_stream_complete",
None,
)
if _deferred_cb is not None:
# Proxy has post-call guardrails — let the closure
# run guardrails on the assembled response, then
# fire logging with guardrail_information populated.
self.logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined]
asyncio.create_task(
_deferred_cb(complete_streaming_response, cache_hit)
)
else:
asyncio.create_task(
self.logging_obj.async_success_handler(
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
)
)
executor.submit(
self.logging_obj.success_handler,
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
)
)
executor.submit(
self.logging_obj.success_handler,
complete_streaming_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
)
raise StopAsyncIteration # Re-raise StopIteration
else:

View file

@ -1,3 +1,4 @@
import copy
import hashlib
import json
from typing import (
@ -833,6 +834,11 @@ class LiteLLMAnthropicMessagesAdapter:
if not schema:
return None
# Deep copy to avoid mutating the original schema
schema = copy.deepcopy(schema)
# OpenAI strict mode requires additionalProperties: false on every object
self._add_additional_properties_false(schema)
# Convert to OpenAI response_format structure
return {
"type": "json_schema",
@ -843,6 +849,40 @@ class LiteLLMAnthropicMessagesAdapter:
},
}
@staticmethod
def _add_additional_properties_false(schema: dict) -> None:
"""
Recursively ensure object schemas comply with OpenAI strict mode.
OpenAI's strict mode requires:
1. 'additionalProperties': false at every object nesting level
2. All property keys listed in 'required'
"""
if not isinstance(schema, dict):
return
if schema.get("type") == "object" and "properties" in schema:
schema["additionalProperties"] = False
schema["required"] = list(schema["properties"].keys())
for prop in schema["properties"].values():
LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(prop)
# Handle array items
if "items" in schema:
LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(schema["items"])
# Handle anyOf/oneOf/allOf
for key in ("anyOf", "oneOf", "allOf"):
if key in schema:
for sub_schema in schema[key]:
LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(sub_schema)
# Handle $defs / definitions
for key in ("$defs", "definitions"):
if key in schema:
for def_schema in schema[key].values():
LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(def_schema)
def _add_system_message_to_messages(
self,
new_messages: List[AllMessageValues],

View file

@ -8,7 +8,7 @@
import asyncio
import contextvars
from functools import partial
from typing import Any, AsyncIterator, Coroutine, Dict, List, Optional, Union
from typing import Any, AsyncIterator, Coroutine, Dict, List, Optional, Union, cast
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -114,6 +114,55 @@ async def _execute_pre_request_hooks(
return request_kwargs
async def _try_websearch_short_circuit(
model: str,
messages: List[Dict],
tools: Optional[List[Dict]],
custom_llm_provider: Optional[str],
stream: Optional[bool],
) -> Optional[Union[AnthropicMessagesResponse, AsyncIterator]]:
"""
Attempt to short-circuit a web-search-only request.
Claude Code sends web search as a separate, standalone /v1/messages
request. For providers that don't natively support web search (e.g.
github_copilot), we detect this pattern, execute the search via
Tavily/Perplexity, and return a synthetic Anthropic response — bypassing
the backend LLM entirely.
Returns the synthetic response if short-circuited, or None to continue
normal processing.
"""
if not litellm.callbacks:
return None
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
for callback in litellm.callbacks:
if not isinstance(callback, WebSearchInterceptionLogger):
continue
response = await callback.try_short_circuit_search(
model=model,
messages=messages,
tools=tools,
custom_llm_provider=custom_llm_provider,
)
if response is not None:
anthropic_response = cast(AnthropicMessagesResponse, response)
if stream:
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
return FakeAnthropicMessagesStreamIterator(anthropic_response)
return anthropic_response
return None
@client
async def anthropic_messages(
max_tokens: int,
@ -138,6 +187,12 @@ async def anthropic_messages(
"""
Async: Make llm api request in Anthropic /messages API spec
"""
# Save original stream flag before pre-request hooks can convert it.
# The websearch interception hook converts stream=True → stream=False
# for the agentic loop, but the short-circuit path needs to know
# whether the caller originally requested streaming.
original_stream = stream
# Execute pre-request hooks to allow CustomLoggers to modify request
request_kwargs = await _execute_pre_request_hooks(
model=model,
@ -151,11 +206,38 @@ async def anthropic_messages(
# Extract modified parameters
tools = request_kwargs.pop("tools", tools)
stream = request_kwargs.pop("stream", stream)
# Propagate the provider derived inside pre-request hooks, if not already set.
# The litellm_params dict may have been overwritten by **kwargs in
# _execute_pre_request_hooks, so fall back to get_llm_provider() if needed.
if not custom_llm_provider:
custom_llm_provider = request_kwargs.get("litellm_params", {}).get(
"custom_llm_provider"
)
if not custom_llm_provider:
try:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
except Exception:
pass
# Remove litellm_params from kwargs (only needed for hooks)
request_kwargs.pop("litellm_params", None)
# Merge back any other modifications
kwargs.update(request_kwargs)
# Short-circuit web-search-only requests: detect the pattern, execute
# search directly via Tavily/Perplexity, and return a synthetic response
# without ever touching the backend LLM or the adapter path.
# Use original_stream (not the hook-converted stream) so streaming
# callers get SSE events instead of a plain dict.
short_circuit_response = await _try_websearch_short_circuit(
model=model,
messages=messages,
tools=tools,
custom_llm_provider=custom_llm_provider,
stream=original_stream,
)
if short_circuit_response is not None:
return short_circuit_response
loop = asyncio.get_event_loop()
kwargs["is_async"] = True

View file

@ -1446,6 +1446,16 @@ class AmazonConverseConfig(BaseConfig):
original_tools, model, headers, additional_request_params
)
# Append cachePoint to tools if cache_control_injection_points has tool_config
cache_injection_points = additional_request_params.pop(
"cache_control_injection_points", None
)
if cache_injection_points and len(bedrock_tools) > 0:
for point in cache_injection_points:
if point.get("location") == "tool_config":
bedrock_tools.append({"cachePoint": {"type": "default"}})
break
bedrock_tool_config: Optional[ToolConfigBlock] = None
if len(bedrock_tools) > 0:
tool_choice_values: ToolChoiceValuesBlock = inference_params.pop(

View file

@ -64,8 +64,15 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
verbose_logger.debug(f"Transformed request: {bedrock_request}")
# Get endpoint URL using simplified function
api_base = litellm_params.get("api_base", None)
aws_bedrock_runtime_endpoint = litellm_params.get(
"aws_bedrock_runtime_endpoint", None
)
endpoint_url = self.get_bedrock_count_tokens_endpoint(
resolved_model, aws_region_name
model=resolved_model,
aws_region_name=aws_region_name,
api_base=api_base,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
)
verbose_logger.debug(f"Making request to: {endpoint_url}")

View file

@ -177,7 +177,11 @@ class BedrockCountTokensConfig(BaseAWSLLM):
return {"input": {"invokeModel": {"body": json.dumps(body_data)}}}
def get_bedrock_count_tokens_endpoint(
self, model: str, aws_region_name: str
self,
model: str,
aws_region_name: str,
api_base: Optional[str] = None,
aws_bedrock_runtime_endpoint: Optional[str] = None,
) -> str:
"""
Construct the AWS Bedrock CountTokens API endpoint using existing LiteLLM functions.
@ -185,6 +189,8 @@ class BedrockCountTokensConfig(BaseAWSLLM):
Args:
model: The resolved model ID from router lookup
aws_region_name: AWS region (e.g., "eu-west-1")
api_base: Optional custom API base URL (takes highest priority)
aws_bedrock_runtime_endpoint: Optional custom Bedrock runtime endpoint
Returns:
Complete endpoint URL for CountTokens API
@ -196,7 +202,11 @@ class BedrockCountTokensConfig(BaseAWSLLM):
if model_id.startswith("bedrock/"):
model_id = model_id[8:] # Remove "bedrock/" prefix
base_url = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com"
base_url, _ = self.get_runtime_endpoint(
api_base=api_base,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
aws_region_name=aws_region_name,
)
endpoint = f"{base_url}/model/{model_id}/count-tokens"
return endpoint

View file

@ -155,9 +155,11 @@ class MoonshotChatConfig(OpenAIGPTConfig):
message that contains tool_calls (multi-turn tool-calling flows).
For each such message that is missing the field:
1. Promote provider_specific_fields["reasoning_content"] if present and non-empty
1. Check if reasoning_content exists at the top level (for Pydantic models
that have the attribute but don't support 'in' operator)
2. Promote provider_specific_fields["reasoning_content"] if present and non-empty
(this is where LiteLLM stores it from a previous response)
2. Otherwise inject a single space — the minimum value the API accepts
3. Otherwise inject a single space — the minimum value the API accepts
Messages that already carry the field, or are not assistant/tool-call messages,
are appended as-is (no copy made).
"""
@ -166,7 +168,7 @@ class MoonshotChatConfig(OpenAIGPTConfig):
if (
msg.get("role") == "assistant"
and msg.get("tool_calls")
and "reasoning_content" not in msg
and not msg.get("reasoning_content") # Check using .get() which works for both dicts and Pydantic models
):
patched = dict(cast(dict, msg))
provider_fields = patched.get("provider_specific_fields") or {}

View file

@ -7,6 +7,7 @@ from typing import (
Any,
AsyncIterator,
Coroutine,
Dict,
Iterator,
List,
Literal,
@ -805,8 +806,8 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
choices = chunk.get("choices", [])
choices = self._map_reasoning_to_reasoning_content(choices)
kwargs = {
"id": chunk["id"],
kwargs: Dict[str, Any] = {
"id": chunk.get("id"),
"object": "chat.completion.chunk",
"created": chunk.get("created"),
"model": chunk.get("model"),

View file

@ -7,7 +7,7 @@ More information on our website: https://endpoints.ai.cloud.ovh.net
from typing import Optional, Union, List
import httpx
from litellm.utils import ModelResponseStream, get_model_info
from litellm.utils import ModelResponseStream, _get_model_info_helper
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm._logging import verbose_logger
from litellm.llms.ovhcloud.utils import OVHCloudException
@ -28,13 +28,17 @@ class OVHCloudChatConfig(OpenAIGPTConfig):
"""
supports_function_calling: Optional[bool] = None
try:
model_info = get_model_info(model, custom_llm_provider="ovhcloud")
supports_function_calling = model_info.get(
"supports_function_calling", False
model_info = _get_model_info_helper(
model, custom_llm_provider="ovhcloud"
)
supports_function_calling = model_info.get(
"supports_function_calling", None
)
if supports_function_calling is None:
supports_function_calling = False
except Exception as e:
verbose_logger.debug(f"Error getting supported OpenAI params: {e}")
pass
supports_function_calling = False
optional_params = super().get_supported_openai_params(model)
if supports_function_calling is not True:

View file

@ -540,6 +540,39 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
assistant_content.append(gemini_tool_call_part)
last_message_with_tool_calls = assistant_msg
## HANDLE SERVER-SIDE TOOL INVOCATIONS (context circulation)
_psf = assistant_msg.get("provider_specific_fields")
if isinstance(_psf, dict):
_ss_invocations = _psf.get("server_side_tool_invocations")
if isinstance(_ss_invocations, list):
for invocation in _ss_invocations:
# Re-inject toolCall part
tc_part: Dict[str, Any] = {
"toolCall": {
"toolType": invocation.get("tool_type"),
"id": invocation.get("id"),
"args": invocation.get("args"),
}
}
if "thought_signature" in invocation:
tc_part["thoughtSignature"] = invocation["thought_signature"]
assistant_content.append(tc_part) # type: ignore
# Re-inject toolResponse part if response is present
if "response" in invocation:
tr_dict: Dict[str, Any] = {
"id": invocation.get("id"),
"response": invocation.get("response"),
}
if invocation.get("tool_type"):
tr_dict["toolType"] = invocation["tool_type"]
tr_part: Dict[str, Any] = {
"toolResponse": tr_dict
}
if "thought_signature" in invocation:
tr_part["thoughtSignature"] = invocation["thought_signature"]
assistant_content.append(tr_part) # type: ignore
msg_i += 1
if assistant_content:
@ -666,6 +699,9 @@ def _transform_request_body( # noqa: PLR0915
)
tools: Optional[Tools] = optional_params.pop("tools", None)
tool_choice: Optional[ToolConfig] = optional_params.pop("tool_choice", None)
include_server_side_tool_invocations: bool = optional_params.pop(
"include_server_side_tool_invocations", False
)
safety_settings: Optional[List[SafetSettingsConfig]] = optional_params.pop(
"safety_settings", None
) # type: ignore
@ -715,6 +751,10 @@ def _transform_request_body( # noqa: PLR0915
data["tools"] = tools
if tool_choice is not None:
data["toolConfig"] = tool_choice
if include_server_side_tool_invocations:
if "toolConfig" not in data:
data["toolConfig"] = {}
data["toolConfig"]["includeServerSideToolInvocations"] = True
if safety_settings is not None:
data["safetySettings"] = safety_settings
if generation_config is not None and len(generation_config) > 0:

View file

@ -316,6 +316,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"audio",
"parallel_tool_calls",
"web_search_options",
"include_server_side_tool_invocations",
]
# Add penalty parameters only for non-preview models
@ -1119,6 +1120,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
optional_params = self._add_tools_to_optional_params(
optional_params, [_tools]
)
elif param == "include_server_side_tool_invocations" and value is True:
optional_params["include_server_side_tool_invocations"] = True
if litellm.vertex_ai_safety_settings is not None:
optional_params["safety_settings"] = litellm.vertex_ai_safety_settings
@ -1360,6 +1363,67 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
signatures.append(signature)
return signatures if signatures else None
@staticmethod
def _extract_server_side_tool_invocations(
parts: List[HttpxPartType],
) -> Optional[List[Dict[str, Any]]]:
"""Extract server-side tool invocations (toolCall/toolResponse) from parts.
These are returned by Gemini when context circulation is enabled
(includeServerSideToolInvocations=true). They represent tools executed
server-side (e.g. Google Search) and must be circulated back in
subsequent turns for multi-turn coherence.
Returns:
List of server-side invocation dicts if any found, None otherwise.
"""
invocations: List[Dict[str, Any]] = []
# Index toolCalls by id so we can pair them with responses
tool_calls_by_id: Dict[str, Dict[str, Any]] = {}
tool_responses_by_id: Dict[str, Dict[str, Any]] = {}
for part in parts:
if "toolCall" in part:
tc = part["toolCall"]
entry: Dict[str, Any] = {
"tool_type": tc.get("toolType"),
"id": tc.get("id"),
"args": tc.get("args"),
}
signature = part.get("thoughtSignature")
if signature is not None:
entry["thought_signature"] = signature
tool_calls_by_id[tc.get("id", "")] = entry
elif "toolResponse" in part:
tr = part["toolResponse"]
entry = {
"id": tr.get("id"),
"tool_type": tr.get("toolType"),
"response": tr.get("response"),
}
signature = part.get("thoughtSignature")
if signature is not None:
entry["thought_signature"] = signature
tool_responses_by_id[tr.get("id", "")] = entry
# Merge calls with their responses
for call_id, call_entry in tool_calls_by_id.items():
merged = dict(call_entry)
resp = tool_responses_by_id.pop(call_id, None)
if resp is not None:
merged["response"] = resp.get("response")
# Keep response signature if call didn't have one
if "thought_signature" not in merged and "thought_signature" in resp:
merged["thought_signature"] = resp["thought_signature"]
invocations.append(merged)
# Any orphan responses (shouldn't happen, but be safe)
for resp_id, resp_entry in tool_responses_by_id.items():
invocations.append(resp_entry)
return invocations if invocations else None
def _extract_image_response_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[List[ImageURLListItem]]:
@ -2018,6 +2082,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None
reasoning_content: Optional[str] = None
thought_signatures: Optional[Any] = None
server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None
for idx, candidate in enumerate(_candidates):
if "content" not in candidate:
@ -2068,6 +2133,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
)
# Extract server-side tool invocations (context circulation)
server_side_tool_invocations = (
VertexGeminiConfig._extract_server_side_tool_invocations(
parts=candidate["content"]["parts"]
)
)
if audio_response is not None:
cast(Dict[str, Any], chat_completion_message)[
"audio"
@ -2139,6 +2211,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
chat_completion_message["provider_specific_fields"] = {}
chat_completion_message["provider_specific_fields"]["thought_signatures"] = thought_signatures # type: ignore
# Store server-side tool invocations in provider_specific_fields
if server_side_tool_invocations is not None:
if "provider_specific_fields" not in chat_completion_message:
chat_completion_message["provider_specific_fields"] = {}
chat_completion_message["provider_specific_fields"]["server_side_tool_invocations"] = server_side_tool_invocations # type: ignore
if isinstance(model_response, ModelResponseStream):
choice = VertexGeminiConfig._create_streaming_choice(
chat_completion_message=chat_completion_message,

View file

@ -152,6 +152,8 @@ def transform_openai_input_gemini_content(
gemini_params = optional_params.copy()
if "dimensions" in gemini_params:
gemini_params["outputDimensionality"] = gemini_params.pop("dimensions")
if "task_type" in gemini_params:
gemini_params["taskType"] = gemini_params.pop("task_type")
requests: List[EmbedContentRequest] = []
if isinstance(input, str):
@ -196,6 +198,8 @@ def transform_openai_input_gemini_embed_content(
gemini_params = optional_params.copy()
if "dimensions" in gemini_params:
gemini_params["outputDimensionality"] = gemini_params.pop("dimensions")
if "task_type" in gemini_params:
gemini_params["taskType"] = gemini_params.pop("task_type")
input_list = [input] if isinstance(input, str) else input
parts: List[PartType] = []

View file

@ -1672,6 +1672,7 @@ class NewTeamRequest(TeamBase):
int
] = None # allow user to set TPM limit for all team members
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
team_member_budget_duration: Optional[str] = None # e.g. "30d", "1mo"
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
enforced_batch_output_expires_after: Optional[dict] = None
enforced_file_expires_after: Optional[dict] = None

View file

@ -539,8 +539,45 @@ def bytes_to_mb(bytes_value: int):
# helpers used by parallel request limiter to handle model rpm/tpm limits for a given api key
def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]:
"""
Return the minimum value of `field` across all deployments for model_name,
or None if no deployment has the field set.
When multiple deployments share the same model name, taking the minimum is
the safest choice for load-balanced setups: it ensures no deployment is
over-consumed regardless of which one actually serves a given request.
"""
from litellm.proxy.proxy_server import llm_router
if llm_router is None:
return None
deployments = llm_router.get_model_list(model_name=model_name)
if not deployments:
return None
limits = []
for deployment in deployments:
raw = deployment.get("litellm_params", {}).get(field)
if raw is not None:
try:
if isinstance(raw, (int, float, str, bytes, bytearray)):
limits.append(int(raw))
except (ValueError, TypeError):
pass
return min(limits) if limits else None
def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]:
return _get_deployment_default_limit(model_name, "default_api_key_rpm_limit")
def _get_deployment_default_tpm_limit(model_name: str) -> Optional[int]:
return _get_deployment_default_limit(model_name, "default_api_key_tpm_limit")
def get_key_model_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
model_name: Optional[str] = None,
) -> Optional[Dict[str, int]]:
"""
Get the model rpm limit for a given api key.
@ -549,6 +586,7 @@ def get_key_model_rpm_limit(
1. Key metadata (model_rpm_limit)
2. Key model_max_budget (rpm_limit per model)
3. Team metadata (model_rpm_limit)
4. Deployment default_api_key_rpm_limit (when model_name is provided)
"""
# 1. Check key metadata first (takes priority)
if user_api_key_dict.metadata:
@ -567,13 +605,22 @@ def get_key_model_rpm_limit(
# 3. Fallback to team metadata
if user_api_key_dict.team_metadata:
return user_api_key_dict.team_metadata.get("model_rpm_limit")
team_limit = user_api_key_dict.team_metadata.get("model_rpm_limit")
if team_limit is not None:
return team_limit
# 4. Fallback to deployment default_api_key_rpm_limit
if model_name is not None:
default_limit = _get_deployment_default_rpm_limit(model_name)
if default_limit is not None:
return {model_name: default_limit}
return None
def get_key_model_tpm_limit(
user_api_key_dict: UserAPIKeyAuth,
model_name: Optional[str] = None,
) -> Optional[Dict[str, int]]:
"""
Get the model tpm limit for a given api key.
@ -582,6 +629,7 @@ def get_key_model_tpm_limit(
1. Key metadata (model_tpm_limit)
2. Key model_max_budget (tpm_limit per model)
3. Team metadata (model_tpm_limit)
4. Deployment default_api_key_tpm_limit (when model_name is provided)
"""
# 1. Check key metadata first (takes priority)
if user_api_key_dict.metadata:
@ -600,7 +648,15 @@ def get_key_model_tpm_limit(
# 3. Fallback to team metadata
if user_api_key_dict.team_metadata:
return user_api_key_dict.team_metadata.get("model_tpm_limit")
team_limit = user_api_key_dict.team_metadata.get("model_tpm_limit")
if team_limit is not None:
return team_limit
# 4. Fallback to deployment default_api_key_tpm_limit
if model_name is not None:
default_limit = _get_deployment_default_tpm_limit(model_name)
if default_limit is not None:
return {model_name: default_limit}
return None

View file

@ -45,7 +45,9 @@ from litellm.proxy.common_utils.callback_utils import (
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.utils import ProxyLogging
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.router import Router
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import ServerToolUse
# Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format)
@ -801,7 +803,7 @@ class ProxyBaseLLMRequestProcessing:
json.dumps(self.data, indent=4, default=str),
)
async def base_process_llm_request(
async def base_process_llm_request( # noqa: PLR0915
self,
request: Request,
fastapi_response: Response,
@ -926,6 +928,26 @@ class ProxyBaseLLMRequestProcessing:
llm_router=llm_router,
)
# Defer async logging when post-call guardrails are configured so the
# StandardLoggingPayload is built after guardrails write to metadata.
# Cache the result to avoid scanning litellm.callbacks twice.
_post_call_guardrails_active = self._has_post_call_guardrails()
# Non-streaming: defer the create_task in wrapper_async so the
# SLP is built after guardrails write to metadata. Streaming
# uses a separate closure mechanism (see below).
#
# Edge case: if _is_streaming_request is False but the response
# turns out to be a CustomStreamWrapper (rare provider behavior),
# wrapper_async exits early before the _defer_async_logging block
# so _enqueue_deferred_logging is never stored — the finally
# block is a no-op. The CSW path handles this correctly via
# _on_deferred_stream_complete, which fires its own logging.
if _post_call_guardrails_active and not self._is_streaming_request(
data=self.data, is_streaming_request=is_streaming_request
):
logging_obj._defer_async_logging = True # type: ignore
tasks = []
# Start the moderation check (during_call_hook) as early as possible
# This gives it a head start to mask/validate input while the proxy handles routing
@ -962,124 +984,229 @@ class ProxyBaseLLMRequestProcessing:
response = responses[1]
hidden_params = getattr(response, "_hidden_params", {}) or {}
model_id = self._get_model_id_from_response(hidden_params, self.data)
_exception_raised = False
try:
hidden_params = getattr(response, "_hidden_params", {}) or {}
model_id = self._get_model_id_from_response(hidden_params, self.data)
cache_key, api_base, response_cost = (
hidden_params.get("cache_key", None) or "",
hidden_params.get("api_base", None) or "",
hidden_params.get("response_cost", None) or "",
)
fastest_response_batch_completion, additional_headers = (
hidden_params.get("fastest_response_batch_completion", None),
hidden_params.get("additional_headers", {}) or {},
)
# Post Call Processing
if llm_router is not None:
self.data["deployment"] = llm_router.get_deployment(model_id=model_id)
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=self.data.get("litellm_call_id", ""), status="success"
cache_key, api_base, response_cost = (
hidden_params.get("cache_key", None) or "",
hidden_params.get("api_base", None) or "",
hidden_params.get("response_cost", None) or "",
)
)
if self._is_streaming_request(
data=self.data, is_streaming_request=is_streaming_request
) or self._is_streaming_response(
response
): # use generate_responses to stream responses
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=logging_obj.litellm_call_id,
model_id=model_id,
cache_key=cache_key,
api_base=api_base,
version=version,
response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
fastest_response_batch_completion=fastest_response_batch_completion,
request_data=self.data,
hidden_params=hidden_params,
litellm_logging_obj=logging_obj,
**additional_headers,
fastest_response_batch_completion, additional_headers = (
hidden_params.get("fastest_response_batch_completion", None),
hidden_params.get("additional_headers", {}) or {},
)
# Call response headers hook for streaming success
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=self.data,
user_api_key_dict=user_api_key_dict,
response=response,
request_headers=dict(request.headers),
# Post Call Processing
if llm_router is not None:
self.data["deployment"] = llm_router.get_deployment(model_id=model_id)
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=self.data.get("litellm_call_id", ""), status="success"
)
)
if callback_headers:
custom_headers.update(callback_headers)
if self._is_streaming_request(
data=self.data, is_streaming_request=is_streaming_request
) or self._is_streaming_response(
response
): # use generate_responses to stream responses
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=logging_obj.litellm_call_id,
model_id=model_id,
cache_key=cache_key,
api_base=api_base,
version=version,
response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
fastest_response_batch_completion=fastest_response_batch_completion,
request_data=self.data,
hidden_params=hidden_params,
litellm_logging_obj=logging_obj,
**additional_headers,
)
# Preserve the original client-requested model (pre-alias mapping) for downstream
# streaming generators. Pre-call processing can rewrite `self.data["model"]` for
# aliasing/routing, but the OpenAI-compatible response `model` field should reflect
# what the client sent.
if requested_model_from_client:
self.data[
"_litellm_client_requested_model"
] = requested_model_from_client
if route_type == "allm_passthrough_route":
# Check if response is an async generator
if self._is_streaming_response(response):
if asyncio.iscoroutine(response):
generator = await response
else:
generator = response
# Call response headers hook for streaming success
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=self.data,
user_api_key_dict=user_api_key_dict,
response=response,
request_headers=dict(request.headers),
)
if callback_headers:
custom_headers.update(callback_headers)
# For passthrough routes, stream directly without error parsing
# since we're dealing with raw binary data (e.g., AWS event streams)
return StreamingResponse(
content=generator,
status_code=status.HTTP_200_OK,
headers=custom_headers,
)
else:
# Traditional HTTP response with aiter_bytes
return StreamingResponse(
content=response.aiter_bytes(),
status_code=response.status_code,
headers=custom_headers,
)
elif route_type == "anthropic_messages":
# Check if response is actually a streaming response (async generator)
# Non-streaming responses (dict) should be returned directly
# This handles cases like websearch_interception agentic loop
# which returns a non-streaming dict even for streaming requests
if self._is_streaming_response(response):
selected_data_generator = (
ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=self.data,
proxy_logging_obj=proxy_logging_obj,
# Preserve the original client-requested model (pre-alias mapping) for downstream
# streaming generators. Pre-call processing can rewrite `self.data["model"]` for
# aliasing/routing, but the OpenAI-compatible response `model` field should reflect
# what the client sent.
if requested_model_from_client:
self.data[
"_litellm_client_requested_model"
] = requested_model_from_client
# Streaming: attach a closure that CSW.__anext__ will call
# at stream end instead of firing logging directly. The
# closure runs ONLY guardrail hooks (not all callbacks) on
# the assembled response so guardrail_information is
# populated, then fires both logging handlers.
# Only for CustomStreamWrapper — raw async generators from
# passthrough routes bypass CSW and would orphan the closure.
from litellm.litellm_core_utils.streaming_handler import (
CustomStreamWrapper,
)
if _post_call_guardrails_active and isinstance(
response, CustomStreamWrapper
):
# Intentionally a live reference (not a copy) — mirrors
# ProxyLogging.post_call_success_hook which also mutates
# data["guardrail_to_apply"] during iteration.
_captured_data = self.data
_captured_user_api_key_dict = user_api_key_dict
_captured_logging_obj = logging_obj
async def _on_deferred_stream_complete(
assembled_response, cache_hit
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=_captured_data,
captured_user_api_key_dict=_captured_user_api_key_dict,
captured_logging_obj=_captured_logging_obj,
assembled_response=assembled_response,
cache_hit=cache_hit,
)
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete # type: ignore[attr-defined]
if route_type == "allm_passthrough_route":
# Check if response is an async generator
if self._is_streaming_response(response):
if asyncio.iscoroutine(response):
generator = await response
else:
generator = response
# For passthrough routes, stream directly without error parsing
# since we're dealing with raw binary data (e.g., AWS event streams)
return StreamingResponse(
content=generator,
status_code=status.HTTP_200_OK,
headers=custom_headers,
)
else:
# Traditional HTTP response with aiter_bytes
return StreamingResponse(
content=response.aiter_bytes(),
status_code=response.status_code,
headers=custom_headers,
)
elif route_type == "anthropic_messages":
# Check if response is actually a streaming response (async generator)
# Non-streaming responses (dict) should be returned directly
# This handles cases like websearch_interception agentic loop
# which returns a non-streaming dict even for streaming requests
if self._is_streaming_response(response):
selected_data_generator = (
ProxyBaseLLMRequestProcessing.async_sse_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=self.data,
proxy_logging_obj=proxy_logging_obj,
)
)
return await create_response(
generator=selected_data_generator,
media_type="text/event-stream",
headers=custom_headers,
)
# Non-streaming response - fall through to normal response handling
elif select_data_generator:
selected_data_generator = select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=self.data,
)
return await create_response(
generator=selected_data_generator,
media_type="text/event-stream",
headers=custom_headers,
)
# Non-streaming response - fall through to normal response handling
elif select_data_generator:
selected_data_generator = select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=self.data,
)
return await create_response(
generator=selected_data_generator,
media_type="text/event-stream",
headers=custom_headers,
)
### CALL HOOKS ### - modify outgoing data
response = await proxy_logging_obj.post_call_success_hook(
data=self.data, user_api_key_dict=user_api_key_dict, response=response
)
### CALL HOOKS ### - modify outgoing data
# If we reach here with a streaming closure still set, it means
# no early-return route consumed the CSW (hypothetical fallthrough).
# Clear the closure so guardrails run inline as before — this
# preserves blocking behavior and avoids double invocation.
if getattr(logging_obj, "_on_deferred_stream_complete", None):
logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined]
response = await proxy_logging_obj.post_call_success_hook(
data=self.data, user_api_key_dict=user_api_key_dict, response=response
)
except Exception:
_exception_raised = True
raise
finally:
# Enqueue deferred logging after post-call guardrails have written
# guardrail_information to metadata. The finally block ensures
# logging fires even if a guardrail raises.
# For streaming early-returns: no closure is stored (wrapper_async
# returns before the deferred block), so _enqueue_fn is None — no-op.
_enqueue_fn = getattr(logging_obj, "_enqueue_deferred_logging", None)
if _enqueue_fn is not None:
logging_obj._enqueue_deferred_logging = None # type: ignore[attr-defined]
try:
_enqueue_fn()
except Exception as e:
verbose_proxy_logger.exception(
"Error firing deferred logging: %s", e
)
# Streaming cleanup: if an exception occurred AND the deferred
# streaming closure is still set, no streaming route will
# consume the CSW — the closure is orphaned. Clear it and
# fire logging directly to avoid silent loss.
#
# On normal streaming returns the closure must stay: CSW calls
# it at stream end. _exception_raised is function-scoped and
# immune to outer exception context, avoiding false positives.
if _exception_raised:
_deferred_fn = getattr(
logging_obj, "_on_deferred_stream_complete", None
)
if _deferred_fn is not None:
logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined]
try:
asyncio.create_task(
logging_obj.async_success_handler(
response,
cache_hit=None,
start_time=None,
end_time=None,
)
)
except Exception as e:
verbose_proxy_logger.exception(
"Error in orphaned streaming async logging: %s", e
)
try:
from litellm.litellm_core_utils.thread_pool_executor import (
executor as _exc,
)
_exc.submit(
logging_obj.success_handler,
response,
cache_hit=None,
start_time=None,
end_time=None,
)
except Exception as e:
verbose_proxy_logger.exception(
"Error in orphaned streaming sync logging: %s", e
)
# Always return the client-requested model name (not provider-prefixed internal identifiers)
# for OpenAI-compatible responses.
@ -1217,6 +1344,131 @@ class ProxyBaseLLMRequestProcessing:
return True
return False
@staticmethod
def _has_post_call_guardrails() -> bool:
"""
Check if any registered callback is a post-call guardrail.
Uses the global litellm.callbacks list rather than per-request
should_run_guardrail() — intentionally conservative so that the
check is simple and stateless. The deferral path produces
identical logging output, just fires it slightly later, so
false-positives are harmless.
"""
for cb in litellm.callbacks:
if isinstance(cb, CustomGuardrail) and cb._event_hook_is_event_type(
GuardrailEventHooks.post_call
):
return True
return False
@staticmethod
async def _run_deferred_stream_guardrails(
captured_data: dict,
captured_user_api_key_dict: "UserAPIKeyAuth",
captured_logging_obj: Any,
assembled_response: Any,
cache_hit: Any,
) -> None:
"""
Run only post-call guardrail hooks on an assembled streaming response,
then fire both async and sync logging handlers.
Called by CSW.__anext__ at stream end via a closure stored on
logging_obj._on_deferred_stream_complete.
This is audit-only — content has already been delivered to the client.
Blocking guardrails that raise HTTPException cannot prevent content
delivery for streaming. Per-chunk filtering should use
async_post_call_streaming_hook instead.
Extracted as a static method so tests can call the production
implementation directly rather than reimplementing the closure.
"""
from litellm.litellm_core_utils.thread_pool_executor import executor
_response = assembled_response
try:
from litellm.proxy.proxy_server import llm_router as _global_llm_router
from litellm.proxy.utils import (
_check_and_merge_model_level_guardrails,
unified_guardrail as _unified_guardrail,
)
guardrail_data = _check_and_merge_model_level_guardrails(
data=captured_data, llm_router=_global_llm_router
)
for cb in litellm.callbacks:
if not isinstance(cb, CustomGuardrail):
continue
if not cb.should_run_guardrail(
data=guardrail_data,
event_type=GuardrailEventHooks.post_call,
):
continue
try:
guardrail_result = None
if "apply_guardrail" in type(cb).__dict__:
guardrail_data["guardrail_to_apply"] = cb
guardrail_result = (
await _unified_guardrail.async_post_call_success_hook(
user_api_key_dict=captured_user_api_key_dict,
data=guardrail_data,
response=_response,
)
)
else:
guardrail_result = await cb.async_post_call_success_hook(
user_api_key_dict=captured_user_api_key_dict,
data=guardrail_data,
response=_response,
)
if guardrail_result is not None:
_response = guardrail_result
except Exception as e:
verbose_proxy_logger.exception(
"Error running post-call guardrail %s on streaming response: %s",
getattr(cb, "guardrail_name", type(cb).__name__),
e,
)
if isinstance(e, HTTPException) and hasattr(
captured_logging_obj, "model_call_details"
):
captured_logging_obj.model_call_details.setdefault(
"metadata", {}
)["guardrail_blocked"] = True
except Exception as e:
verbose_proxy_logger.exception(
"Error in deferred streaming guardrail initialization: %s", e,
)
finally:
try:
asyncio.create_task(
captured_logging_obj.async_success_handler(
_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
)
)
except Exception as e:
verbose_proxy_logger.exception(
"Error in deferred streaming async logging: %s", e,
)
try:
executor.submit(
captured_logging_obj.success_handler,
_response,
cache_hit=cache_hit,
start_time=None,
end_time=None,
)
except Exception as e:
verbose_proxy_logger.exception(
"Error in deferred streaming sync logging: %s", e,
)
async def _handle_llm_api_exception(
self,
e: Exception,

View file

@ -32,8 +32,17 @@ def remove_sensitive_info_from_deployment(
deployment_dict["litellm_params"].pop("aws_access_key_id", None)
deployment_dict["litellm_params"].pop("aws_secret_access_key", None)
# Rate-limit config fields must never be masked — they are integers, not credentials.
# The field names contain "key" which matches the masker's sensitive pattern, so we
# explicitly exclude them here rather than widening the global non_sensitive_overrides.
_rate_limit_config_keys = {
"default_api_key_tpm_limit",
"default_api_key_rpm_limit",
}
_excluded = (excluded_keys or set()) | _rate_limit_config_keys
deployment_dict["litellm_params"] = SENSITIVE_DATA_MASKER.mask_dict(
deployment_dict["litellm_params"], excluded_keys=excluded_keys
deployment_dict["litellm_params"], excluded_keys=_excluded
)
return deployment_dict

View file

@ -5,6 +5,7 @@ This file contains the PrismaWrapper class, which is used to wrap the Prisma cli
import asyncio
import os
import random
import signal
import subprocess
import time
import urllib
@ -45,6 +46,46 @@ class PrismaWrapper:
self._reconnection_lock = asyncio.Lock()
self._last_refresh_time: Optional[datetime] = None
def _get_engine_pid(self) -> int:
"""Get the PID of the current Prisma engine subprocess, or 0 if unavailable."""
try:
engine = self._original_prisma._engine
process = getattr(engine, "process", None) if engine is not None else None
if process is not None:
return process.pid
except (AttributeError, TypeError):
pass
return 0
@staticmethod
async def _kill_engine_process(pid: int) -> None:
"""Force-kill an orphaned engine subprocess to prevent DB connection pool leaks.
Called when disconnect() fails and the old engine process may still be
holding open connections. Sends SIGTERM for graceful shutdown, waits
briefly, then SIGKILL as a backstop.
"""
if pid <= 0:
return
try:
os.kill(pid, signal.SIGTERM)
except (ProcessLookupError, PermissionError, OSError):
return # Already dead or inaccessible
verbose_proxy_logger.warning(
"Sent SIGTERM to orphaned prisma-query-engine PID %s after failed disconnect.",
pid,
)
# Brief wait for graceful shutdown, then force-kill
await asyncio.sleep(0.5)
try:
os.kill(pid, getattr(signal, "SIGKILL", signal.SIGTERM))
verbose_proxy_logger.warning(
"Sent SIGKILL to prisma-query-engine PID %s (did not exit after SIGTERM).",
pid,
)
except (ProcessLookupError, PermissionError, OSError):
pass # Exited after SIGTERM — expected
def _extract_token_from_db_url(self, db_url: Optional[str]) -> Optional[str]:
"""
Extract the token (password) from the DATABASE_URL.
@ -179,10 +220,13 @@ class PrismaWrapper:
"""Disconnect and reconnect the Prisma client with a new database URL."""
from prisma import Prisma # type: ignore
old_engine_pid = self._get_engine_pid()
try:
await self._original_prisma.disconnect()
except Exception as e:
verbose_proxy_logger.warning(f"Failed to disconnect Prisma client: {e}")
await self._kill_engine_process(old_engine_pid)
if http_client is not None:
self._original_prisma = Prisma(http=http_client)

View file

@ -27,10 +27,17 @@ async def get_ui_config():
admin_ui_disabled = os.getenv("DISABLE_ADMIN_UI", "false").lower() == "true"
sso_configured = _has_user_setup_sso()
from litellm.proxy.proxy_server import proxy_config
is_control_plane = len(proxy_config.worker_registry) > 0
return UiDiscoveryEndpoints(
server_root_path=get_server_root_path(),
proxy_base_url=get_proxy_base_url(),
auto_redirect_to_sso=sso_configured and auto_redirect_ui_login_to_sso,
admin_ui_disabled=admin_ui_disabled,
sso_configured=sso_configured,
is_control_plane=is_control_plane,
workers=proxy_config.worker_registry if is_control_plane else [],
)

View file

@ -295,16 +295,17 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
)
# Check if request under RPM/TPM per model for a given API Key
if (
get_key_model_tpm_limit(user_api_key_dict) is not None
or get_key_model_rpm_limit(user_api_key_dict) is not None
):
_model = data.get("model", None)
_model = data.get("model", None)
_tpm_limit_for_key_model = get_key_model_tpm_limit(
user_api_key_dict, model_name=_model
)
_rpm_limit_for_key_model = get_key_model_rpm_limit(
user_api_key_dict, model_name=_model
)
if _tpm_limit_for_key_model is not None or _rpm_limit_for_key_model is not None:
request_count_api_key = (
f"{api_key}::{_model}::{precise_minute}::request_count"
)
_tpm_limit_for_key_model = get_key_model_tpm_limit(user_api_key_dict)
_rpm_limit_for_key_model = get_key_model_rpm_limit(user_api_key_dict)
tpm_limit_for_model = None
rpm_limit_for_model = None
@ -477,6 +478,15 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
kwargs["litellm_params"]["metadata"].get("user_api_key_metadata", {})
or {}
)
user_api_key_team_metadata = kwargs["litellm_params"]["metadata"].get(
"user_api_key_team_metadata", None
)
user_api_key_dict = UserAPIKeyAuth(
api_key=user_api_key,
metadata=user_api_key_metadata,
model_max_budget=user_api_key_model_max_budget,
team_metadata=user_api_key_team_metadata,
)
# ------------
# Setup values
@ -538,6 +548,16 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
# Update usage - model group + API Key
# ------------
model_group = get_model_group_from_litellm_kwargs(kwargs)
_success_tpm_limit = (
get_key_model_tpm_limit(user_api_key_dict, model_name=model_group)
if model_group is not None
else None
)
_success_rpm_limit = (
get_key_model_rpm_limit(user_api_key_dict, model_name=model_group)
if model_group is not None
else None
)
if (
user_api_key is not None
and model_group is not None
@ -545,6 +565,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
"model_rpm_limit" in user_api_key_metadata
or "model_tpm_limit" in user_api_key_metadata
or user_api_key_model_max_budget is not None
or _success_tpm_limit is not None
or _success_rpm_limit is not None
)
):
request_count_api_key = (

View file

@ -687,8 +687,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if not requested_model:
return
_tpm_limit_for_key_model = get_key_model_tpm_limit(user_api_key_dict)
_rpm_limit_for_key_model = get_key_model_rpm_limit(user_api_key_dict)
_tpm_limit_for_key_model = get_key_model_tpm_limit(
user_api_key_dict, model_name=requested_model
)
_rpm_limit_for_key_model = get_key_model_rpm_limit(
user_api_key_dict, model_name=requested_model
)
if _tpm_limit_for_key_model is None and _rpm_limit_for_key_model is None:
return

View file

@ -1214,17 +1214,9 @@ if MCP_AVAILABLE:
"error": "User does not have permission to create mcp servers. You can only create mcp servers if you are a PROXY_ADMIN."
},
)
elif payload.server_id is not None:
# fail if the mcp server with id already exists
mcp_server = await get_mcp_server(prisma_client, payload.server_id)
if mcp_server is not None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": f"MCP Server with id {payload.server_id} already exists. Cannot create another."
},
)
elif (
# Block reserved special server IDs
if (
SpecialMCPServerName.all_team_servers == payload.server_id
or SpecialMCPServerName.all_proxy_servers == payload.server_id
):
@ -1235,6 +1227,17 @@ if MCP_AVAILABLE:
},
)
if payload.server_id is not None:
# fail if the mcp server with id already exists
mcp_server = await get_mcp_server(prisma_client, payload.server_id)
if mcp_server is not None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": f"MCP Server with id {payload.server_id} already exists. Cannot create another."
},
)
# TODO: audit log for create
# Admin-created servers are always active — clear any submission lifecycle

View file

@ -724,6 +724,7 @@ async def new_team( # noqa: PLR0915
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission.
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
- team_member_budget_duration: Optional[str] - The duration of the budget for the team member. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
- team_member_rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for individual team members.
- team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.
- team_member_key_duration: Optional[str] - The duration for a team member's key. e.g. "1d", "1w", "1mo"
@ -934,6 +935,7 @@ async def new_team( # noqa: PLR0915
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,
):
data_json = await TeamMemberBudgetHandler.create_team_member_budget_table(
data=data,
@ -942,6 +944,7 @@ async def new_team( # noqa: PLR0915
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,
)
## ADD TO TEAM TABLE

View file

@ -16,6 +16,7 @@ import os
import secrets
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
from urllib.parse import urlencode, urlparse
if TYPE_CHECKING:
import httpx
@ -301,6 +302,7 @@ async def google_login(
source: Optional[str] = None,
key: Optional[str] = None,
existing_key: Optional[str] = None,
return_to: Optional[str] = None,
): # noqa: PLR0915
"""
Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env
@ -394,13 +396,23 @@ async def google_login(
is True
):
verbose_proxy_logger.info(f"Redirecting to SSO login for {redirect_url}")
return await SSOAuthenticationHandler.get_sso_login_redirect(
sso_redirect = await SSOAuthenticationHandler.get_sso_login_redirect(
redirect_url=redirect_url,
microsoft_client_id=microsoft_client_id,
google_client_id=google_client_id,
generic_client_id=generic_client_id,
state=cli_state,
)
if return_to is not None and sso_redirect is not None:
SSOAuthenticationHandler._validate_return_to(return_to)
sso_redirect.set_cookie(
key="litellm_cp_return_to",
value=return_to,
max_age=600,
httponly=True,
samesite="lax",
)
return sso_redirect
elif ui_username is not None:
# No Google, Microsoft SSO
# Use UI Credentials set in .env
@ -1312,12 +1324,17 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
request=request, key=key_id, existing_key=existing_key, result=result
)
# Control-plane cross-origin: read return_to from cookie.
# Starlette's cookie_parser already handles RFC 2109 unquoting.
cp_return_to: Optional[str] = request.cookies.get("litellm_cp_return_to")
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=result,
request=request,
received_response=received_response,
generic_client_id=generic_client_id,
ui_access_mode=ui_access_mode,
return_to=cp_return_to,
)
@ -1760,6 +1777,38 @@ class SSOAuthenticationHandler:
Handler for SSO Authentication across all SSO providers
"""
@staticmethod
def _validate_return_to(return_to: str) -> None:
"""
Validate that return_to matches the configured control_plane_url origin.
Raises HTTPException(400) if:
- control_plane_url is not configured in general_settings
- return_to origin does not match control_plane_url origin
"""
from litellm.proxy.proxy_server import general_settings
control_plane_url = general_settings.get("control_plane_url")
if control_plane_url is None:
raise HTTPException(
status_code=400,
detail="return_to is not allowed: control_plane_url is not configured",
)
def _origin(url: str) -> tuple:
parsed = urlparse(url)
scheme = (parsed.scheme or "").lower()
hostname = (parsed.hostname or "").lower()
default_port = 443 if scheme == "https" else 80
port = parsed.port if parsed.port is not None else default_port
return (scheme, hostname, port)
if _origin(return_to) != _origin(control_plane_url):
raise HTTPException(
status_code=400,
detail="return_to does not match the configured control_plane_url",
)
@staticmethod
async def get_sso_login_redirect(
redirect_url: str,
@ -2358,6 +2407,7 @@ class SSOAuthenticationHandler:
received_response: Optional[dict] = None,
generic_client_id: Optional[str] = None,
ui_access_mode: Optional[Dict] = None,
return_to: Optional[str] = None,
) -> RedirectResponse:
import jwt
@ -2367,6 +2417,7 @@ class SSOAuthenticationHandler:
master_key,
premium_user,
proxy_logging_obj,
redis_usage_cache,
user_api_key_cache,
user_custom_sso,
)
@ -2534,6 +2585,36 @@ class SSOAuthenticationHandler:
master_key or "",
algorithm="HS256",
)
# Control-plane cross-origin: store JWT behind a single-use opaque
# code (60s TTL) so the token never appears in browser history / logs.
# The control plane redeems it via POST /v3/login/exchange.
if return_to is not None:
SSOAuthenticationHandler._validate_return_to(return_to)
code = secrets.token_urlsafe(32)
cache_key = f"login_code:{code}"
cache_value = {"token": jwt_token, "redirect_url": return_to}
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(
key=cache_key, value=cache_value, ttl=60
)
else:
await user_api_key_cache.async_set_cache(
key=cache_key, value=cache_value, ttl=60
)
separator = "&" if "?" in return_to else "?"
redirect_url = (
return_to + separator + urlencode({"login": "success", "code": code})
)
verbose_proxy_logger.info(
"Cross-origin SSO: redirecting to control plane with login code"
)
redirect_response = RedirectResponse(url=redirect_url, status_code=303)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
if user_id is not None and isinstance(user_id, str):
litellm_dashboard_ui += "?login=success"
verbose_proxy_logger.info(f"Redirecting to {litellm_dashboard_ui}")

View file

@ -541,6 +541,7 @@ from litellm.types.llms.anthropic import (
AnthropicResponseUsageBlock,
)
from litellm.types.llms.openai import HttpxBinaryResponseContent
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
ModelGroupInfoProxy,
)
@ -1546,6 +1547,7 @@ user_custom_key_generate = None
# Sentinel: prevents PKCE-no-Redis advisory from re-logging on config hot-reload.
# Tests that need to reset it can patch 'litellm.proxy.proxy_server._pkce_no_redis_warning_emitted'.
_pkce_no_redis_warning_emitted: bool = False
_cp_no_redis_warning_emitted: bool = False
user_custom_sso = None
user_custom_ui_sso_sign_in_handler = None
use_background_health_checks = None
@ -2295,6 +2297,7 @@ class ProxyConfig:
self.config: Dict[str, Any] = {}
self._last_semantic_filter_config: Optional[Dict[str, Any]] = None
self._last_hashicorp_vault_config: Optional[Dict[str, Any]] = None
self.worker_registry: List["WorkerRegistryEntry"] = []
def is_yaml(self, config_file_path: str) -> bool:
if not os.path.isfile(config_file_path):
@ -3095,6 +3098,21 @@ class ProxyConfig:
"Set PKCE_STRICT_CACHE_MISS=true to fail fast with a 401 on cache misses "
"instead of continuing without a code_verifier."
)
### CONTROL PLANE CODE-EXCHANGE PREREQUISITE CHECK ###
cp_url = general_settings.get("control_plane_url")
if cp_url and redis_usage_cache is None:
global _cp_no_redis_warning_emitted
if not _cp_no_redis_warning_emitted:
_cp_no_redis_warning_emitted = True
verbose_proxy_logger.warning(
"control_plane_url is configured but Redis is not configured for LiteLLM caching. "
"Login codes (SSO and /v3/login) will not be shared across instances — "
"the /v3/login/exchange call may land on a different pod and fail with 401. "
"Configure Redis via the 'cache' section in your proxy config, "
"or ensure sticky sessions for single-instance deployments."
)
### STORE MODEL IN DB ### feature flag for `/model/new`
store_model_in_db = general_settings.get("store_model_in_db", False)
if store_model_in_db is None:
@ -3385,7 +3403,15 @@ class ProxyConfig:
litellm.vector_store_registry.load_vector_stores_from_config(
vector_store_registry_config
)
pass
## WORKER REGISTRY (Control Plane)
worker_registry_config = config.get("worker_registry", None)
if worker_registry_config:
self.worker_registry = [
WorkerRegistryEntry(**e) for e in worker_registry_config
]
else:
self.worker_registry = []
async def _init_policy_engine(
self,
@ -11095,6 +11121,165 @@ async def login_v2(request: Request): # noqa: PLR0915
)
@router.post(
"/v3/login", include_in_schema=False
) # control-plane login — always returns token in body for cross-origin use
async def login_v3(request: Request): # noqa: PLR0915
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.utils import get_custom_url
try:
if not general_settings.get("control_plane_url"):
raise ProxyException(
message="/v3/login is only available on workers with control_plane_url configured",
type=ProxyErrorTypes.not_found_error,
param="control_plane_url",
code=status.HTTP_404_NOT_FOUND,
)
body = await request.json()
username = str(body.get("username"))
password = str(body.get("password"))
login_result = await authenticate_user(
username=username,
password=password,
master_key=master_key,
prisma_client=prisma_client,
)
returned_ui_token_object = create_ui_token_object(
login_result=login_result,
general_settings=general_settings,
premium_user=premium_user,
)
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
litellm_dashboard_ui += "ui/"
else:
litellm_dashboard_ui += "/ui/"
litellm_dashboard_ui += "?login=success"
# Store JWT behind a single-use opaque code (60s TTL)
code = secrets.token_urlsafe(32)
cache_key = f"login_code:{code}"
cache_value = {"token": jwt_token, "redirect_url": litellm_dashboard_ui}
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(
key=cache_key, value=cache_value, ttl=60
)
else:
await user_api_key_cache.async_set_cache(
key=cache_key, value=cache_value, ttl=60
)
return JSONResponse(
content={"code": code, "expires_in": 60},
status_code=status.HTTP_200_OK,
)
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.login_v3(): Exception occurred - {}".format(
str(e)
)
)
if isinstance(e, ProxyException):
raise e
elif isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", str(e)),
type=ProxyErrorTypes.auth_error,
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
)
else:
error_msg = f"{str(e)}"
raise ProxyException(
message=error_msg,
type=ProxyErrorTypes.auth_error,
param="None",
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@router.post(
"/v3/login/exchange", include_in_schema=False
) # exchange single-use opaque code for JWT
async def login_v3_exchange(request: Request):
try:
if not general_settings.get("control_plane_url"):
raise ProxyException(
message="/v3/login/exchange is only available on workers with control_plane_url configured",
type=ProxyErrorTypes.not_found_error,
param="control_plane_url",
code=status.HTTP_404_NOT_FOUND,
)
body = await request.json()
code = body.get("code")
if not code:
raise ProxyException(
message="Missing 'code' parameter",
type=ProxyErrorTypes.auth_error,
param="code",
code=status.HTTP_400_BAD_REQUEST,
)
cache_key = f"login_code:{code}"
if redis_usage_cache is not None:
cached_data = await redis_usage_cache.async_get_cache(key=cache_key)
else:
cached_data = await user_api_key_cache.async_get_cache(key=cache_key)
if not cached_data or not isinstance(cached_data, dict):
raise ProxyException(
message="Invalid or expired login code",
type=ProxyErrorTypes.auth_error,
param="code",
code=status.HTTP_401_UNAUTHORIZED,
)
# Single-use: delete immediately
if redis_usage_cache is not None:
await redis_usage_cache.async_delete_cache(key=cache_key)
else:
await user_api_key_cache.async_delete_cache(key=cache_key)
json_response = JSONResponse(
content={
"token": cached_data["token"],
"redirect_url": cached_data["redirect_url"],
},
status_code=status.HTTP_200_OK,
)
json_response.set_cookie(key="token", value=cached_data["token"])
return json_response
except ProxyException:
raise
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.login_v3_exchange(): Exception occurred - {}".format(
str(e)
)
)
raise ProxyException(
message=str(e),
type=ProxyErrorTypes.auth_error,
param="None",
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
@app.get("/onboarding/get_token", include_in_schema=False)
async def onboarding(invite_link: str, request: Request):
"""

View file

@ -1922,6 +1922,7 @@ class ProxyLogging:
original_exception,
traceback.format_exc(),
),
daemon=True,
).start()
async def post_call_success_hook(
@ -4010,13 +4011,15 @@ class PrismaClient:
)
async def _do_direct_reconnect() -> None:
old_pid = self._get_engine_pid()
try:
await self.db.disconnect()
except Exception as disconnect_err:
verbose_proxy_logger.debug(
"Prisma DB disconnect before reconnect failed (ignored): %s",
verbose_proxy_logger.warning(
"Prisma DB disconnect before reconnect failed: %s",
disconnect_err,
)
await PrismaWrapper._kill_engine_process(old_pid)
await self.db.connect()
await self.db.query_raw("SELECT 1")

View file

@ -16,4 +16,13 @@ class CacheControlMessageInjectionPoint(TypedDict):
control: Optional[ChatCompletionCachedContent]
CacheControlInjectionPoint = CacheControlMessageInjectionPoint
class CacheControlToolConfigInjectionPoint(TypedDict):
"""Type for tool_config-level injection points (Bedrock)."""
location: Literal["tool_config"]
CacheControlInjectionPoint = Union[
CacheControlMessageInjectionPoint,
CacheControlToolConfigInjectionPoint,
]

View file

@ -58,12 +58,26 @@ class HttpxBlobType(TypedDict, total=False):
data: str
class HttpxServerSideToolCall(TypedDict, total=False):
toolType: str
id: str
args: dict
class HttpxServerSideToolResponse(TypedDict, total=False):
toolType: str
id: str
response: Union[str, dict]
class HttpxPartType(TypedDict, total=False):
text: str
inlineData: HttpxBlobType
fileData: FileDataType
functionCall: HttpxFunctionCall
functionResponse: FunctionResponse
toolCall: HttpxServerSideToolCall
toolResponse: HttpxServerSideToolResponse
executableCode: HttpxExecutableCode
codeExecutionResult: HttpxCodeExecutionResult
thought: bool
@ -244,8 +258,9 @@ class Tools(TypedDict, total=False):
retrieval: Retrieval
class ToolConfig(TypedDict):
class ToolConfig(TypedDict, total=False):
functionCallingConfig: FunctionCallingConfig
includeServerSideToolInvocations: bool
class TTL(TypedDict, total=False):

View file

@ -0,0 +1,14 @@
from pydantic import BaseModel, field_validator
class WorkerRegistryEntry(BaseModel):
worker_id: str
name: str
url: str
@field_validator("url")
@classmethod
def url_must_be_http(cls, v: str) -> str:
if not v.startswith(("http://", "https://")):
raise ValueError("Worker URL must start with http:// or https://")
return v

View file

@ -1,7 +1,9 @@
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
class UiDiscoveryEndpoints(BaseModel):
server_root_path: str
@ -9,3 +11,5 @@ class UiDiscoveryEndpoints(BaseModel):
auto_redirect_to_sso: bool
admin_ui_disabled: bool
sso_configured: bool
is_control_plane: bool = False
workers: List[WorkerRegistryEntry] = []

View file

@ -188,6 +188,11 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
max_file_size_mb: Optional[float] = None
# Proxy-wide default rate limits applied to any API key using this deployment
# when the key does not have a model-specific tpm/rpm limit configured.
default_api_key_tpm_limit: Optional[int] = None
default_api_key_rpm_limit: Optional[int] = None
# Deployment budgets
max_budget: Optional[float] = None
budget_duration: Optional[str] = None

View file

@ -1944,15 +1944,38 @@ def client(original_function): # noqa: PLR0915
)
# LOG SUCCESS - handle streaming success logging in the _next_ object
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
# NOTE: streaming requests return early (before this point) via
# CustomStreamWrapper, so this block is non-streaming only.
if getattr(logging_obj, "_defer_async_logging", False):
# Proxy has post-call guardrails that must complete before the
# SLP is built. Store a closure the proxy will call after
# post_call_success_hook so guardrail_information is in metadata.
# Only create_task is deferred; sync callbacks fire immediately
# (below, outside the if/else) for billing/rate-limiting.
def _enqueue_deferred_logging() -> None:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging # type: ignore
else:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
)
# Sync callbacks always fire immediately regardless of deferral
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,

View file

@ -6152,7 +6152,8 @@
"max_query_tokens": 4096,
"max_tokens": 32768,
"mode": "rerank",
"output_cost_per_token": 0.0
"output_cost_per_token": 0.0,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076"
},
"azure_ai/cohere-rerank-v4.0-fast": {
"input_cost_per_query": 0.002,
@ -6163,7 +6164,8 @@
"max_query_tokens": 4096,
"max_tokens": 32768,
"mode": "rerank",
"output_cost_per_token": 0.0
"output_cost_per_token": 0.0,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076"
},
"azure_ai/deepseek-v3.2": {
"input_cost_per_token": 5.8e-07,
@ -6173,6 +6175,7 @@
"max_tokens": 163840,
"mode": "chat",
"output_cost_per_token": 1.68e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -6187,6 +6190,7 @@
"max_tokens": 163840,
"mode": "chat",
"output_cost_per_token": 1.68e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,

8
poetry.lock generated
View file

@ -3209,15 +3209,15 @@ openai = ["openai (>=0.27.8)"]
[[package]]
name = "litellm-enterprise"
version = "0.1.33"
version = "0.1.35"
description = "Package for LiteLLM Enterprise features"
optional = true
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
{file = "litellm_enterprise-0.1.33-py3-none-any.whl", hash = "sha256:ae262ecfca680a235095becd6215e412e5ceba90efef739e61e6096b121188a2"},
{file = "litellm_enterprise-0.1.33.tar.gz", hash = "sha256:5e3c0de9c4b54694ebb3017c8e18ee1d40e02ebef86e9ebd9c006e445885d5a0"},
{file = "litellm_enterprise-0.1.35-py3-none-any.whl", hash = "sha256:8d2d9c925de8ee35e308c0f4975483b60f5e22beb50506e261e555e466f019c5"},
{file = "litellm_enterprise-0.1.35.tar.gz", hash = "sha256:b752d07e538424743fcc08ba0d3d9d83d1f04a45c115811ac7828d789b6d87cc"},
]
[[package]]
@ -8018,4 +8018,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "2cf958f1a04fd5f1ab0e5cfc33bdbf441b518ed6c82d0f2546bf64cd3d2f89be"
content-hash = "f0977419272b446bc2df0e062406c8f7fe03566bd38fdb8395418bdc6da3fe20"

View file

@ -63,7 +63,7 @@ mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "^0.4.58", optional = true}
rich = {version = "^13.7.1", optional = true}
litellm-enterprise = {version = "^0.1.33", optional = true}
litellm-enterprise = {version = "0.1.35", optional = true}
diskcache = {version = "^5.6.1", optional = true}
polars = {version = "^1.31.0", optional = true, python = ">=3.10"}
semantic-router = {version = ">=0.1.12", optional = true, python = ">=3.9,<3.14"}

View file

@ -80,4 +80,4 @@ pypdf>=6.7.3 # for PDF text extraction in RAG ingestion (CVE-2026-27888)
########################
# LITELLM ENTERPRISE DEPENDENCIES
########################
litellm-enterprise==0.1.34
litellm-enterprise==0.1.35

View file

@ -22,6 +22,7 @@ from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation
_is_multimodal_input,
_parse_data_url,
process_embed_content_response,
transform_openai_input_gemini_content,
transform_openai_input_gemini_embed_content,
)
from litellm.types.utils import EmbeddingResponse
@ -396,6 +397,44 @@ def test_transform_with_optional_params():
assert result["taskType"] == "SEMANTIC_SIMILARITY"
def test_task_type_mapped_to_camel_case_batch():
"""Test that snake_case task_type is converted to camelCase taskType for batchEmbedContents."""
result = transform_openai_input_gemini_content(
input="test text",
model="text-embedding-004",
optional_params={"task_type": "RETRIEVAL_DOCUMENT"},
)
for request in result["requests"]:
assert "taskType" in request
assert request["taskType"] == "RETRIEVAL_DOCUMENT"
assert "task_type" not in request
def test_task_type_mapped_to_camel_case_embed_content():
"""Test that snake_case task_type is converted to camelCase taskType for embedContent."""
result = transform_openai_input_gemini_embed_content(
input=["test text"],
model="gemini-embedding-2-preview",
optional_params={"task_type": "RETRIEVAL_DOCUMENT"},
resolved_files=None,
)
assert "taskType" in result
assert result["taskType"] == "RETRIEVAL_DOCUMENT"
assert "task_type" not in result
def test_task_type_camel_case_passthrough():
"""Test that camelCase taskType passed directly is preserved."""
result = transform_openai_input_gemini_embed_content(
input=["test text"],
model="gemini-embedding-2-preview",
optional_params={"taskType": "SEMANTIC_SIMILARITY"},
resolved_files=None,
)
assert result["taskType"] == "SEMANTIC_SIMILARITY"
assert "task_type" not in result
def test_dimensions_mapped_to_output_dimensionality():
"""Test that OpenAI 'dimensions' param is mapped to Gemini 'outputDimensionality'."""
input_data = ["test text"]

View file

@ -11,6 +11,7 @@ counting, the test will be skipped.
import os
import sys
from typing import Any, Dict, List
from unittest.mock import patch
import pytest
@ -99,3 +100,66 @@ class TestBedrockTokenCounter(BaseTokenCounterTest):
assert result.total_tokens > 0, f"Token count should be > 0, got {result.total_tokens}"
assert result.tokenizer_type is not None, "tokenizer_type should be set"
assert result.error is not True, f"Token counting should not error: {result.error_message}"
class TestBedrockCountTokensEndpoint:
"""Unit tests for custom endpoint URL resolution in BedrockCountTokensConfig."""
def _make_handler(self):
from litellm.llms.bedrock.count_tokens.transformation import (
BedrockCountTokensConfig,
)
return BedrockCountTokensConfig()
def test_default_endpoint(self):
handler = self._make_handler()
url = handler.get_bedrock_count_tokens_endpoint(
model="amazon.nova-lite-v1:0",
aws_region_name="us-east-1",
)
assert url == "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.nova-lite-v1:0/count-tokens"
def test_api_base_overrides_default(self):
handler = self._make_handler()
custom_base = "https://vpce-xxx.bedrock-runtime.us-east-1.vpce.amazonaws.com"
url = handler.get_bedrock_count_tokens_endpoint(
model="amazon.nova-lite-v1:0",
aws_region_name="us-east-1",
api_base=custom_base,
)
assert url == f"{custom_base}/model/amazon.nova-lite-v1:0/count-tokens"
def test_aws_bedrock_runtime_endpoint_overrides_default(self):
handler = self._make_handler()
custom_endpoint = "https://vpce-yyy.bedrock-runtime.eu-west-1.vpce.amazonaws.com"
url = handler.get_bedrock_count_tokens_endpoint(
model="amazon.nova-lite-v1:0",
aws_region_name="eu-west-1",
aws_bedrock_runtime_endpoint=custom_endpoint,
)
assert url == f"{custom_endpoint}/model/amazon.nova-lite-v1:0/count-tokens"
def test_api_base_takes_priority_over_aws_bedrock_runtime_endpoint(self):
handler = self._make_handler()
api_base = "https://api-base.example.com"
runtime_endpoint = "https://runtime-endpoint.example.com"
url = handler.get_bedrock_count_tokens_endpoint(
model="amazon.nova-lite-v1:0",
aws_region_name="us-east-1",
api_base=api_base,
aws_bedrock_runtime_endpoint=runtime_endpoint,
)
assert url == f"{api_base}/model/amazon.nova-lite-v1:0/count-tokens"
def test_env_var_overrides_default(self, monkeypatch):
monkeypatch.setenv(
"AWS_BEDROCK_RUNTIME_ENDPOINT",
"https://env-endpoint.bedrock-runtime.us-west-2.amazonaws.com",
)
handler = self._make_handler()
url = handler.get_bedrock_count_tokens_endpoint(
model="amazon.nova-lite-v1:0",
aws_region_name="us-west-2",
)
assert url.startswith("https://env-endpoint.bedrock-runtime.us-west-2.amazonaws.com")

View file

@ -0,0 +1,384 @@
"""
Unit tests for WebSearch Short-Circuit
Tests the short-circuit path that detects web-search-only /v1/messages requests
and executes the search directly without routing through the backend LLM.
"""
from unittest.mock import AsyncMock, patch
import pytest
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
# ---------------------------------------------------------------------------
# Detection tests
# ---------------------------------------------------------------------------
class TestTryShortCircuitSearch:
"""Tests for WebSearchInterceptionLogger.try_short_circuit_search"""
@pytest.mark.asyncio
async def test_short_circuits_single_web_search_tool(self):
"""Single web_search_20250305 tool → short-circuit fires"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.return_value = (
"Title: Result\nURL: https://example.com\nSnippet: test"
)
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[
{"role": "user", "content": "Search for Claude Code releases"}
],
tools=[
{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}
],
custom_llm_provider="github_copilot",
)
assert result is not None
assert result["type"] == "message"
assert result["role"] == "assistant"
assert result["stop_reason"] == "end_turn"
assert len(result["content"]) == 1
assert result["content"][0]["type"] == "text"
assert "Result" in result["content"][0]["text"]
mock_search.assert_called_once_with("Search for Claude Code releases")
@pytest.mark.asyncio
async def test_does_not_short_circuit_mixed_tools(self):
"""Mix of web_search and other tools → NOT short-circuited"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "Do something"}],
tools=[
{"type": "web_search_20250305", "name": "web_search", "max_uses": 8},
{"name": "Read", "description": "Read a file", "input_schema": {}},
],
custom_llm_provider="github_copilot",
)
assert result is None
@pytest.mark.asyncio
async def test_does_not_short_circuit_no_tools(self):
"""No tools → NOT short-circuited"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "Hello"}],
tools=None,
custom_llm_provider="github_copilot",
)
assert result is None
@pytest.mark.asyncio
async def test_does_not_short_circuit_empty_tools(self):
"""Empty tools list → NOT short-circuited"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "Hello"}],
tools=[],
custom_llm_provider="github_copilot",
)
assert result is None
@pytest.mark.asyncio
async def test_does_not_short_circuit_wrong_provider(self):
"""Provider not in enabled_providers → NOT short-circuited"""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "Search for something"}],
tools=[
{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}
],
custom_llm_provider="github_copilot",
)
assert result is None
@pytest.mark.asyncio
async def test_does_not_short_circuit_bedrock(self):
"""Bedrock has native agentic loop support → NOT short-circuited.
Providers with a BaseAnthropicMessagesConfig (bedrock, vertex_ai, etc.)
use the agentic loop which includes a follow-up LLM synthesis step.
The short-circuit must not fire for them.
"""
logger = WebSearchInterceptionLogger(
enabled_providers=["bedrock", "github_copilot"]
)
result = await logger.try_short_circuit_search(
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=[{"role": "user", "content": "Search for something"}],
tools=[
{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}
],
custom_llm_provider="bedrock",
)
assert result is None
@pytest.mark.asyncio
async def test_does_not_short_circuit_no_messages(self):
"""Empty messages → NOT short-circuited"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[],
tools=[
{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}
],
custom_llm_provider="github_copilot",
)
assert result is None
@pytest.mark.asyncio
async def test_search_failure_returns_error_text(self):
"""Search failure → response with error message, not exception"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.side_effect = RuntimeError("Tavily API error")
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "Search for something"}],
tools=[
{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}
],
custom_llm_provider="github_copilot",
)
assert result is not None
assert "Search failed" in result["content"][0]["text"]
@pytest.mark.asyncio
async def test_response_has_valid_structure(self):
"""Synthetic response has all required AnthropicMessagesResponse fields"""
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.return_value = "search results here"
result = await logger.try_short_circuit_search(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "Search query"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
)
assert result is not None
# Required fields
assert "id" in result
assert result["id"].startswith("msg_")
assert result["type"] == "message"
assert result["role"] == "assistant"
assert result["model"] == "github_copilot/claude-sonnet-4"
assert result["stop_reason"] == "end_turn"
assert result["stop_sequence"] is None
assert "usage" in result
assert "content" in result
# ---------------------------------------------------------------------------
# Query extraction tests
# ---------------------------------------------------------------------------
# ---------------------------------------------------------------------------
# Integration with entry point
# ---------------------------------------------------------------------------
class TestShortCircuitEntryPoint:
"""Tests for _try_websearch_short_circuit in the /v1/messages handler"""
@pytest.mark.asyncio
async def test_returns_none_when_no_callbacks(self):
"""No callbacks configured → returns None"""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
_try_websearch_short_circuit,
)
with patch("litellm.callbacks", []):
result = await _try_websearch_short_circuit(
model="test",
messages=[],
tools=[],
custom_llm_provider="github_copilot",
stream=False,
)
assert result is None
@pytest.mark.asyncio
async def test_returns_dict_when_not_streaming(self):
"""Non-streaming short-circuit → returns dict"""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
_try_websearch_short_circuit,
)
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.return_value = "results"
with patch("litellm.callbacks", [logger]):
result = await _try_websearch_short_circuit(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "search query"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
stream=False,
)
assert isinstance(result, dict)
assert result["content"][0]["text"] == "results"
@pytest.mark.asyncio
async def test_returns_stream_iterator_when_streaming(self):
"""Streaming short-circuit → returns FakeAnthropicMessagesStreamIterator"""
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
_try_websearch_short_circuit,
)
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.return_value = "streaming results"
with patch("litellm.callbacks", [logger]):
result = await _try_websearch_short_circuit(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "search query"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
stream=True,
)
assert isinstance(result, FakeAnthropicMessagesStreamIterator)
# Verify stream produces valid SSE events
chunks = []
async for chunk in result:
chunks.append(chunk)
assert len(chunks) > 0
# First chunk should be message_start
assert b"event: message_start" in chunks[0]
# Last chunk should be message_stop
assert b"event: message_stop" in chunks[-1]
# Should contain the search results text
all_data = b"".join(chunks)
assert b"streaming results" in all_data
@pytest.mark.asyncio
async def test_skips_non_websearch_callbacks(self):
"""Non-WebSearchInterceptionLogger callbacks are ignored"""
from unittest.mock import MagicMock
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
_try_websearch_short_circuit,
)
other_callback = MagicMock()
with patch("litellm.callbacks", [other_callback]):
result = await _try_websearch_short_circuit(
model="test",
messages=[{"role": "user", "content": "search"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
stream=False,
)
assert result is None
@pytest.mark.asyncio
async def test_uses_original_stream_not_hook_converted(self):
"""Verify that the entry point passes original_stream to the short-circuit.
The pre-request hook converts stream=True → stream=False for the agentic
loop. The short-circuit must use the ORIGINAL stream value so streaming
callers get SSE events instead of a plain dict.
"""
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
_try_websearch_short_circuit,
)
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.return_value = "streaming results"
with patch("litellm.callbacks", [logger]):
# Simulate what anthropic_messages() does: original_stream=True
# is passed to the short-circuit, even though the hook would have
# already converted stream to False in request_kwargs.
result = await _try_websearch_short_circuit(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "search query"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
stream=True, # original_stream, NOT the hook-converted value
)
# Must return a stream iterator, not a plain dict
assert isinstance(result, FakeAnthropicMessagesStreamIterator)
@pytest.mark.asyncio
async def test_short_circuits_with_provider_from_model_string(self):
"""Provider embedded in model string (custom_llm_provider=None) should
still fire the short-circuit when the caller propagates the derived
provider.
"""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
_try_websearch_short_circuit,
)
logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"])
with patch.object(
logger, "_execute_search", new_callable=AsyncMock
) as mock_search:
mock_search.return_value = "results"
with patch("litellm.callbacks", [logger]):
# Simulate the caller having derived custom_llm_provider from
# the model string before calling _try_websearch_short_circuit
result = await _try_websearch_short_circuit(
model="github_copilot/claude-sonnet-4",
messages=[{"role": "user", "content": "search query"}],
tools=[{"type": "web_search_20250305", "name": "web_search"}],
custom_llm_provider="github_copilot",
stream=False,
)
assert result is not None
assert result["content"][0]["text"] == "results"

View file

@ -2248,3 +2248,54 @@ def test_failure_handler_runs_sync_callbacks_for_non_pass_through_requests(
)
dummy_logger.log_failure_event.assert_called_once()
def test_merge_hidden_params_from_response_into_metadata_populates_metadata():
"""Streaming completion path should mirror non-stream: metadata.hidden_params from response."""
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
logging_obj = LiteLLMLoggingObj(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="acompletion",
start_time=time.time(),
litellm_call_id="merge-hp-test",
function_id="merge-hp-fn",
)
logging_obj.model_call_details = {
"litellm_params": {"metadata": {}},
}
class _Resp:
_hidden_params = {"response_cost": 0.001, "model_id": "mid-test"}
logging_obj._merge_hidden_params_from_response_into_metadata(_Resp())
meta = logging_obj.model_call_details["litellm_params"]["metadata"]
assert meta["hidden_params"]["response_cost"] == 0.001
assert meta["hidden_params"]["model_id"] == "mid-test"
def test_merge_hidden_params_from_response_into_metadata_no_op_when_empty():
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
logging_obj = LiteLLMLoggingObj(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="acompletion",
start_time=time.time(),
litellm_call_id="merge-hp-empty",
function_id="merge-hp-empty-fn",
)
logging_obj.model_call_details = {
"litellm_params": {"metadata": {"existing": True}},
}
class _NoHp:
_hidden_params = {}
logging_obj._merge_hidden_params_from_response_into_metadata(_NoHp())
assert "hidden_params" not in logging_obj.model_call_details["litellm_params"][
"metadata"
]

View file

@ -1984,3 +1984,130 @@ def test_translate_anthropic_to_openai_with_mixed_tools():
# tool_name_mapping should be empty for short tool names
assert tool_name_mapping == {}
class TestTranslateAnthropicOutputFormatToOpenAI:
"""Tests for translate_anthropic_output_format_to_openai adding additionalProperties: false."""
def setup_method(self):
self.adapter = LiteLLMAnthropicMessagesAdapter()
def test_simple_object_adds_additional_properties_false(self):
output_format = {
"type": "json_schema",
"schema": {
"type": "object",
"properties": {"name": {"type": "string"}},
},
}
result = self.adapter.translate_anthropic_output_format_to_openai(output_format)
assert result is not None
schema = result["json_schema"]["schema"]
assert schema["additionalProperties"] is False
assert schema["required"] == ["name"]
def test_nested_objects_adds_additional_properties_false(self):
output_format = {
"type": "json_schema",
"schema": {
"type": "object",
"properties": {
"user": {
"type": "object",
"properties": {
"name": {"type": "string"},
"address": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
}
},
},
}
result = self.adapter.translate_anthropic_output_format_to_openai(output_format)
assert result is not None
schema = result["json_schema"]["schema"]
assert schema["additionalProperties"] is False
assert schema["required"] == ["user"]
assert schema["properties"]["user"]["additionalProperties"] is False
assert schema["properties"]["user"]["required"] == ["name", "address"]
assert schema["properties"]["user"]["properties"]["address"]["additionalProperties"] is False
assert schema["properties"]["user"]["properties"]["address"]["required"] == ["city"]
def test_array_items_object_adds_additional_properties_false(self):
output_format = {
"type": "json_schema",
"schema": {
"type": "object",
"properties": {
"items": {
"type": "array",
"items": {
"type": "object",
"properties": {"id": {"type": "integer"}},
},
}
},
},
}
result = self.adapter.translate_anthropic_output_format_to_openai(output_format)
assert result is not None
schema = result["json_schema"]["schema"]
assert schema["additionalProperties"] is False
assert schema["properties"]["items"]["items"]["additionalProperties"] is False
def test_does_not_mutate_original_schema(self):
original_schema = {
"type": "object",
"properties": {"name": {"type": "string"}},
}
output_format = {"type": "json_schema", "schema": original_schema}
self.adapter.translate_anthropic_output_format_to_openai(output_format)
assert "additionalProperties" not in original_schema
assert "required" not in original_schema
def test_defs_adds_additional_properties_false(self):
output_format = {
"type": "json_schema",
"schema": {
"type": "object",
"properties": {"ref": {"$ref": "#/$defs/Item"}},
"$defs": {
"Item": {
"type": "object",
"properties": {"value": {"type": "string"}},
}
},
},
}
result = self.adapter.translate_anthropic_output_format_to_openai(output_format)
assert result is not None
schema = result["json_schema"]["schema"]
assert schema["$defs"]["Item"]["additionalProperties"] is False
assert schema["$defs"]["Item"]["required"] == ["value"]
def test_incomplete_required_gets_completed(self):
"""OpenAI strict mode requires ALL properties in required."""
output_format = {
"type": "json_schema",
"schema": {
"type": "object",
"properties": {
"name": {"type": "string"},
"age": {"type": "integer"},
"email": {"type": "string"},
},
"required": ["name"], # only 1 of 3
},
}
result = self.adapter.translate_anthropic_output_format_to_openai(output_format)
assert result is not None
schema = result["json_schema"]["schema"]
assert schema["additionalProperties"] is False
assert sorted(schema["required"]) == ["age", "email", "name"]
def test_invalid_output_format_returns_none(self):
assert self.adapter.translate_anthropic_output_format_to_openai("invalid") is None
assert self.adapter.translate_anthropic_output_format_to_openai({"type": "text"}) is None
assert self.adapter.translate_anthropic_output_format_to_openai({"type": "json_schema"}) is None

View file

@ -336,7 +336,7 @@ class TestAnthropicFilesHandler:
"extra_body": None
}
with patch.object(handler.anthropic_model_info, "get_api_key", return_value=None):
with patch.object(handler.anthropic_model_info, "get_auth_header", return_value=None):
with pytest.raises(ValueError, match="Missing Anthropic API Key"):
await handler.afile_content(
file_content_request=file_content_request,

View file

@ -3803,3 +3803,100 @@ def test_streaming_without_json_mode_passes_all_tools():
assert tool_use_delta is not None
assert tool_use_delta["function"]["arguments"] == '{"data": 1}'
def test_cache_control_injection_tool_config():
"""Test that cache_control_injection_points with location=tool_config appends cachePoint to tools."""
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string"},
},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config"},
],
}
result = config._transform_request(
model="anthropic.claude-3-5-haiku-20241022-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tool_config = result["toolConfig"]
tools = tool_config["tools"]
# Last element should be a cachePoint block
assert tools[-1] == {"cachePoint": {"type": "default"}}
# First element should be the actual tool
assert "toolSpec" in tools[0]
def test_cache_control_injection_tool_config_no_tools():
"""Test that tool_config injection is ignored when no tools are provided."""
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "Hello"},
]
optional_params = {
"cache_control_injection_points": [
{"location": "tool_config"},
],
}
result = config._transform_request(
model="anthropic.claude-3-5-haiku-20241022-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
assert "toolConfig" not in result
def test_cache_control_injection_tool_config_not_added_without_injection_point():
"""Test that cachePoint is NOT appended when cache_control_injection_points doesn't include tool_config."""
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "message", "role": "system"},
],
}
result = config._transform_request(
model="anthropic.claude-3-5-haiku-20241022-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
# No cachePoint should be appended
assert all("cachePoint" not in tool for tool in tools)

View file

@ -550,4 +550,72 @@ class TestMoonshotConfig:
# reasoning_content must not have been injected
for msg in result["messages"]:
assert "reasoning_content" not in msg
assert "reasoning_content" not in msg
def test_reasoning_content_preserved_on_pydantic_message_object(self):
"""reasoning_content on Pydantic Message objects is preserved (not overwritten with placeholder).
Regression test for: https://github.com/BerriAI/litellm/issues/23765
The issue was that 'reasoning_content' in msg doesn't work for Pydantic models
because they don't support the 'in' operator the same way as dicts.
"""
from litellm.types.utils import Message
config = MoonshotChatConfig()
# Create a Pydantic Message object with reasoning_content (as would come from API response)
message_with_reasoning = Message(
role="assistant",
content=None,
reasoning_content="<thinking>User wants weather</thinking>",
tool_calls=[
{"id": "call_1", "type": "function", "function": {"name": "fn", "arguments": "{}"}}
],
)
messages = [message_with_reasoning]
result = config.fill_reasoning_content(messages)
# reasoning_content should be preserved, not replaced with placeholder
assert result[0].get("reasoning_content") == "<thinking>User wants weather</thinking>"
def test_reasoning_content_preserved_in_multi_turn_flow(self):
"""reasoning_content is preserved through multi-turn conversation flow.
This tests the complete flow: API response -> Message object -> dict -> fill_reasoning_content
"""
from litellm.types.utils import Message
from litellm.utils import convert_to_dict
config = MoonshotChatConfig()
# Simulate API response with reasoning_content
api_response = {
"role": "assistant",
"content": None,
"reasoning_content": "<thinking>Planning to call weather tool</thinking>",
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{}'}}
],
}
# Convert to Message object (as LiteLLM does)
message_obj = Message(**api_response)
# Convert back to dict (when building next request)
message_dict = convert_to_dict(message_obj)
# Build multi-turn conversation
messages = [
{"role": "user", "content": "What's the weather?"},
message_dict,
{"role": "tool", "tool_call_id": "call_1", "content": '{"temp": 72}'},
{"role": "user", "content": "Thanks!"},
]
# Apply fill_reasoning_content
result = config.fill_reasoning_content(messages)
# reasoning_content should be preserved in the assistant message
assert result[1].get("reasoning_content") == "<thinking>Planning to call weather tool</thinking>"

View file

@ -281,6 +281,45 @@ class TestOpenAIChatCompletionStreamingHandler:
# Verify that reasoning_content is not set (it should be deleted by Delta.__init__)
assert not hasattr(parsed_chunk.choices[0].delta, "reasoning_content")
def test_chunk_parser_without_id_field(self):
"""
Test that chunk_parser works when chunk is missing the 'id' field.
Some OpenAI-compatible providers (e.g., MiniMax) return streaming chunks
without an 'id' field in certain cases. This should not raise KeyError.
Regression test for: KeyError: 'id' when using MiniMax m2.5 model
"""
handler = OpenAIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Simulate a chunk without 'id' field (as returned by MiniMax)
chunk = {
"object": "chat.completion.chunk",
"created": 1769511767,
"model": "minimax/m2.5",
"choices": [
{
"delta": {
"content": "Hello",
"role": "assistant",
},
"finish_reason": None,
"index": 0,
}
],
}
# Parse the chunk - should not raise KeyError
parsed_chunk = handler.chunk_parser(chunk)
# Verify that content is present and id was auto-generated
assert parsed_chunk.choices[0].delta.content == "Hello"
assert parsed_chunk.choices[0].delta.role == "assistant"
# ModelResponseStream auto-generates an id when None is passed
assert parsed_chunk.id is not None
class TestPromptCacheKeyIntegration:
"""Tests for prompt_cache_key support"""

View file

@ -0,0 +1,234 @@
"""
Tests for Gemini context circulation (server-side tool invocations).
When includeServerSideToolInvocations=true is set, Gemini returns toolCall/toolResponse
parts for server-side tools (e.g. Google Search). These must be:
1. Extracted from the response into provider_specific_fields["server_side_tool_invocations"]
2. Re-injected as raw toolCall/toolResponse parts when converting messages back to Gemini format
3. The includeServerSideToolInvocations flag must be passed through to toolConfig
"""
import json
from typing import Any, Dict, List
from unittest.mock import MagicMock
import pytest
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
from litellm.llms.vertex_ai.gemini.transformation import (
_gemini_convert_messages_with_history,
)
from litellm.types.llms.vertex_ai import HttpxPartType
# --- Response extraction tests ---
class TestExtractServerSideToolInvocations:
"""Test _extract_server_side_tool_invocations from response parts."""
def test_extracts_tool_call_and_response(self):
"""Basic case: one toolCall + one toolResponse with same id."""
parts: List[HttpxPartType] = [
{
"thoughtSignature": "sig_call_1",
"toolCall": {
"toolType": "GOOGLE_SEARCH_WEB",
"id": "abc123",
"args": {"queries": ["weather Buenos Aires"]},
},
},
{
"thoughtSignature": "sig_resp_1",
"toolResponse": {
"toolType": "GOOGLE_SEARCH_WEB",
"id": "abc123",
"response": {"weather": "Sunny, 20°C"},
},
},
{
"text": "The weather in Buenos Aires is sunny.",
"thoughtSignature": "sig_text",
},
]
result = VertexGeminiConfig._extract_server_side_tool_invocations(parts)
assert result is not None
assert len(result) == 1
assert result[0]["tool_type"] == "GOOGLE_SEARCH_WEB"
assert result[0]["id"] == "abc123"
assert result[0]["args"] == {"queries": ["weather Buenos Aires"]}
assert result[0]["response"] == {"weather": "Sunny, 20°C"}
assert result[0]["thought_signature"] == "sig_call_1"
def test_returns_none_when_no_server_side_tools(self):
"""No toolCall/toolResponse parts → returns None."""
parts: List[HttpxPartType] = [
{"text": "Hello world", "thoughtSignature": "sig1"},
{
"functionCall": {
"name": "get_weather",
"args": {"location": "Paris"},
},
"thoughtSignature": "sig2",
},
]
result = VertexGeminiConfig._extract_server_side_tool_invocations(parts)
assert result is None
def test_multiple_server_side_invocations(self):
"""Multiple toolCall/toolResponse pairs."""
parts: List[HttpxPartType] = [
{
"toolCall": {
"toolType": "GOOGLE_SEARCH_WEB",
"id": "search1",
"args": {"queries": ["query1"]},
},
"thoughtSignature": "sig1",
},
{
"toolResponse": {"toolType": "GOOGLE_SEARCH_WEB", "id": "search1", "response": "result1"},
"thoughtSignature": "sig2",
},
{
"toolCall": {
"toolType": "GOOGLE_SEARCH_WEB",
"id": "search2",
"args": {"queries": ["query2"]},
},
"thoughtSignature": "sig3",
},
{
"toolResponse": {"toolType": "GOOGLE_SEARCH_WEB", "id": "search2", "response": "result2"},
"thoughtSignature": "sig4",
},
]
result = VertexGeminiConfig._extract_server_side_tool_invocations(parts)
assert result is not None
assert len(result) == 2
assert result[0]["id"] == "search1"
assert result[0]["response"] == "result1"
assert result[1]["id"] == "search2"
assert result[1]["response"] == "result2"
def test_tool_call_without_response(self):
"""toolCall without matching toolResponse is still captured."""
parts: List[HttpxPartType] = [
{
"toolCall": {
"toolType": "CODE_EXECUTION",
"id": "exec1",
"args": {"code": "print('hello')"},
},
},
]
result = VertexGeminiConfig._extract_server_side_tool_invocations(parts)
assert result is not None
assert len(result) == 1
assert result[0]["id"] == "exec1"
assert "response" not in result[0]
# --- Input re-injection tests ---
class TestReInjectServerSideToolInvocations:
"""Test that server_side_tool_invocations are re-injected into Gemini parts."""
def test_roundtrip_single_invocation(self):
"""Server-side invocations from assistant message are converted back to Gemini parts."""
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": "It's sunny in Buenos Aires.",
"provider_specific_fields": {
"server_side_tool_invocations": [
{
"tool_type": "GOOGLE_SEARCH_WEB",
"id": "abc123",
"args": {"queries": ["weather Buenos Aires"]},
"response": {"weather": "Sunny, 20°C"},
"thought_signature": "sig_abc",
}
]
},
},
{"role": "user", "content": "Thanks!"},
]
contents = _gemini_convert_messages_with_history(messages)
# Find the model turn
model_turn = [c for c in contents if c["role"] == "model"]
assert len(model_turn) == 1
parts = model_turn[0]["parts"]
# Should have: text part + toolCall part + toolResponse part
tool_call_parts = [p for p in parts if "toolCall" in p]
tool_response_parts = [p for p in parts if "toolResponse" in p]
assert len(tool_call_parts) == 1
assert tool_call_parts[0]["toolCall"]["toolType"] == "GOOGLE_SEARCH_WEB"
assert tool_call_parts[0]["toolCall"]["id"] == "abc123"
assert tool_call_parts[0]["toolCall"]["args"] == {"queries": ["weather Buenos Aires"]}
assert tool_call_parts[0]["thoughtSignature"] == "sig_abc"
assert len(tool_response_parts) == 1
assert tool_response_parts[0]["toolResponse"]["id"] == "abc123"
assert tool_response_parts[0]["toolResponse"]["toolType"] == "GOOGLE_SEARCH_WEB"
assert tool_response_parts[0]["toolResponse"]["response"] == {"weather": "Sunny, 20°C"}
def test_no_invocations_no_extra_parts(self):
"""Without server_side_tool_invocations, no extra parts are added."""
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there!"},
{"role": "user", "content": "Bye"},
]
contents = _gemini_convert_messages_with_history(messages)
model_turn = [c for c in contents if c["role"] == "model"]
assert len(model_turn) == 1
parts = model_turn[0]["parts"]
assert len(parts) == 1
assert "text" in parts[0]
assert "toolCall" not in parts[0]
# --- toolConfig flag tests ---
class TestIncludeServerSideToolInvocationsConfig:
"""Test that the flag is passed through to toolConfig."""
def test_flag_added_to_tool_config(self):
"""include_server_side_tool_invocations=True should be mapped to optional_params."""
config = VertexGeminiConfig()
non_default_params = {"include_server_side_tool_invocations": True}
optional_params: Dict[str, Any] = {}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-3-flash-preview",
drop_params=False,
)
assert result["include_server_side_tool_invocations"] is True
def test_flag_in_supported_params(self):
"""include_server_side_tool_invocations should be in supported params."""
config = VertexGeminiConfig()
supported = config.get_supported_openai_params(model="gemini-3-flash-preview")
assert "include_server_side_tool_invocations" in supported

View file

@ -2,7 +2,8 @@
Unit tests for auth_utils functions related to rate limiting and customer ID extraction.
"""
from unittest.mock import patch
from typing import Optional
from unittest.mock import MagicMock, patch
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
@ -70,6 +71,19 @@ class TestGetKeyModelRpmLimit:
assert result is None
def test_team_metadata_empty_rpm_dict_falls_through_to_deployment_default(self):
"""Explicitly empty team model_rpm_limit ({}) should be returned as-is, not fallen through."""
# An empty dict is a valid team limit map (no per-model limits configured).
# It should be returned directly rather than falling through to deployment defaults,
# so a team with an empty map is treated as unconstrained at the team level.
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-123",
team_metadata={"model_rpm_limit": {}},
)
result = get_key_model_rpm_limit(user_api_key_dict)
assert result == {}
class TestGetKeyModelTpmLimit:
"""Tests for get_key_model_tpm_limit function."""
@ -136,6 +150,33 @@ class TestGetKeyModelTpmLimit:
assert result == {"gpt-4": 10000}
def test_team_metadata_empty_tpm_dict_falls_through_to_deployment_default(self):
"""Explicitly empty team model_tpm_limit ({}) should be returned as-is, not fallen through."""
# An empty dict is a valid team limit map (no per-model limits configured).
# It should be returned directly rather than falling through to deployment defaults,
# so a team with an empty map is treated as unconstrained at the team level.
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-123",
team_metadata={"model_tpm_limit": {}},
)
result = get_key_model_tpm_limit(user_api_key_dict)
assert result == {}
def test_skips_deployments_with_malformed_limit_value(self):
"""Deployments with non-integer-parseable limit values are skipped without raising."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
{"model_name": "model1", "litellm_params": {"default_api_key_tpm_limit": "not-a-number"}},
_make_deployment_dict("model1", tpm=500),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
# The malformed deployment is skipped; the valid one provides 500
assert result == {"model1": 500}
class TestGetCustomerIdFromStandardHeaders:
"""Tests for _get_customer_id_from_standard_headers helper function."""
@ -315,3 +356,196 @@ def test_get_end_user_id_falls_back_to_deprecated_user_header_name():
result = get_end_user_id_from_request_body(request_body={}, request_headers=headers)
assert result == "user-legacy"
def _make_deployment_dict(model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None) -> dict:
"""Helper to build a minimal deployment dict as returned by router.get_model_list."""
litellm_params: dict = {"model": model_name}
if tpm is not None:
litellm_params["default_api_key_tpm_limit"] = tpm
if rpm is not None:
litellm_params["default_api_key_rpm_limit"] = rpm
return {"model_name": model_name, "litellm_params": litellm_params}
_ROUTER_PATCH = "litellm.proxy.proxy_server.llm_router"
class TestDeploymentDefaultRpmLimit:
"""Tests for deployment default_api_key_rpm_limit fallback in get_key_model_rpm_limit."""
def test_returns_deployment_default_when_key_has_no_limits(self):
"""Case 2 from spec: key has no model-specific limits, falls back to deployment default."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", rpm=200)
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 200}
def test_key_model_limit_takes_priority_over_deployment_default(self):
"""Case 1 from spec: key model-specific limit wins over deployment default."""
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-123",
metadata={"model_rpm_limit": {"model1": 10}},
)
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", rpm=200)
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 10}
def test_returns_none_when_no_deployment_default_and_no_key_limits(self):
"""Returns None when neither the key nor the deployment has any rpm limit."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1") # no rpm default
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
assert result is None
def test_returns_none_without_model_name_even_when_deployment_has_default(self):
"""No model_name means deployment fallback is skipped."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", rpm=200)
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict)
assert result is None
def test_returns_none_when_llm_router_is_none(self):
"""No router means deployment fallback returns None gracefully."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
with patch(_ROUTER_PATCH, None):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
assert result is None
def test_returns_minimum_across_multiple_deployments(self):
"""When multiple deployments share a model name, the minimum rpm limit is used."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", rpm=200),
_make_deployment_dict("model1", rpm=50),
_make_deployment_dict("model1", rpm=150),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 50}
def test_ignores_deployments_without_default_when_others_have_it(self):
"""Deployments missing the field are skipped; min is taken over those that have it."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1"), # no rpm default
_make_deployment_dict("model1", rpm=75),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 75}
def test_skips_deployments_with_malformed_limit_value(self):
"""Deployments with non-integer-parseable limit values are skipped without raising."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
{"model_name": "model1", "litellm_params": {"default_api_key_rpm_limit": "not-a-number"}},
_make_deployment_dict("model1", rpm=100),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
# The malformed deployment is skipped; the valid one provides 100
assert result == {"model1": 100}
class TestDeploymentDefaultTpmLimit:
"""Tests for deployment default_api_key_tpm_limit fallback in get_key_model_tpm_limit."""
def test_returns_deployment_default_when_key_has_no_limits(self):
"""Case 2 from spec: key has no model-specific limits, falls back to deployment default."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", tpm=100)
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 100}
def test_key_model_limit_takes_priority_over_deployment_default(self):
"""Case 1 from spec: key model-specific limit wins over deployment default."""
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-123",
metadata={"model_tpm_limit": {"model1": 20}},
)
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", tpm=100)
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 20}
def test_returns_none_when_no_deployment_default_and_no_key_limits(self):
"""Returns None when neither the key nor the deployment has any tpm limit."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1") # no tpm default
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
assert result is None
def test_returns_none_without_model_name_even_when_deployment_has_default(self):
"""No model_name means deployment fallback is skipped."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", tpm=100)
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict)
assert result is None
def test_returns_none_when_llm_router_is_none(self):
"""No router means deployment fallback returns None gracefully."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
with patch(_ROUTER_PATCH, None):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
assert result is None
def test_returns_minimum_across_multiple_deployments(self):
"""When multiple deployments share a model name, the minimum tpm limit is used."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1", tpm=1000),
_make_deployment_dict("model1", tpm=300),
_make_deployment_dict("model1", tpm=700),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 300}
def test_ignores_deployments_without_default_when_others_have_it(self):
"""Deployments missing the field are skipped; min is taken over those that have it."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
_make_deployment_dict("model1"), # no tpm default
_make_deployment_dict("model1", tpm=400),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
assert result == {"model1": 400}

View file

@ -1,7 +1,8 @@
import json
import os
import signal
import sys
from unittest.mock import AsyncMock, Mock, patch
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import pytest
from fastapi.testclient import TestClient
@ -14,6 +15,14 @@ sys.path.insert(
from litellm.proxy.db.prisma_client import PrismaWrapper, should_update_prisma_schema
@pytest.fixture(autouse=True)
def mock_prisma_binary():
"""Mock prisma.Prisma to avoid requiring generated Prisma binaries for unit tests."""
mock_module = MagicMock()
with patch.dict(sys.modules, {"prisma": mock_module}):
yield mock_module
def test_should_update_prisma_schema(monkeypatch):
# CASE 1: Environment variable behavior
# When DISABLE_SCHEMA_UPDATE is not set -> should update
@ -73,4 +82,79 @@ async def test_recreate_prisma_client_successful_disconnect():
# Verify that the new client replaced the original
assert wrapper._original_prisma != mock_prisma
assert hasattr(wrapper._original_prisma, 'connect')
assert hasattr(wrapper._original_prisma, 'connect')
@pytest.mark.asyncio
async def test_recreate_prisma_client_kills_old_engine_on_disconnect_failure(
mock_prisma_binary,
):
"""When disconnect() fails, recreate_prisma_client must SIGTERM/SIGKILL the old engine PID."""
mock_prisma = AsyncMock()
mock_prisma.disconnect.side_effect = Exception("engine hung")
# Simulate engine subprocess with a known PID
mock_engine = MagicMock()
mock_engine.process.pid = 12345
mock_prisma._engine = mock_engine
wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False)
# Configure the mock Prisma constructor
mock_new_prisma = AsyncMock()
mock_prisma_binary.Prisma.return_value = mock_new_prisma
with (
patch("os.kill") as mock_kill,
patch("asyncio.sleep", new_callable=AsyncMock),
):
await wrapper.recreate_prisma_client("postgresql://new")
# Verify old engine was killed
mock_kill.assert_any_call(12345, signal.SIGTERM)
# Verify new client was created and connected
mock_new_prisma.connect.assert_awaited_once()
@pytest.mark.asyncio
async def test_recreate_prisma_client_skips_kill_on_successful_disconnect(
mock_prisma_binary,
):
"""When disconnect() succeeds, no kill should be attempted."""
mock_prisma = AsyncMock()
mock_prisma.disconnect.return_value = None
wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False)
mock_new_prisma = AsyncMock()
mock_prisma_binary.Prisma.return_value = mock_new_prisma
with patch("os.kill") as mock_kill:
await wrapper.recreate_prisma_client("postgresql://new")
mock_kill.assert_not_called()
mock_new_prisma.connect.assert_awaited_once()
@pytest.mark.asyncio
async def test_recreate_prisma_client_handles_missing_engine_pid(
mock_prisma_binary,
):
"""When engine PID is unavailable (no _engine attr), kill is skipped gracefully."""
mock_prisma = AsyncMock()
mock_prisma.disconnect.side_effect = Exception("engine hung")
mock_prisma._engine = None # No engine subprocess
wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False)
mock_new_prisma = AsyncMock()
mock_prisma_binary.Prisma.return_value = mock_new_prisma
with (
patch("os.kill") as mock_kill,
patch("asyncio.sleep", new_callable=AsyncMock),
):
await wrapper.recreate_prisma_client("postgresql://new")
mock_kill.assert_not_called() # PID was 0, kill skipped
mock_new_prisma.connect.assert_awaited_once()

View file

@ -1,5 +1,6 @@
import asyncio
import os
import signal
import sys
import time
from unittest.mock import AsyncMock, MagicMock, patch
@ -279,3 +280,37 @@ async def test_db_health_watchdog_start_stop_lifecycle(mock_proxy_logging):
await client.stop_db_health_watchdog_task()
assert client._db_health_watchdog_task is None
assert dummy_task.cancelled() is True
@pytest.mark.asyncio
async def test_lightweight_reconnect_kills_engine_on_disconnect_failure(mock_proxy_logging):
"""Lightweight reconnect must kill the old engine PID when disconnect() fails."""
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
client.db.disconnect = AsyncMock(side_effect=Exception("disconnect failed"))
client.db.connect = AsyncMock(return_value=None)
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
with (
patch.object(client, "_get_engine_pid", return_value=9999),
patch("os.kill") as mock_kill,
patch("asyncio.sleep", new_callable=AsyncMock),
):
await client._run_reconnect_cycle(timeout_seconds=5.0)
mock_kill.assert_any_call(9999, signal.SIGTERM)
client.db.connect.assert_awaited_once()
client.db.query_raw.assert_awaited_once_with("SELECT 1")
@pytest.mark.asyncio
async def test_lightweight_reconnect_skips_kill_on_successful_disconnect(mock_proxy_logging):
"""Lightweight reconnect must NOT kill when disconnect() succeeds."""
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
client.db.disconnect = AsyncMock(return_value=None)
client.db.connect = AsyncMock(return_value=None)
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
with patch("os.kill") as mock_kill:
await client._run_reconnect_cycle(timeout_seconds=5.0)
mock_kill.assert_not_called()

View file

@ -1,6 +1,6 @@
import os
import sys
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import pytest
from fastapi import FastAPI
@ -11,6 +11,7 @@ sys.path.insert(
)
from litellm.proxy.discovery_endpoints.ui_discovery_endpoints import router
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
def test_ui_discovery_endpoints_with_defaults():
@ -245,9 +246,9 @@ def test_ui_discovery_endpoints_with_admin_ui_enabled():
patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), \
patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False), \
patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False):
response = client.get("/.well-known/litellm-ui-config")
assert response.status_code == 200
data = response.json()
assert data["server_root_path"] == "/"
@ -256,3 +257,53 @@ def test_ui_discovery_endpoints_with_admin_ui_enabled():
assert data["admin_ui_disabled"] is False
assert data["sso_configured"] is False
def test_ui_discovery_endpoints_is_control_plane_true_when_workers_configured():
app = FastAPI()
app.include_router(router)
client = TestClient(app)
mock_config = MagicMock()
mock_config.worker_registry = [
WorkerRegistryEntry(
worker_id="team-a", name="Team A", url="https://worker-1:4001"
),
]
with patch("litellm.proxy.utils.get_server_root_path", return_value="/"), \
patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), \
patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False), \
patch("litellm.proxy.proxy_server.proxy_config", mock_config), \
patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False):
response = client.get("/.well-known/litellm-ui-config")
assert response.status_code == 200
data = response.json()
assert data["is_control_plane"] is True
assert len(data["workers"]) == 1
assert data["workers"][0]["worker_id"] == "team-a"
assert data["workers"][0]["name"] == "Team A"
assert data["workers"][0]["url"] == "https://worker-1:4001"
def test_ui_discovery_endpoints_is_control_plane_false_when_no_workers():
app = FastAPI()
app.include_router(router)
client = TestClient(app)
mock_config = MagicMock()
mock_config.worker_registry = []
with patch("litellm.proxy.utils.get_server_root_path", return_value="/"), \
patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), \
patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False), \
patch("litellm.proxy.proxy_server.proxy_config", mock_config), \
patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False):
response = client.get("/.well-known/litellm-ui-config")
assert response.status_code == 200
data = response.json()
assert data["is_control_plane"] is False
assert data["workers"] == []

View file

@ -0,0 +1,944 @@
"""
Tests for deferred logging with post-call guardrails.
When post-call guardrails are configured, the async logging task is deferred
until after guardrails complete. This ensures the StandardLoggingPayload
is built with guardrail_information populated.
Non-streaming: create_task in wrapper_async is replaced by a closure that
the proxy fires in a try/finally after post_call_success_hook.
Streaming: a closure on logging_obj is called by CSW.__anext__ at stream end.
The closure runs ONLY guardrail hooks (not all callbacks), then fires
both logging handlers.
"""
import asyncio
import os
import sys
from typing import Any
from unittest.mock import MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class PostCallGuardrail(CustomGuardrail):
"""A post-call guardrail."""
def __init__(self):
super().__init__(
guardrail_name="post-call",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
return response
class PreCallGuardrail(CustomGuardrail):
"""A pre-call-only guardrail — should NOT trigger deferral."""
def __init__(self):
super().__init__(
guardrail_name="pre-call",
default_on=True,
event_hook=GuardrailEventHooks.pre_call,
)
class AllEventsGuardrail(CustomGuardrail):
"""A guardrail with event_hook=None (runs on all events)."""
def __init__(self):
super().__init__(
guardrail_name="all-events",
default_on=True,
event_hook=None,
)
# ---------------------------------------------------------------------------
# 1. _has_post_call_guardrails detection
# ---------------------------------------------------------------------------
class TestHasPostCallGuardrails:
def test_returns_true_for_post_call_guardrail(self):
with patch("litellm.callbacks", [PostCallGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True
def test_returns_true_for_event_hook_none(self):
"""event_hook=None means 'all events', including post_call."""
with patch("litellm.callbacks", [AllEventsGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True
def test_returns_false_for_pre_call_only(self):
with patch("litellm.callbacks", [PreCallGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False
def test_returns_false_for_no_callbacks(self):
with patch("litellm.callbacks", []):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False
def test_ignores_non_guardrail_callbacks(self):
"""String callbacks and CustomLogger instances are not guardrails."""
with patch("litellm.callbacks", ["langfuse", CustomLogger()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False
def test_returns_true_for_list_with_post_call(self):
"""event_hook as a list containing post_call should trigger deferral."""
class ListGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="list-post",
default_on=True,
event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call],
)
with patch("litellm.callbacks", [ListGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True
def test_returns_false_for_list_without_post_call(self):
"""event_hook as a list without post_call should not trigger deferral."""
class ListGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="list-pre",
default_on=True,
event_hook=[GuardrailEventHooks.pre_call],
)
with patch("litellm.callbacks", [ListGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False
# ---------------------------------------------------------------------------
# 2. Non-streaming: deferral flag → closure stored, create_task skipped
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_deferred_flag_stores_and_executes_closure():
"""
When _defer_async_logging is True on logging_obj:
1. wrapper_async stores a callable closure instead of calling create_task
2. Calling the closure fires create_task
3. Sync callbacks fire immediately (not deferred)
"""
mock_logging_obj = MagicMock()
mock_logging_obj._defer_async_logging = True
mock_logging_obj._enqueue_deferred_logging = None
await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
litellm_logging_obj=mock_logging_obj,
)
# Closure was stored
enqueue_fn = mock_logging_obj._enqueue_deferred_logging
assert callable(enqueue_fn), "Closure should be stored on logging_obj"
# Sync callbacks fired immediately
mock_logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
# Calling the closure fires create_task
created_tasks = []
real_create_task = asyncio.create_task
def tracking_create_task(coro):
task = real_create_task(coro)
created_tasks.append(task)
return task
with patch("asyncio.create_task", side_effect=tracking_create_task):
enqueue_fn()
assert len(created_tasks) >= 1, "Closure should fire asyncio.create_task"
for task in created_tasks:
if not task.done():
task.cancel()
try:
await task
except (asyncio.CancelledError, Exception):
pass
# ---------------------------------------------------------------------------
# 3. Non-streaming regression: without flag, create_task fires normally
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_no_flag_fires_create_task_normally():
"""Without _defer_async_logging, wrapper_async calls create_task as before."""
created_tasks = []
real_create_task = asyncio.create_task
def tracking_create_task(coro):
task = real_create_task(coro)
created_tasks.append(task)
return task
with patch("asyncio.create_task", side_effect=tracking_create_task):
await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
)
assert len(created_tasks) >= 1
for task in created_tasks:
if not task.done():
task.cancel()
try:
await task
except (asyncio.CancelledError, Exception):
pass
# ---------------------------------------------------------------------------
# 4. Non-streaming: deferred logging fires even if guardrail raises
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_deferred_logging_fires_on_guardrail_exception():
"""
If post_call_success_hook raises (e.g., guardrail blocks content),
the deferred logging closure must still fire (via try/finally).
"""
from fastapi import HTTPException # noqa: local import for test isolation
enqueue_called = False
def mock_enqueue():
nonlocal enqueue_called
enqueue_called = True
class BlockingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="blocker",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
raise HTTPException(status_code=400, detail="Content blocked")
guardrail = BlockingGuardrail()
logging_obj = MagicMock()
logging_obj._enqueue_deferred_logging = mock_enqueue
with patch("litellm.callbacks", [guardrail]):
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
with pytest.raises(HTTPException):
try:
await proxy_logging.post_call_success_hook(
data={"model": "gpt-4", "metadata": {}},
response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(api_key="test"),
)
finally:
# Mirrors the proxy's finally block
_enqueue_fn = getattr(logging_obj, "_enqueue_deferred_logging", None)
if _enqueue_fn is not None:
logging_obj._enqueue_deferred_logging = None
_enqueue_fn()
assert enqueue_called is True
assert logging_obj._enqueue_deferred_logging is None
# ---------------------------------------------------------------------------
# 5. Streaming: closure defers logging at stream end
# ---------------------------------------------------------------------------
class TestDeferredStreamingClosure:
@pytest.mark.asyncio
async def test_streaming_closure_defers_logging(self):
"""When _on_deferred_stream_complete is set, CSW calls the closure
instead of firing async_success_handler directly."""
mock_logging_obj = MagicMock()
callback_called = False
callback_args = {}
async def mock_callback(assembled_response, cache_hit):
nonlocal callback_called, callback_args
callback_called = True
callback_args = {"response": assembled_response, "cache_hit": cache_hit}
mock_logging_obj._on_deferred_stream_complete = mock_callback
resp = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
stream=True,
litellm_logging_obj=mock_logging_obj,
)
async for _ in resp:
pass
await asyncio.sleep(0)
assert callback_called is True, "Closure should be called at stream end"
assert callback_args["response"] is not None
assert mock_logging_obj._on_deferred_stream_complete is None
@pytest.mark.asyncio
async def test_streaming_no_closure_fires_normally(self):
"""Regression: without closure, CSW fires logging immediately."""
created_tasks = []
real_create_task = asyncio.create_task
def tracking_create_task(coro):
task = real_create_task(coro)
created_tasks.append(task)
return task
resp = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
stream=True,
)
with patch("asyncio.create_task", side_effect=tracking_create_task):
async for _ in resp:
pass
assert len(created_tasks) >= 1
for task in created_tasks:
if not task.done():
task.cancel()
try:
await task
except (asyncio.CancelledError, Exception):
pass
@pytest.mark.asyncio
async def test_closure_runs_only_guardrail_hooks(self):
"""The closure must call only CustomGuardrail hooks, not all callbacks.
This is the key v2 change — PR #23929 called post_call_success_hook
which ran ALL callbacks, causing behavioral changes for streaming."""
guardrail_called = False
logger_called = False
class TrackingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="tracker",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
nonlocal guardrail_called
guardrail_called = True
return response
class TrackingLogger(CustomLogger):
async def async_post_call_success_hook(
self, user_api_key_dict, data, response
):
nonlocal logger_called
logger_called = True
return response
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
pass
mock_logging_obj.async_success_handler = track_async_success
tracking_guardrail = TrackingGuardrail()
tracking_logger = TrackingLogger()
# Use the real production static method via a thin closure
_captured_data = {"model": "gpt-4", "metadata": {}}
_captured_user_api_key_dict = UserAPIKeyAuth(api_key="test")
async def _on_deferred_stream_complete(assembled_response, cache_hit):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=_captured_data,
captured_user_api_key_dict=_captured_user_api_key_dict,
captured_logging_obj=mock_logging_obj,
assembled_response=assembled_response,
cache_hit=cache_hit,
)
mock_logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
with patch("litellm.callbacks", [tracking_guardrail, tracking_logger]):
resp = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
stream=True,
litellm_logging_obj=mock_logging_obj,
)
async for _ in resp:
pass
await asyncio.sleep(0)
await asyncio.sleep(0)
assert guardrail_called is True, "Guardrail hook should be called"
assert logger_called is False, "Non-guardrail logger should NOT be called by closure"
@pytest.mark.asyncio
async def test_closure_passes_guardrail_modified_response_to_logging(self):
"""The production _run_deferred_stream_guardrails must pass the
guardrail-modified response to async_success_handler."""
logged_response = None
modified_response = MagicMock()
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
nonlocal logged_response
logged_response = args[0] if args else None
mock_logging_obj.async_success_handler = track_async_success
class ModifyingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="modifier",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
return modified_response
guardrail = ModifyingGuardrail()
with patch("litellm.callbacks", [guardrail]):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={"model": "gpt-4", "metadata": {}},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert logged_response is modified_response, \
"Logging must receive the guardrail-modified response"
@pytest.mark.asyncio
async def test_closure_logs_even_on_guardrail_exception(self):
"""If a guardrail raises HTTPException, the production
_run_deferred_stream_guardrails must still fire logging
and set guardrail_blocked in metadata."""
from fastapi import HTTPException # noqa: local import for test isolation
logging_called = False
class BlockingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="blocker",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
raise HTTPException(status_code=400, detail="Blocked")
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
nonlocal logging_called
logging_called = True
mock_logging_obj.async_success_handler = track_async_success
guardrail = BlockingGuardrail()
with patch("litellm.callbacks", [guardrail]):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={"model": "gpt-4", "metadata": {}},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert logging_called is True, \
"Logging must fire even when guardrail raises HTTPException"
assert mock_logging_obj.model_call_details["metadata"].get(
"guardrail_blocked"
) is True, "guardrail_blocked must be set for HTTPException"
@pytest.mark.asyncio
async def test_transient_error_does_not_set_guardrail_blocked(self):
"""Transient errors (not HTTPException) should NOT set
guardrail_blocked. Uses the production _run_deferred_stream_guardrails."""
class TransientErrorGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="transient",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
raise ConnectionError("Network timeout")
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
pass
mock_logging_obj.async_success_handler = track_async_success
guardrail = TransientErrorGuardrail()
with patch("litellm.callbacks", [guardrail]):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={"model": "gpt-4", "metadata": {}},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
await asyncio.sleep(0)
assert mock_logging_obj.model_call_details["metadata"].get(
"guardrail_blocked"
) is not True, "guardrail_blocked must NOT be set for transient errors"
@pytest.mark.asyncio
async def test_production_closure_integration(self):
"""Integration test: calls the real _run_deferred_stream_guardrails
static method and verifies it calls guardrail hooks and passes
the modified response to logging."""
hook_called = False
logged_response = None
modified_response = MagicMock()
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
nonlocal logged_response
logged_response = args[0] if args else None
mock_logging_obj.async_success_handler = track_async_success
class TestGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="test",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
nonlocal hook_called
hook_called = True
return modified_response
guardrail = TestGuardrail()
async def _on_deferred_stream_complete(assembled_response, cache_hit):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={"model": "gpt-4", "metadata": {}},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=assembled_response,
cache_hit=cache_hit,
)
mock_logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
with patch("litellm.callbacks", [guardrail]):
resp = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
stream=True,
litellm_logging_obj=mock_logging_obj,
)
async for _ in resp:
pass
await asyncio.sleep(0)
await asyncio.sleep(0)
assert hook_called is True, \
"Production closure must call guardrail hook"
assert logged_response is modified_response, \
"Production closure must pass guardrail-modified response to logging"
@pytest.mark.asyncio
async def test_apply_guardrail_path_uses_unified_guardrail(self):
"""Guardrails that define apply_guardrail should be dispatched through
UnifiedLLMGuardrails.async_post_call_success_hook via the real
_run_deferred_stream_guardrails static method."""
from litellm.types.utils import GenericGuardrailAPIInputs
unified_hook_called = False
class ApplyGuardrailType(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="apply-type",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def apply_guardrail(
self, inputs, request_data, input_type, logging_obj=None
) -> GenericGuardrailAPIInputs:
nonlocal unified_hook_called
unified_hook_called = True
return inputs
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
logged_response = None
async def track_async_success(*args, **kwargs):
nonlocal logged_response
logged_response = args[0] if args else None
mock_logging_obj.async_success_handler = track_async_success
guardrail = ApplyGuardrailType()
async def _on_deferred_stream_complete(assembled_response, cache_hit):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={"model": "gpt-4", "metadata": {}},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=assembled_response,
cache_hit=cache_hit,
)
mock_logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete
with patch("litellm.callbacks", [guardrail]):
resp = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello!",
stream=True,
litellm_logging_obj=mock_logging_obj,
)
async for _ in resp:
pass
await asyncio.sleep(0)
await asyncio.sleep(0)
assert unified_hook_called is True, \
"apply_guardrail guardrails must be dispatched through UnifiedLLMGuardrails"
assert logged_response is not None, \
"Logging must fire after unified guardrail path"
@pytest.mark.asyncio
async def test_hooks_receive_merged_guardrail_data(self):
"""Hooks must receive guardrail_data (the merged dict from
_check_and_merge_model_level_guardrails), not the original
captured_data. This ensures model-level non-default guardrails
are visible to any inner should_run_guardrail re-checks.
Uses a deep-copy mock to break the shallow-copy side-effect that
would otherwise mask the bug — verifying the code is explicitly
correct, not correct-by-accident."""
import copy
hook_received_data = None
class InspectingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="inspector",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
nonlocal hook_received_data
hook_received_data = data
return response
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
pass
mock_logging_obj.async_success_handler = track_async_success
guardrail = InspectingGuardrail()
captured_data = {"model": "gpt-4", "metadata": {"existing_key": "value"}}
def mock_merge(data, llm_router):
"""Return a fully independent dict (deep copy) so the original
captured_data is NOT mutated. This simulates a correct merge
implementation and proves _run_deferred_stream_guardrails uses
the return value, not the original data."""
merged = copy.deepcopy(data)
merged["metadata"]["guardrails"] = ["model-guardrail"]
merged["_merged_marker"] = True
return merged
with patch("litellm.callbacks", [guardrail]), \
patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=captured_data,
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
assert hook_received_data is not None, "Guardrail hook must be called"
assert hook_received_data.get("_merged_marker") is True, \
"Hook must receive guardrail_data (merged), not original captured_data"
assert "model-guardrail" in hook_received_data.get("metadata", {}).get(
"guardrails", []
), "Hook data must contain model-level guardrails"
@pytest.mark.asyncio
async def test_apply_guardrail_path_receives_merged_guardrail_data(self):
"""The apply_guardrail path (through UnifiedLLMGuardrails) must also
receive guardrail_data so that the inner should_run_guardrail re-check
inside UnifiedLLMGuardrails sees model-level guardrails.
This is the specific scenario Greptile flagged: a default_on=False
guardrail configured at the model level would pass the outer gate but
be silently skipped at execution time if captured_data (unmerged) were
passed instead of guardrail_data (merged)."""
import copy
from litellm.types.utils import GenericGuardrailAPIInputs
unified_received_data = None
class ModelLevelApplyGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="model-apply-guardrail",
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def apply_guardrail(
self, inputs, request_data, input_type, logging_obj=None
) -> GenericGuardrailAPIInputs:
return inputs
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
pass
mock_logging_obj.async_success_handler = track_async_success
guardrail = ModelLevelApplyGuardrail()
captured_data = {"model": "gpt-4", "metadata": {}}
def mock_merge(data, llm_router):
merged = copy.deepcopy(data)
merged["metadata"]["guardrails"] = ["model-apply-guardrail"]
merged["_merged_marker"] = True
return merged
# Capture what UnifiedLLMGuardrails.async_post_call_success_hook receives
original_unified_hook = None
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
original_unified_hook = UnifiedLLMGuardrails.async_post_call_success_hook
async def tracking_unified_hook(self, user_api_key_dict, data, response):
nonlocal unified_received_data
unified_received_data = data
return response
with patch("litellm.callbacks", [guardrail]), \
patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
), \
patch.object(
UnifiedLLMGuardrails,
"async_post_call_success_hook",
tracking_unified_hook,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=captured_data,
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
assert unified_received_data is not None, \
"UnifiedLLMGuardrails must be called for apply_guardrail guardrails"
assert unified_received_data.get("_merged_marker") is True, \
"UnifiedLLMGuardrails must receive guardrail_data (merged), not captured_data"
assert "model-apply-guardrail" in unified_received_data.get(
"metadata", {}
).get("guardrails", []), \
"UnifiedLLMGuardrails data must contain model-level guardrails"
@pytest.mark.asyncio
async def test_multiple_guardrails_all_receive_merged_data(self):
"""When multiple guardrails are configured, ALL of them must receive
guardrail_data (merged), not just the first one."""
import copy
received_data_per_guardrail = {}
class TaggedGuardrail(CustomGuardrail):
def __init__(self, name):
super().__init__(
guardrail_name=name,
default_on=True,
event_hook=GuardrailEventHooks.post_call,
)
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any
) -> Any:
received_data_per_guardrail[self.guardrail_name] = data
return response
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
pass
mock_logging_obj.async_success_handler = track_async_success
guardrail_a = TaggedGuardrail("guardrail-a")
guardrail_b = TaggedGuardrail("guardrail-b")
captured_data = {"model": "gpt-4", "metadata": {}}
def mock_merge(data, llm_router):
merged = copy.deepcopy(data)
merged["metadata"]["guardrails"] = ["guardrail-a", "guardrail-b"]
merged["_merged_marker"] = True
return merged
with patch("litellm.callbacks", [guardrail_a, guardrail_b]), \
patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=captured_data,
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
for name in ("guardrail-a", "guardrail-b"):
assert name in received_data_per_guardrail, \
f"{name} must be called"
assert received_data_per_guardrail[name].get("_merged_marker") is True, \
f"{name} must receive guardrail_data (merged), not captured_data"
@pytest.mark.asyncio
async def test_logging_fires_even_if_guardrail_init_raises(self):
"""If _check_and_merge_model_level_guardrails raises during
initialization, logging must still fire via the try/finally guard.
This prevents silent logging loss on transient init errors."""
logging_called = False
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {"metadata": {}}
async def track_async_success(*args, **kwargs):
nonlocal logging_called
logging_called = True
mock_logging_obj.async_success_handler = track_async_success
def exploding_merge(data, llm_router):
raise RuntimeError("Simulated init failure")
with patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=exploding_merge,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data={"model": "gpt-4", "metadata": {}},
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
captured_logging_obj=mock_logging_obj,
assembled_response=MagicMock(),
cache_hit=False,
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert logging_called is True, \
"Logging must fire even when guardrail initialization raises"

View file

@ -6441,3 +6441,53 @@ async def test_list_team_v1_batches_key_queries():
assert result[0].keys == [key1, key2]
assert result[1].team_id == "team-2"
assert result[1].keys == [key3]
def test_new_team_request_accepts_team_member_budget_duration():
"""Test that NewTeamRequest does not silently drop team_member_budget_duration."""
from litellm.proxy._types import NewTeamRequest
request = NewTeamRequest(
team_member_budget=20.0,
team_member_budget_duration="30d",
)
assert request.team_member_budget == 20.0
assert request.team_member_budget_duration == "30d"
@pytest.mark.asyncio
async def test_create_team_member_budget_table_with_duration():
"""Verify that create_team_member_budget_table passes budget_duration
through to the new_budget call when team_member_budget_duration is provided."""
from litellm.proxy._types import NewTeamRequest, UserAPIKeyAuth, LitellmUserRoles
from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler
mock_budget_response = MagicMock(budget_id="budget-abc")
mock_admin = UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
)
data = NewTeamRequest(
team_alias="test-team",
team_member_budget=20.0,
team_member_budget_duration="30d",
)
with patch(
"litellm.proxy.management_endpoints.budget_management_endpoints.new_budget",
new_callable=AsyncMock,
return_value=mock_budget_response,
) as mock_new_budget:
result = await TeamMemberBudgetHandler.create_team_member_budget_table(
data=data,
new_team_data_json={"metadata": None},
user_api_key_dict=mock_admin,
team_member_budget=20.0,
team_member_budget_duration="30d",
)
mock_new_budget.assert_awaited_once()
budget_request = mock_new_budget.call_args.kwargs["budget_obj"]
assert budget_request.budget_duration == "30d"
assert budget_request.max_budget == 20.0
assert result["metadata"]["team_member_budget_id"] == "budget-abc"

View file

@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import Request
from fastapi import HTTPException, Request
from litellm._uuid import uuid
@ -5160,3 +5160,99 @@ def test_generic_response_convertor_extra_attributes_missing_field(monkeypatch):
assert result.extra_fields["missing_field"] is None
assert result.extra_fields["another_missing"] is None
class TestValidateReturnTo:
"""Tests for SSOAuthenticationHandler._validate_return_to"""
def test_rejects_when_no_control_plane_url_configured(self, monkeypatch):
"""return_to should be rejected if control_plane_url is not in general_settings."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", {}
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui")
assert exc_info.value.status_code == 400
assert "not configured" in exc_info.value.detail
def test_allows_matching_origin(self, monkeypatch):
"""return_to matching the configured control_plane_url origin should pass."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
# Should not raise
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui?page=models")
def test_allows_matching_origin_with_trailing_slash(self, monkeypatch):
"""Trailing slash on control_plane_url should not affect origin comparison."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com/"},
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui")
def test_rejects_prefix_attack(self, monkeypatch):
"""return_to like cp.example.com.evil.com must be rejected (not just prefix match)."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("https://cp.example.com.evil.com/steal")
assert exc_info.value.status_code == 400
def test_rejects_different_origin(self, monkeypatch):
"""return_to pointing to a completely different domain should be rejected."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("https://evil.com/phish")
assert exc_info.value.status_code == 400
def test_case_insensitive_hostname(self, monkeypatch):
"""Hostname comparison should be case-insensitive per RFC 3986."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://CP.Example.COM"},
)
# Should not raise
SSOAuthenticationHandler._validate_return_to("https://cp.example.com/ui")
def test_rejects_scheme_mismatch(self, monkeypatch):
"""http:// must be rejected when control_plane_url uses https://."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("http://cp.example.com/ui")
assert exc_info.value.status_code == 400
def test_rejects_port_mismatch(self, monkeypatch):
"""Non-default port must be rejected."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
with pytest.raises(HTTPException) as exc_info:
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:8443/ui")
assert exc_info.value.status_code == 400
def test_allows_explicit_default_port(self, monkeypatch):
"""https://host:443 should match https://host (default port normalisation)."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:443/ui")
def test_allows_matching_custom_port(self, monkeypatch):
"""Both sides on the same custom port should match."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com:3000"},
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:3000/ui")

View file

@ -0,0 +1,167 @@
"""
Tests verifying that default_api_key_tpm_limit and default_api_key_rpm_limit set in
litellm_params are returned by the /model/info endpoint.
"""
from typing import Optional
from unittest.mock import MagicMock, patch
import pytest
from litellm.proxy.proxy_server import _get_proxy_model_info
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
def _make_deployment(
model_name: str,
default_tpm: Optional[int] = None,
default_rpm: Optional[int] = None,
) -> Deployment:
params: dict = {"model": f"openai/{model_name}"}
if default_tpm is not None:
params["default_api_key_tpm_limit"] = default_tpm
if default_rpm is not None:
params["default_api_key_rpm_limit"] = default_rpm
return Deployment(
model_name=model_name,
litellm_params=LiteLLM_Params(**params),
model_info=ModelInfo(),
)
class TestModelInfoDefaultLimitsInResponse:
"""
Verify _get_proxy_model_info (the helper used by the /model/info endpoint) returns
default_api_key_tpm_limit and default_api_key_rpm_limit from litellm_params.
"""
def test_default_tpm_and_rpm_present_in_model_info_response(self):
"""Both defaults should appear in the litellm_params section of the response."""
deployment = _make_deployment("model1", default_tpm=100, default_rpm=200)
model_dict = deployment.model_dump(exclude_none=True)
result = _get_proxy_model_info(model=model_dict)
litellm_params = result["litellm_params"]
assert litellm_params.get("default_api_key_tpm_limit") == 100
assert litellm_params.get("default_api_key_rpm_limit") == 200
def test_default_tpm_only_present_when_only_tpm_configured(self):
"""Only the configured default appears; the other stays absent."""
deployment = _make_deployment("model1", default_tpm=500)
model_dict = deployment.model_dump(exclude_none=True)
result = _get_proxy_model_info(model=model_dict)
litellm_params = result["litellm_params"]
assert litellm_params.get("default_api_key_tpm_limit") == 500
assert "default_api_key_rpm_limit" not in litellm_params
def test_default_rpm_only_present_when_only_rpm_configured(self):
"""Only the configured default appears; the other stays absent."""
deployment = _make_deployment("model1", default_rpm=300)
model_dict = deployment.model_dump(exclude_none=True)
result = _get_proxy_model_info(model=model_dict)
litellm_params = result["litellm_params"]
assert litellm_params.get("default_api_key_rpm_limit") == 300
assert "default_api_key_tpm_limit" not in litellm_params
def test_defaults_absent_when_not_configured(self):
"""Neither field appears when not set on the deployment."""
deployment = _make_deployment("model1")
model_dict = deployment.model_dump(exclude_none=True)
result = _get_proxy_model_info(model=model_dict)
litellm_params = result["litellm_params"]
assert "default_api_key_tpm_limit" not in litellm_params
assert "default_api_key_rpm_limit" not in litellm_params
def test_defaults_not_masked_or_stripped_by_sensitive_data_filter(self):
"""
default_api_key_tpm_limit / default_api_key_rpm_limit must not be
treated as sensitive and must survive remove_sensitive_info_from_deployment.
They contain "key" which normally triggers masking; the call site explicitly
excludes these two fields via excluded_keys rather than widening the global
non_sensitive_overrides.
"""
deployment = _make_deployment("model1", default_tpm=100, default_rpm=200)
model_dict = deployment.model_dump(exclude_none=True)
result = _get_proxy_model_info(model=model_dict)
# Values should be unchanged integers, not masked strings
assert result["litellm_params"]["default_api_key_tpm_limit"] == 100
assert result["litellm_params"]["default_api_key_rpm_limit"] == 200
class TestModelInfoEndpointWithRouter:
"""
Integration-style tests simulating the /model/info endpoint reading from the router.
"""
@pytest.mark.asyncio
async def test_model_info_endpoint_returns_defaults_for_specific_model_id(self):
"""
When litellm_model_id is provided, the endpoint should return the deployment's
default limits in litellm_params.
"""
from litellm.proxy.proxy_server import model_info_v1
from litellm.proxy._types import UserAPIKeyAuth
deployment = _make_deployment("model1", default_tpm=100, default_rpm=200)
mock_router = MagicMock()
mock_router.get_deployment.return_value = deployment
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test")
with patch("litellm.proxy.proxy_server.llm_router", mock_router), \
patch("litellm.proxy.proxy_server.llm_model_list", []), \
patch("litellm.proxy.proxy_server.user_model", None):
response = await model_info_v1(
user_api_key_dict=user_api_key_dict,
litellm_model_id="some-model-id",
)
assert len(response["data"]) == 1
litellm_params = response["data"][0]["litellm_params"]
assert litellm_params.get("default_api_key_tpm_limit") == 100
assert litellm_params.get("default_api_key_rpm_limit") == 200
@pytest.mark.asyncio
async def test_model_info_endpoint_returns_defaults_in_full_model_list(self):
"""
Without litellm_model_id, the endpoint iterates all models. Each deployment's
default limits should appear in its litellm_params entry.
"""
from litellm.proxy.proxy_server import model_info_v1
from litellm.proxy._types import UserAPIKeyAuth
deployment = _make_deployment("model1", default_tpm=100, default_rpm=200)
deployment_dict = deployment.model_dump(exclude_none=True)
mock_router = MagicMock()
mock_router.get_model_names.return_value = ["model1"]
mock_router.get_model_access_groups.return_value = {}
mock_router.get_model_list.return_value = [deployment_dict]
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test")
with patch("litellm.proxy.proxy_server.llm_router", mock_router), \
patch("litellm.proxy.proxy_server.llm_model_list", [deployment_dict]), \
patch("litellm.proxy.proxy_server.user_model", None), \
patch("litellm.proxy.proxy_server.get_key_models", return_value=["model1"]), \
patch("litellm.proxy.proxy_server.get_team_models", return_value=["model1"]), \
patch("litellm.proxy.proxy_server.get_complete_model_list", return_value=["model1"]):
response = await model_info_v1(
user_api_key_dict=user_api_key_dict,
litellm_model_id=None,
)
assert len(response["data"]) >= 1
litellm_params = response["data"][0]["litellm_params"]
assert litellm_params.get("default_api_key_tpm_limit") == 100
assert litellm_params.get("default_api_key_rpm_limit") == 200

View file

@ -236,6 +236,217 @@ def test_login_v2_returns_json_on_invalid_json_body(monkeypatch):
assert isinstance(data["error"], dict)
def test_login_v3_rejected_without_control_plane_url(monkeypatch):
"""v3/login returns 404 when control_plane_url is not configured."""
mock_prisma_client = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
client = TestClient(app)
response = client.post(
"/v3/login",
json={"username": "alice", "password": "secret"},
)
assert response.status_code == 404
assert "control_plane_url" in response.json()["error"]["message"]
def test_login_v3_returns_code(monkeypatch):
"""v3/login returns an opaque code, not the JWT directly."""
mock_prisma_client = MagicMock()
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.create_ui_token_object",
MagicMock(return_value={"user_id": "test-user"}),
)
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_config = MagicMock()
mock_config.worker_registry = []
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
client = TestClient(app)
response = client.post(
"/v3/login",
json={"username": "alice", "password": "secret"},
)
assert response.status_code == 200
data = response.json()
assert "code" in data
assert data["expires_in"] == 60
assert "token" not in data
def test_login_v3_exchange_happy_path(monkeypatch):
"""Full flow: v3/login returns code, v3/login/exchange redeems it for JWT."""
mock_prisma_client = MagicMock()
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.create_ui_token_object",
MagicMock(return_value={"user_id": "test-user"}),
)
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_config = MagicMock()
mock_config.worker_registry = []
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
client = TestClient(app)
# Step 1: login — get code
login_response = client.post(
"/v3/login",
json={"username": "alice", "password": "secret"},
)
assert login_response.status_code == 200
code = login_response.json()["code"]
# Step 2: exchange — get JWT
exchange_response = client.post(
"/v3/login/exchange",
json={"code": code},
)
assert exchange_response.status_code == 200
exchange_data = exchange_response.json()
assert exchange_data["token"] == "signed-token"
assert "redirect_url" in exchange_data
assert exchange_response.cookies.get("token") == "signed-token"
def test_login_v3_exchange_single_use(monkeypatch):
"""Code can only be redeemed once."""
mock_prisma_client = MagicMock()
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
AsyncMock(return_value={"user_id": "test-user"}),
)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.create_ui_token_object",
MagicMock(return_value={"user_id": "test-user"}),
)
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_config = MagicMock()
mock_config.worker_registry = []
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
client = TestClient(app)
login_response = client.post(
"/v3/login",
json={"username": "alice", "password": "secret"},
)
code = login_response.json()["code"]
# First exchange succeeds
first = client.post("/v3/login/exchange", json={"code": code})
assert first.status_code == 200
# Second exchange fails
second = client.post("/v3/login/exchange", json={"code": code})
assert second.status_code == 401
def test_login_v3_exchange_invalid_code(monkeypatch):
"""Random code returns 401."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
client = TestClient(app)
response = client.post(
"/v3/login/exchange",
json={"code": "nonexistent-code"},
)
assert response.status_code == 401
def test_login_v3_exchange_rejected_without_control_plane_url(monkeypatch):
"""v3/login/exchange returns 404 when control_plane_url is not configured."""
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
client = TestClient(app)
response = client.post(
"/v3/login/exchange",
json={"code": "some-code"},
)
assert response.status_code == 404
assert "control_plane_url" in response.json()["error"]["message"]
def test_login_v3_returns_json_on_proxy_exception(monkeypatch):
"""Test that /v3/login returns JSON error when ProxyException is raised"""
from litellm.proxy._types import ProxyErrorTypes, ProxyException
mock_prisma_client = MagicMock()
mock_authenticate_user = AsyncMock(
side_effect=ProxyException(
message="Invalid credentials",
type=ProxyErrorTypes.auth_error,
param="password",
code=401,
)
)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
mock_authenticate_user,
)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"control_plane_url": "https://cp.example.com"},
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
client = TestClient(app)
response = client.post(
"/v3/login",
json={"username": "alice", "password": "wrong"},
)
assert response.status_code == 401
assert response.headers["content-type"] == "application/json"
data = response.json()
assert "error" in data
assert data["error"]["message"] == "Invalid credentials"
assert data["error"]["type"] == "auth_error"
def test_fallback_login_has_no_deprecation_banner(client_no_auth):
response = client_no_auth.get("/fallback/login")

View file

@ -1,16 +1,10 @@
"use client";
import APIReferenceView from "@/app/(dashboard)/api-reference/APIReferenceView";
import { useState } from "react";
interface ProxySettings {
PROXY_BASE_URL: string;
PROXY_LOGOUT_URL: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
}
import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings";
const APIReferencePage = () => {
const [proxySettings, setProxySettings] = useState<ProxySettings>({ PROXY_BASE_URL: "", PROXY_LOGOUT_URL: "" });
const proxySettings = useProxySettings();
return <APIReferenceView proxySettings={proxySettings} />;
};

View file

@ -195,7 +195,7 @@ const menuItems: MenuItemCfg[] = [
icon: <UserOutlined style={{ fontSize: 18 }} />,
roles: all_admin_roles,
},
{ key: "14", page: "api_ref", label: "API Reference", icon: <ApiOutlined style={{ fontSize: 18 }} /> },
{ key: "14", page: "api-reference", label: "API Reference", icon: <ApiOutlined style={{ fontSize: 18 }} /> },
{
key: "16",
page: "model-hub-table",

View file

@ -0,0 +1,34 @@
import { describe, it, expect } from "vitest";
import { createQueryKeys } from "./queryKeysFactory";
describe("createQueryKeys", () => {
const keys = createQueryKeys("books");
it("should return the resource name as the base key", () => {
expect(keys.all).toEqual(["books"]);
});
it("should generate a lists key", () => {
expect(keys.lists()).toEqual(["books", "list"]);
});
it("should generate a list key with params", () => {
expect(keys.list({ page: 1, limit: 10 })).toEqual([
"books",
"list",
{ params: { page: 1, limit: 10 } },
]);
});
it("should generate a list key with undefined params when none provided", () => {
expect(keys.list()).toEqual(["books", "list", { params: undefined }]);
});
it("should generate a details key", () => {
expect(keys.details()).toEqual(["books", "detail"]);
});
it("should generate a detail key for a specific ID", () => {
expect(keys.detail("123")).toEqual(["books", "detail", "123"]);
});
});

View file

@ -3,8 +3,8 @@ import { loginCall, LoginRequest } from "@/components/networking";
export const useLogin = () => {
return useMutation({
mutationFn: async ({ username, password }: LoginRequest) => {
const result = await loginCall(username, password);
mutationFn: async ({ username, password, useV3 }: LoginRequest) => {
const result = await loginCall(username, password, useV3);
return result;
},
});

View file

@ -0,0 +1,21 @@
import { useState, useEffect } from "react";
import { fetchProxySettings } from "@/utils/proxyUtils";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
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,
});
useEffect(() => {
if (!accessToken) return;
fetchProxySettings(accessToken).then((settings) => {
if (settings) setProxySettings(settings);
});
}, [accessToken]);
return proxySettings;
}

View file

@ -28,6 +28,8 @@ const mockUIConfig: LiteLLMWellKnownUiConfig = {
proxy_base_url: "https://proxy.example.com",
auto_redirect_to_sso: true,
admin_ui_disabled: false,
is_control_plane: false,
workers: [],
};
describe("useUIConfig", () => {
@ -102,6 +104,8 @@ describe("useUIConfig", () => {
auto_redirect_to_sso: false,
sso_configured: false,
admin_ui_disabled: true,
is_control_plane: false,
workers: [],
};
// Mock successful API call with different data

View file

@ -0,0 +1,54 @@
import { render, screen } from "@testing-library/react";
import React from "react";
import { describe, expect, it, vi } from "vitest";
import TeamsHeaderTabs from "./TeamsHeaderTabs";
vi.mock("@tremor/react", () => ({
TabGroup: ({ children, ...props }: any) => <div data-testid="tab-group" {...props}>{children}</div>,
TabList: ({ children, ...props }: any) => <div data-testid="tab-list" {...props}>{children}</div>,
Tab: ({ children, ...props }: any) => <button {...props}>{children}</button>,
TabPanels: ({ children, ...props }: any) => <div data-testid="tab-panels" {...props}>{children}</div>,
Text: ({ children, ...props }: any) => <span {...props}>{children}</span>,
Icon: ({ onClick, ...props }: any) => <button data-testid="refresh-icon" onClick={onClick} />,
}));
vi.mock("@heroicons/react/outline", () => ({
RefreshIcon: () => <svg data-testid="refresh-svg" />,
}));
const renderTabs = (props: Partial<Parameters<typeof TeamsHeaderTabs>[0]> = {}) => {
const defaults = {
lastRefreshed: "",
onRefresh: vi.fn(),
userRole: "Internal User",
children: <div data-testid="panel-content">Panel</div>,
};
return render(<TeamsHeaderTabs {...defaults} {...props} />);
};
describe("TeamsHeaderTabs", () => {
it("should render 'Your Teams' and 'Available Teams' tabs", () => {
renderTabs();
expect(screen.getByText("Your Teams")).toBeInTheDocument();
expect(screen.getByText("Available Teams")).toBeInTheDocument();
});
it("should render 'Default Team Settings' tab when user is Admin", () => {
renderTabs({ userRole: "Admin" });
expect(screen.getByText("Default Team Settings")).toBeInTheDocument();
});
it("should not render 'Default Team Settings' tab for non-admin users", () => {
renderTabs({ userRole: "Internal User" });
expect(screen.queryByText("Default Team Settings")).not.toBeInTheDocument();
});
it("should display last refreshed time when provided", () => {
renderTabs({ lastRefreshed: "2024-06-01 12:00:00" });
expect(screen.getByText("Last Refreshed: 2024-06-01 12:00:00")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,129 @@
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import React from "react";
import { describe, expect, it, vi } from "vitest";
import { Team } from "@/components/key_team_helpers/key_list";
import TeamsTable from "./TeamsTable";
vi.mock("@tremor/react", () => ({
Button: React.forwardRef<HTMLButtonElement, any>(({ children, ...props }, ref) =>
React.createElement("button", { ...props, ref }, children),
),
Icon: ({ onClick, ...props }: any) => <button data-testid={props["data-testid"] || "icon-btn"} onClick={onClick} aria-label={props["aria-label"]} />,
Table: ({ children }: any) => <table>{children}</table>,
TableHead: ({ children }: any) => <thead>{children}</thead>,
TableBody: ({ children }: any) => <tbody>{children}</tbody>,
TableRow: ({ children }: any) => <tr>{children}</tr>,
TableHeaderCell: ({ children }: any) => <th>{children}</th>,
TableCell: ({ children, ...props }: any) => <td {...props}>{children}</td>,
Text: ({ children }: any) => <span>{children}</span>,
}));
vi.mock("antd", () => ({
Tooltip: ({ children }: any) => <>{children}</>,
}));
vi.mock("@heroicons/react/outline", () => ({
PencilAltIcon: () => <svg data-testid="pencil-icon" />,
TrashIcon: () => <svg data-testid="trash-icon" />,
}));
vi.mock("@/utils/dataUtils", () => ({
formatNumberWithCommas: (val: number, decimals: number) =>
val != null ? val.toFixed(decimals) : "N/A",
}));
vi.mock("@/app/(dashboard)/teams/components/TeamsTable/ModelsCell", () => ({
default: ({ team }: any) => <td data-testid="models-cell">{team.models.join(",")}</td>,
}));
vi.mock("@/app/(dashboard)/teams/components/TeamsTable/YourRoleCell/YourRoleCell", () => ({
default: ({ team }: any) => <td data-testid="role-cell">{team.team_id}</td>,
}));
const makeTeam = (overrides: Partial<Team> = {}): Team => ({
team_id: "team-abc1234",
team_alias: "Platform",
models: ["gpt-4"],
max_budget: 500,
budget_duration: null,
tpm_limit: null,
rpm_limit: null,
organization_id: "org-1",
created_at: "2024-06-01T00:00:00Z",
keys: [],
members_with_roles: [],
spend: 123.4567,
...overrides,
});
const defaultPerTeamInfo = {
"team-abc1234": {
keys: [{ token: "tok-1" } as any, { token: "tok-2" } as any],
team_info: {
members_with_roles: [{ user_id: "u1", role: "admin" } as any],
},
},
};
const renderTable = (overrides: Partial<Parameters<typeof TeamsTable>[0]> = {}) => {
const defaults = {
teams: [makeTeam()],
currentOrg: null,
perTeamInfo: defaultPerTeamInfo,
userRole: "Admin",
userId: "user-1",
setSelectedTeamId: vi.fn(),
setEditTeam: vi.fn(),
onDeleteTeam: vi.fn(),
};
return render(<TeamsTable {...defaults} {...overrides} />);
};
describe("TeamsTable", () => {
it("should render table headers", () => {
renderTable();
expect(screen.getByText("Team Name")).toBeInTheDocument();
expect(screen.getByText("Team ID")).toBeInTheDocument();
expect(screen.getByText("Created")).toBeInTheDocument();
expect(screen.getByText("Spend (USD)")).toBeInTheDocument();
expect(screen.getByText("Budget (USD)")).toBeInTheDocument();
expect(screen.getByText("Models")).toBeInTheDocument();
expect(screen.getByText("Organization")).toBeInTheDocument();
expect(screen.getByText("Your Role")).toBeInTheDocument();
expect(screen.getByText("Info")).toBeInTheDocument();
});
it("should render team rows with team data", () => {
renderTable();
expect(screen.getByText("Platform")).toBeInTheDocument();
expect(screen.getByText("team-ab...")).toBeInTheDocument();
expect(screen.getByText("org-1")).toBeInTheDocument();
});
it("should show edit and delete icons for Admin users", () => {
renderTable({ userRole: "Admin" });
expect(screen.getAllByTestId("icon-btn").length).toBeGreaterThanOrEqual(2);
});
it("should not show edit and delete icons for non-Admin users", () => {
renderTable({ userRole: "Internal User" });
// Only the team ID button should be present, no icon-btn for edit/delete
const iconBtns = screen.queryAllByTestId("icon-btn");
expect(iconBtns).toHaveLength(0);
});
it("should call setSelectedTeamId when team ID button is clicked", async () => {
const user = userEvent.setup();
const setSelectedTeamId = vi.fn();
renderTable({ setSelectedTeamId });
await user.click(screen.getByText("team-ab..."));
expect(setSelectedTeamId).toHaveBeenCalledWith("team-abc1234");
});
});

View file

@ -41,6 +41,17 @@ vi.mock("@/app/(dashboard)/hooks/login/useLogin", () => ({
})),
}));
vi.mock("@/hooks/useWorker", () => ({
useWorker: vi.fn(() => ({
isControlPlane: false,
workers: [],
selectedWorkerId: null,
selectedWorker: null,
selectWorker: vi.fn(),
disconnectFromWorker: vi.fn(),
})),
}));
import { useUIConfig } from "@/app/(dashboard)/hooks/uiConfig/useUIConfig";
import { getCookie } from "@/utils/cookieUtils";
import { isJwtExpired } from "@/utils/jwtUtils";
@ -108,7 +119,7 @@ describe("LoginPage", () => {
);
await waitFor(() => {
expect(mockReplace).toHaveBeenCalledWith("http://localhost:4000/ui");
expect(mockReplace).toHaveBeenCalledWith("/ui");
});
});
@ -189,7 +200,7 @@ describe("LoginPage", () => {
);
await waitFor(() => {
expect(mockReplace).toHaveBeenCalledWith("http://localhost:4000/ui");
expect(mockReplace).toHaveBeenCalledWith("/ui");
});
expect(mockPush).not.toHaveBeenCalled();

View file

@ -3,14 +3,15 @@
import { useLogin } from "@/app/(dashboard)/hooks/login/useLogin";
import { useUIConfig } from "@/app/(dashboard)/hooks/uiConfig/useUIConfig";
import LoadingScreen from "@/components/common_components/LoadingScreen";
import { getProxyBaseUrl } from "@/components/networking";
import { getCookie } from "@/utils/cookieUtils";
import { exchangeLoginCode, getProxyBaseUrl, switchToWorkerUrl } from "@/components/networking";
import { clearTokenCookies, getCookie } from "@/utils/cookieUtils";
import { isJwtExpired } from "@/utils/jwtUtils";
import { consumeReturnUrl, getReturnUrl, isValidReturnUrl } from "@/utils/returnUrlUtils";
import { InfoCircleOutlined } from "@ant-design/icons";
import { Alert, Button, Card, Form, Input, Popover, Space, Typography } from "antd";
import { InfoCircleOutlined, CloudServerOutlined } from "@ant-design/icons";
import { Alert, Button, Card, Form, Input, Popover, Select, Space, Typography } from "antd";
import { useRouter } from "next/navigation";
import { useEffect, useState } from "react";
import { useWorker } from "@/hooks/useWorker";
function LoginPageContent() {
const [username, setUsername] = useState("");
@ -19,6 +20,17 @@ function LoginPageContent() {
const { data: uiConfig, isLoading: isConfigLoading } = useUIConfig();
const loginMutation = useLogin();
const router = useRouter();
const { workers, selectWorker } = useWorker();
const [selectedWorkerId, setSelectedWorkerId] = useState<string | null>(null);
// Pre-select worker from URL param (e.g. /ui/login?worker=team-b)
useEffect(() => {
const params = new URLSearchParams(window.location.search);
const workerParam = params.get("worker");
if (workerParam) {
setSelectedWorkerId(workerParam);
}
}, []);
useEffect(() => {
if (isConfigLoading) {
@ -31,6 +43,44 @@ function LoginPageContent() {
return;
}
// Cross-origin SSO: worker redirected back with a single-use code.
// Exchange it for the JWT via the worker's /v3/login/exchange endpoint.
const params = new URLSearchParams(window.location.search);
const ssoCode = params.get("code");
if (ssoCode) {
const workerUrl = localStorage.getItem("litellm_worker_url");
exchangeLoginCode(ssoCode, workerUrl).then(() => {
params.delete("code");
const cleanSearch = params.toString();
window.history.replaceState(null, "", window.location.pathname + (cleanSearch ? `?${cleanSearch}` : ""));
router.replace("/ui/?login=success");
});
return;
}
// Backwards compat: handle direct token in URL (legacy flow)
const urlToken = params.get("token");
if (urlToken && !isJwtExpired(urlToken)) {
document.cookie = `token=${urlToken}; path=/; SameSite=Lax`;
params.delete("token");
const cleanSearch = params.toString();
window.history.replaceState(
null,
"",
window.location.pathname + (cleanSearch ? `?${cleanSearch}` : ""),
);
router.replace("/ui/?login=success");
return;
}
// If switching workers on a control plane, clear the old token and show login
const switchingWorker = params.has("worker");
if (switchingWorker && uiConfig?.is_control_plane) {
clearTokenCookies();
setIsLoading(false);
return;
}
const rawToken = getCookie("token");
if (rawToken && !isJwtExpired(rawToken)) {
// User already logged in - redirect to return URL or default
@ -38,7 +88,7 @@ function LoginPageContent() {
if (returnUrl) {
router.replace(returnUrl);
} else {
router.replace(`${getProxyBaseUrl()}/ui`);
router.replace("/ui");
}
return;
}
@ -58,16 +108,35 @@ function LoginPageContent() {
}, [isConfigLoading, router, uiConfig]);
const handleSubmit = () => {
// If a worker is selected, point proxyBaseUrl at it before login
const selectedWorker = workers.find((w) => w.worker_id === selectedWorkerId);
if (selectedWorker) {
switchToWorkerUrl(selectedWorker.url);
}
loginMutation.mutate(
{ username, password },
{ username, password, useV3: !!selectedWorker },
{
onSuccess: (data) => {
// Check if we have a return URL to use instead of the default redirect
const returnUrl = consumeReturnUrl();
if (returnUrl) {
router.push(returnUrl);
// Update the worker context with the selected worker
if (selectedWorker) {
selectWorker(selectedWorker.worker_id);
// Stay on the CP's UI — proxyBaseUrl already points at the worker
router.push("/ui/?login=success");
} else {
router.push(data.redirect_url);
// Normal (non-control-plane) login — follow the server's redirect
const returnUrl = consumeReturnUrl();
if (returnUrl) {
router.push(returnUrl);
} else {
router.push(data.redirect_url);
}
}
},
onError: () => {
// Reset proxyBaseUrl on login failure
if (selectedWorker) {
switchToWorkerUrl(null);
}
},
},
@ -154,6 +223,22 @@ function LoginPageContent() {
{error && <Alert message={error} type="error" showIcon />}
<Form onFinish={handleSubmit} layout="vertical" requiredMark={true}>
{uiConfig?.is_control_plane && workers.length > 0 && (
<Form.Item label="Worker" style={{ marginBottom: 16 }}>
<Select
value={selectedWorkerId || undefined}
onChange={(value) => setSelectedWorkerId(value)}
placeholder="Choose a worker to connect to"
size="large"
suffixIcon={<CloudServerOutlined />}
options={workers.map((w) => ({
label: w.name,
value: w.worker_id,
}))}
/>
</Form.Item>
)}
<Form.Item
label="Username"
name="username"
@ -209,10 +294,20 @@ function LoginPageContent() {
</Popover>
) : (
<Button
disabled={isLoginLoading}
onClick={() =>
router.push(`${getProxyBaseUrl()}/sso/key/generate`)
}
disabled={isLoginLoading || (!!selectedWorkerId && workers.length === 0)}
onClick={() => {
const selectedWorker = workers.find((w) => w.worker_id === selectedWorkerId);
if (selectedWorker) {
// Store worker selection so useWorker hook restores it after redirect
localStorage.setItem("litellm_selected_worker_id", selectedWorkerId!);
switchToWorkerUrl(selectedWorker.url);
}
// SSO on the worker (or this instance if no worker), always
// include return_to so the callback redirects back here
const ssoBase = selectedWorker?.url ?? getProxyBaseUrl();
const returnTo = encodeURIComponent(window.location.origin + "/ui/login");
router.push(`${ssoBase}/sso/key/generate?return_to=${returnTo}`);
}}
block
size="large"
>

View file

@ -0,0 +1,95 @@
import { render, screen } from "@testing-library/react";
import React from "react";
import { describe, expect, it, vi } from "vitest";
import { OnboardingForm } from "./OnboardingForm";
const mockUseOnboardingCredentials = vi.fn();
const mockClaimToken = vi.fn();
vi.mock("next/navigation", () => ({
useSearchParams: () => new URLSearchParams("invitation_id=inv-123"),
}));
vi.mock("jwt-decode", () => ({
jwtDecode: vi.fn(() => ({
user_email: "alice@example.com",
user_id: "user-1",
key: "access-tok",
})),
}));
vi.mock("@/app/(dashboard)/hooks/onboarding/useOnboarding", () => ({
useOnboardingCredentials: (...args: unknown[]) => mockUseOnboardingCredentials(...args),
useClaimOnboardingToken: () => ({ mutate: mockClaimToken, isPending: false }),
}));
vi.mock("@/components/networking", () => ({
getProxyBaseUrl: vi.fn(() => ""),
}));
vi.mock("./OnboardingLoadingView", () => ({
OnboardingLoadingView: () => <div data-testid="loading-view">Loading</div>,
}));
vi.mock("./OnboardingErrorView", () => ({
OnboardingErrorView: () => <div data-testid="error-view">Error</div>,
}));
vi.mock("./OnboardingFormBody", () => ({
OnboardingFormBody: ({ variant, userEmail }: { variant: string; userEmail: string }) => (
<div data-testid="form-body" data-variant={variant} data-email={userEmail}>
Form Body
</div>
),
}));
describe("OnboardingForm", () => {
it("should render loading view when credentials are loading", () => {
mockUseOnboardingCredentials.mockReturnValue({
data: undefined,
isLoading: true,
isError: false,
});
render(<OnboardingForm variant="signup" />);
expect(screen.getByTestId("loading-view")).toBeInTheDocument();
});
it("should render error view when credentials fail to load", () => {
mockUseOnboardingCredentials.mockReturnValue({
data: undefined,
isLoading: false,
isError: true,
});
render(<OnboardingForm variant="signup" />);
expect(screen.getByTestId("error-view")).toBeInTheDocument();
});
it("should render form body with decoded email when credentials are loaded", () => {
mockUseOnboardingCredentials.mockReturnValue({
data: { token: "fake-jwt-token" },
isLoading: false,
isError: false,
});
render(<OnboardingForm variant="signup" />);
expect(screen.getByTestId("form-body")).toBeInTheDocument();
expect(screen.getByTestId("form-body")).toHaveAttribute("data-email", "alice@example.com");
});
it("should pass variant prop to OnboardingFormBody", () => {
mockUseOnboardingCredentials.mockReturnValue({
data: { token: "fake-jwt-token" },
isLoading: false,
isError: false,
});
render(<OnboardingForm variant="reset_password" />);
expect(screen.getByTestId("form-body")).toHaveAttribute("data-variant", "reset_password");
});
});

View file

@ -1,6 +1,5 @@
"use client";
import APIReferenceView from "@/app/(dashboard)/api-reference/APIReferenceView";
import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider";
import OldModelDashboard from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView";
import PlaygroundPage from "@/app/(dashboard)/playground/page";
@ -9,7 +8,7 @@ import AgentsPanel from "@/components/agents";
import BudgetPanel from "@/components/budgets/budget_panel";
import CacheDashboard from "@/components/cache_dashboard";
import ClaudeCodePluginsPanel from "@/components/claude_code_plugins";
import { fetchTeams } from "@/components/common_components/fetch_teams";
import { teamListCall as v2TeamListCall } from "@/app/(dashboard)/hooks/teams/useTeams";
import LoadingScreen from "@/components/common_components/LoadingScreen";
import { CostTrackingSettings } from "@/components/CostTrackingSettings";
import GeneralSettings from "@/components/general_settings";
@ -48,7 +47,7 @@ import { buildLoginUrlWithReturn, consumeReturnUrl, normalizeUrlForCompare, stor
import { formatUserRole, isAdminRole } from "@/utils/roles";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { jwtDecode } from "jwt-decode";
import { useSearchParams } from "next/navigation";
import { useRouter, useSearchParams } from "next/navigation";
import { Suspense, useEffect, useMemo, useRef, useState } from "react";
import { ConfigProvider, theme } from "antd";
@ -75,6 +74,16 @@ interface ProxySettings {
LITELLM_UI_API_DOC_BASE_URL?: string | null;
}
/**
* Map of legacy query-param page keys → new path-based route segments.
* When a user visits ?page=<key>, they are redirected to /ui/<value>.
* Add entries here as pages are migrated from the if/else chain to path-based routes.
*/
const LEGACY_REDIRECTS: Record<string, string> = {
api_ref: "api-reference",
"api-reference": "api-reference",
};
function CreateKeyPageContent() {
const [userRole, setUserRole] = useState("");
const [premiumUser, setPremiumUser] = useState(false);
@ -90,6 +99,7 @@ function CreateKeyPageContent() {
});
const [showSSOBanner, setShowSSOBanner] = useState<boolean>(true);
const router = useRouter();
const searchParams = useSearchParams()!;
const [modelData, setModelData] = useState<any>({ data: [] });
const [token, setToken] = useState<string | null>(null);
@ -243,6 +253,15 @@ function CreateKeyPageContent() {
}
}, [redirectToLogin]);
// Redirect legacy query-param pages to their new path-based routes
const isLegacyRedirect = page in LEGACY_REDIRECTS;
useEffect(() => {
if (!authLoading && isLegacyRedirect) {
const base = (proxyBaseUrl || "") + "/ui";
router.replace(`${base}/${LEGACY_REDIRECTS[page]}`);
}
}, [authLoading, isLegacyRedirect, page, router]);
// Check for a stored return URL after successful authentication
// This handles the case where user comes back from SSO and we need to redirect to the original URL
useEffect(() => {
@ -339,7 +358,9 @@ function CreateKeyPageContent() {
fetchUserModels(userID, userRole, accessToken, setUserModels);
}
if (accessToken && userID && userRole) {
fetchTeams(accessToken, userID, userRole, null, setTeams);
v2TeamListCall(accessToken, 1, 100, {
userID: userRole !== "Admin" && userRole !== "Admin Viewer" ? userID : null,
}).then((response) => setTeams(response.teams ?? [])).catch(console.error);
}
if (accessToken) {
fetchOrganizations(accessToken, setOrganizations);
@ -427,7 +448,7 @@ function CreateKeyPageContent() {
setShowClaudeCodePrompt(true);
};
if (authLoading || redirectToLogin) {
if (authLoading || redirectToLogin || isLegacyRedirect) {
return <LoadingScreen />;
}
@ -536,8 +557,6 @@ function CreateKeyPageContent() {
<AdminPanel
proxySettings={proxySettings}
/>
) : page == "api_ref" ? (
<APIReferenceView proxySettings={proxySettings} />
) : page == "logging-and-alerts" ? (
<Settings userID={userID} userRole={userRole} accessToken={accessToken} premiumUser={premiumUser} />
) : page == "budgets" ? (

View file

@ -1,5 +1,6 @@
import React from "react";
import { Modal, Form, message } from "antd";
import { Modal, Form } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import {
AccessGroupBaseForm,
AccessGroupFormValues,
@ -37,7 +38,7 @@ export function AccessGroupCreateModal({
createMutation.mutate(params, {
onSuccess: () => {
message.success("Access group created successfully");
MessageManager.success("Access group created successfully");
form.resetFields();
onSuccess?.();
onCancel();

View file

@ -1,5 +1,6 @@
import React, { useEffect } from "react";
import { Modal, Form, message } from "antd";
import { Modal, Form } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import {
AccessGroupBaseForm,
AccessGroupFormValues,
@ -55,7 +56,7 @@ export function AccessGroupEditModal({
{ accessGroupId: accessGroup.access_group_id, params },
{
onSuccess: () => {
message.success("Access group updated successfully");
MessageManager.success("Access group updated successfully");
onSuccess?.();
onCancel();
},

View file

@ -3,7 +3,6 @@ import {
Modal,
Typography,
Divider,
message,
Table,
Select,
InputNumber,
@ -14,6 +13,7 @@ import {
import { userBulkUpdateUserCall, teamBulkMemberAddCall, Member } from "./networking";
import { UserEditView } from "./user_edit_view";
import NotificationsManager from "./molecules/notifications_manager";
import MessageManager from "@/components/molecules/message_manager";
const { Text, Title } = Typography;
@ -188,7 +188,7 @@ const BulkEditUserModal: React.FC<BulkEditUserModalProps> = ({
}
if (failedTeams.length > 0) {
message.warning(`Failed to add users to ${failedTeams.length} team(s)`);
MessageManager.warning(`Failed to add users to ${failedTeams.length} team(s)`);
}
}

View file

@ -1,4 +1,5 @@
import { Form, Modal, Input, message } from "antd";
import { Form, Modal, Input } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { useEffect } from "react";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useCloudZeroCreate } from "@/app/(dashboard)/hooks/cloudzero/useCloudZeroCreate";
@ -31,7 +32,7 @@ export default function CloudZeroCreationModal({ open, onOk, onCancel }: CloudZe
},
{
onSuccess: () => {
message.success("CloudZero integration created successfully");
MessageManager.success("CloudZero integration created successfully");
form.resetFields();
onOk();
},
@ -39,7 +40,7 @@ export default function CloudZeroCreationModal({ open, onOk, onCancel }: CloudZe
if (error?.errorFields) {
return;
}
message.error(error?.message || "Failed to create CloudZero integration");
MessageManager.error(error?.message || "Failed to create CloudZero integration");
},
},
);
@ -47,7 +48,7 @@ export default function CloudZeroCreationModal({ open, onOk, onCancel }: CloudZe
if (error?.errorFields) {
return;
}
message.error(error?.message || "Failed to create CloudZero integration");
MessageManager.error(error?.message || "Failed to create CloudZero integration");
}
};

View file

@ -3,7 +3,8 @@ import { useCloudZeroExport } from "@/app/(dashboard)/hooks/cloudzero/useCloudZe
import { useCloudZeroDeleteSettings } from "@/app/(dashboard)/hooks/cloudzero/useCloudZeroSettings";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
import { Alert, Button, Card, Descriptions, Divider, message, Popconfirm, Tag } from "antd";
import { Alert, Button, Card, Descriptions, Divider, Popconfirm, Tag } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { CheckCircle, Edit, Play, Trash2, Upload } from "lucide-react";
import { useState } from "react";
import CloudZeroUpdateModal from "./CloudZeroUpdateModal";
@ -30,10 +31,10 @@ export function CloudZeroIntegrationSettings({ settings, onSettingsUpdated }: Cl
{ limit: 10 },
{
onSuccess: (data) => {
message.success("Dry run completed successfully");
MessageManager.success("Dry run completed successfully");
},
onError: (error) => {
message.error(error?.message || "Failed to perform dry run");
MessageManager.error(error?.message || "Failed to perform dry run");
},
},
);
@ -48,10 +49,10 @@ export function CloudZeroIntegrationSettings({ settings, onSettingsUpdated }: Cl
{ operation: "replace_hourly" },
{
onSuccess: () => {
message.success("Data successfully exported to CloudZero");
MessageManager.success("Data successfully exported to CloudZero");
},
onError: (error) => {
message.error(error?.message || "Failed to export data");
MessageManager.error(error?.message || "Failed to export data");
},
},
);
@ -79,12 +80,12 @@ export function CloudZeroIntegrationSettings({ settings, onSettingsUpdated }: Cl
deleteMutation.mutate(undefined, {
onSuccess: () => {
message.success("CloudZero integration deleted successfully");
MessageManager.success("CloudZero integration deleted successfully");
setIsDeleteModalOpen(false);
onSettingsUpdated();
},
onError: (error) => {
message.error(error?.message || "Failed to delete CloudZero integration");
MessageManager.error(error?.message || "Failed to delete CloudZero integration");
},
});
};

View file

@ -1,6 +1,7 @@
import { useCloudZeroUpdateSettings } from "@/app/(dashboard)/hooks/cloudzero/useCloudZeroSettings";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { Form, Input, message, Modal } from "antd";
import { Form, Input, Modal } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { useEffect } from "react";
import { CloudZeroSettings } from "./types";
@ -39,7 +40,7 @@ export default function CloudZeroUpdateModal({ open, onOk, onCancel, settings }:
},
{
onSuccess: () => {
message.success("CloudZero integration updated successfully");
MessageManager.success("CloudZero integration updated successfully");
form.resetFields();
onOk();
},
@ -47,7 +48,7 @@ export default function CloudZeroUpdateModal({ open, onOk, onCancel, settings }:
if (error?.errorFields) {
return;
}
message.error(error?.message || "Failed to update CloudZero integration");
MessageManager.error(error?.message || "Failed to update CloudZero integration");
},
},
);
@ -55,7 +56,7 @@ export default function CloudZeroUpdateModal({ open, onOk, onCancel, settings }:
if (error?.errorFields) {
return;
}
message.error(error?.message || "Failed to update CloudZero integration");
MessageManager.error(error?.message || "Failed to update CloudZero integration");
}
};

View file

@ -0,0 +1,38 @@
"use client";
import React from "react";
import { Select } from "antd";
import { CloudServerOutlined } from "@ant-design/icons";
import { useWorker } from "@/hooks/useWorker";
interface WorkerDropdownProps {
onWorkerSwitch: (workerId: string) => void;
}
const WorkerDropdown: React.FC<WorkerDropdownProps> = ({ onWorkerSwitch }) => {
const { isControlPlane, selectedWorker, workers } = useWorker();
if (!isControlPlane || !selectedWorker) return null;
return (
<Select
showSearch
filterOption={(input, option) =>
(option?.label as string ?? "").toLowerCase().includes(input.toLowerCase())
}
value={selectedWorker.worker_id}
style={{ minWidth: 180 }}
suffixIcon={<CloudServerOutlined />}
options={workers.map((w) => ({
label: w.name,
value: w.worker_id,
disabled: w.worker_id === selectedWorker.worker_id,
}))}
onChange={(newWorkerId) => {
onWorkerSwitch(newWorkerId);
}}
/>
);
};
export default WorkerDropdown;

View file

@ -18,8 +18,8 @@ vi.mock("./networking", () => ({
getPoliciesList: vi.fn().mockResolvedValue({ policies: [] }),
}));
vi.mock("./common_components/fetch_teams", () => ({
fetchTeams: vi.fn(),
vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({
teamListCall: vi.fn().mockResolvedValue({ teams: [], total: 0, page: 1, page_size: 100, total_pages: 0 }),
}));
vi.mock("./molecules/notifications_manager", () => ({
@ -375,6 +375,9 @@ describe("OldTeams - handleCreate organization handling", () => {
organizations={[]}
/>,
);
await waitFor(() => {
expect(screen.getByTestId("delete-team-button")).toBeInTheDocument();
});
const deleteTeamButton = screen.getByTestId("delete-team-button");
act(() => {
fireEvent.click(deleteTeamButton);
@ -389,7 +392,7 @@ describe("OldTeams - empty state", () => {
mockUseOrganizations.mockReturnValue({ data: [] });
});
it("should display empty state message when teams array is empty", () => {
it("should display empty state message when teams array is empty", async () => {
renderWithQueryClient(
<OldTeams
teams={[]}
@ -402,11 +405,13 @@ describe("OldTeams - empty state", () => {
/>,
);
expect(screen.getByText("No teams found")).toBeInTheDocument();
expect(screen.getByText("Adjust your filters or create a new team")).toBeInTheDocument();
await waitFor(() => {
expect(screen.getByText("No teams yet")).toBeInTheDocument();
});
expect(screen.getByText("Create your first team to organize members and manage access to models.")).toBeInTheDocument();
});
it("should display empty state message when teams is null", () => {
it("should display empty state message when teams is null", async () => {
renderWithQueryClient(
<OldTeams
teams={null}
@ -419,11 +424,13 @@ describe("OldTeams - empty state", () => {
/>,
);
expect(screen.getByText("No teams found")).toBeInTheDocument();
expect(screen.getByText("Adjust your filters or create a new team")).toBeInTheDocument();
await waitFor(() => {
expect(screen.getByText("No teams yet")).toBeInTheDocument();
});
expect(screen.getByText("Create your first team to organize members and manage access to models.")).toBeInTheDocument();
});
it("should not display empty state when teams array has items", () => {
it("should not display empty state when teams array has items", async () => {
renderWithQueryClient(
<OldTeams
teams={[
@ -451,9 +458,11 @@ describe("OldTeams - empty state", () => {
/>,
);
expect(screen.queryByText("No teams found")).not.toBeInTheDocument();
expect(screen.queryByText("Adjust your filters or create a new team")).not.toBeInTheDocument();
expect(screen.getByText("Test Team")).toBeInTheDocument();
await waitFor(() => {
expect(screen.getByText("Test Team")).toBeInTheDocument();
});
expect(screen.queryByText("No teams yet")).not.toBeInTheDocument();
expect(screen.queryByText("Create your first team to organize members and manage access to models.")).not.toBeInTheDocument();
});
});
@ -621,12 +630,9 @@ describe("OldTeams - premium props", () => {
/>,
);
const truncatedTeamId = "team-123456789".slice(0, 7);
const teamButton = await screen.findByRole("button", {
name: new RegExp(`${truncatedTeamId}\\.\\.\\.`),
});
const teamIdElement = await screen.findByText("team-123456789");
act(() => {
fireEvent.click(teamButton);
fireEvent.click(teamIdElement);
});
await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled());
@ -798,7 +804,7 @@ describe("OldTeams - access_group_ids in team create", () => {
/>,
);
const createButton = screen.getByRole("button", { name: /create new team/i });
const createButton = screen.getAllByRole("button", { name: /create team/i })[0];
act(() => {
fireEvent.click(createButton);
});
@ -823,7 +829,8 @@ describe("OldTeams - access_group_ids in team create", () => {
const accessGroupInput = screen.getByTestId("access-group-selector");
fireEvent.change(accessGroupInput, { target: { value: "ag-1,ag-2" } });
const createTeamSubmitButton = screen.getByRole("button", { name: /create team/i });
const createTeamSubmitButtons = screen.getAllByRole("button", { name: /create team/i });
const createTeamSubmitButton = createTeamSubmitButtons[createTeamSubmitButtons.length - 1];
fireEvent.click(createTeamSubmitButton);
await waitFor(() => {
@ -865,7 +872,7 @@ describe("OldTeams - models dropdown options", () => {
expect(fetchAvailableModelsForTeamOrKey).toHaveBeenCalled();
});
const createButton = screen.getByRole("button", { name: /create new team/i });
const createButton = screen.getAllByRole("button", { name: /create team/i })[0];
act(() => {
fireEvent.click(createButton);
});
@ -884,7 +891,7 @@ describe("OldTeams - organization alias display", () => {
mockUseOrganizations.mockReturnValue({ data: [] });
});
it("should display organization alias instead of organization id", () => {
it("should display organization alias instead of organization id", async () => {
const mockOrganizations = [
{
organization_id: "org-123",
@ -934,11 +941,13 @@ describe("OldTeams - organization alias display", () => {
/>,
);
expect(screen.getByText("Test Organization")).toBeInTheDocument();
await waitFor(() => {
expect(screen.getByText("Test Organization")).toBeInTheDocument();
});
expect(screen.queryByText("org-123")).not.toBeInTheDocument();
});
it("should display organization id when alias is not found", () => {
it("should display organization id when alias is not found", async () => {
mockUseOrganizations.mockReturnValue({ data: [] });
renderWithQueryClient(
@ -968,10 +977,12 @@ describe("OldTeams - organization alias display", () => {
/>,
);
expect(screen.getByText("org-unknown")).toBeInTheDocument();
await waitFor(() => {
expect(screen.getByText("org-unknown")).toBeInTheDocument();
});
});
it("should display N/A when organization_id is null", () => {
it("should display N/A when organization_id is null", async () => {
mockUseOrganizations.mockReturnValue({ data: [] });
renderWithQueryClient(
@ -1001,6 +1012,9 @@ describe("OldTeams - organization alias display", () => {
/>,
);
expect(screen.getByText("N/A")).toBeInTheDocument();
await waitFor(() => {
// When organization_id is null, the table shows "—" in the Organization column
expect(screen.getAllByText("—").length).toBeGreaterThan(0);
});
});
});

File diff suppressed because it is too large Load diff

View file

@ -1,5 +1,6 @@
import { Modal, Form, Button, Typography, message } from "antd";
import { Modal, Form, Button, Typography } from "antd";
import { FolderAddOutlined } from "@ant-design/icons";
import MessageManager from "@/components/molecules/message_manager";
import {
useCreateProject,
ProjectCreateParams,
@ -32,12 +33,12 @@ export function CreateProjectModal({
createMutation.mutate(params, {
onSuccess: () => {
message.success("Project created successfully");
MessageManager.success("Project created successfully");
form.resetFields();
onClose();
},
onError: (error) => {
message.error(error.message || "Failed to create project");
MessageManager.error(error.message || "Failed to create project");
},
});
} catch (error) {

View file

@ -1,6 +1,7 @@
import { useEffect } from "react";
import { Modal, Form, Button, Typography, message } from "antd";
import { Modal, Form, Button, Typography } from "antd";
import { SaveOutlined } from "@ant-design/icons";
import MessageManager from "@/components/molecules/message_manager";
import { ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects";
import {
useUpdateProject,
@ -80,12 +81,12 @@ export function EditProjectModal({
{ projectId: project.project_id, params },
{
onSuccess: () => {
message.success("Project updated successfully");
MessageManager.success("Project updated successfully");
onSuccess?.();
onClose();
},
onError: (error) => {
message.error(error.message || "Failed to update project");
MessageManager.error(error.message || "Failed to update project");
},
},
);

View file

@ -1,5 +1,6 @@
import React, { useState } from "react";
import { Button, Input, Typography, Spin, message } from "antd";
import { Button, Input, Typography, Spin } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { SearchOutlined, LoadingOutlined } from "@ant-design/icons";
import { searchToolQueryCall } from "../networking";
import NotificationsManager from "../molecules/notifications_manager";
@ -39,7 +40,7 @@ export const SearchToolTester: React.FC<SearchToolTesterProps> = ({ searchToolNa
const handleSearch = async () => {
if (!query.trim()) {
message.warning("Please enter a search query");
MessageManager.warning("Please enter a search query");
return;
}

View file

@ -5,8 +5,9 @@
*/
import { Button as TremorButton } from "@tremor/react";
import { Button, message } from "antd";
import { Button } from "antd";
import React, { useEffect, useState } from "react";
import MessageManager from "@/components/molecules/message_manager";
import NotificationManager from "../../../molecules/notifications_manager";
import { fetchAvailableModels, ModelGroup } from "../../../playground/llm_calls/fetch_models";
import { AddFallbacksModal } from "./AddFallbacksModal";
@ -90,7 +91,7 @@ export default function AddFallbacks({
(g) => !g.primaryModel || g.fallbackModels.length === 0,
);
if (invalidGroups.length > 0) {
message.error(
MessageManager.error(
`Please complete configuration for all groups. ${invalidGroups.length} group(s) incomplete.`,
);
return;

View file

@ -5,9 +5,10 @@
*/
import { Button } from "@tremor/react";
import { message, Tabs } from "antd";
import { Tabs } from "antd";
import { Plus } from "lucide-react";
import React, { useEffect, useState } from "react";
import MessageManager from "@/components/molecules/message_manager";
import { FallbackGroup, FallbackGroupConfig } from "./FallbackGroupConfig";
interface FallbackSelectionFormProps {
@ -60,7 +61,7 @@ export function FallbackSelectionForm({
const handleRemoveGroup = (targetId: string) => {
if (groups.length === 1) {
message.warning("At least one group is required");
MessageManager.warning("At least one group is required");
return;
}
const newGroups = groups.filter((g) => g.id !== targetId);

View file

@ -1,5 +1,6 @@
import React, { useState, useEffect } from "react";
import { Modal, Form, message, Select, Input, Steps, Radio, Tag, Divider, Switch, InputNumber, Collapse } from "antd";
import { Modal, Form, Select, Input, Steps, Radio, Tag, Divider, Switch, InputNumber, Collapse } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { Button } from "@tremor/react";
import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined, InfoCircleOutlined } from "@ant-design/icons";
import CreatedKeyDisplay from "../shared/CreatedKeyDisplay";
@ -216,7 +217,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
const handleCreateAgent = async () => {
if (!accessToken) {
message.error("No access token available");
MessageManager.error("No access token available");
return;
}
@ -226,7 +227,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
const values = { ...form.getFieldsValue(true) };
const agentData = buildAgentData(values);
if (!agentData) {
message.error("Failed to build agent data");
MessageManager.error("Failed to build agent data");
setIsSubmitting(false);
return;
}
@ -301,7 +302,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
setCreatedKeyValue(keyResponse.key || null);
} else if (keyAssignOption === "existing_key") {
if (!selectedExistingKey) {
message.error("Please select an existing key to assign");
MessageManager.error("Please select an existing key to assign");
setIsSubmitting(false);
return;
}
@ -318,7 +319,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
} catch (error) {
console.error("Error creating agent:", error);
const errorMessage = error instanceof Error ? error.message : String(error);
message.error(errorMessage ? `Failed to create agent: ${errorMessage}` : "Failed to create agent");
MessageManager.error(errorMessage ? `Failed to create agent: ${errorMessage}` : "Failed to create agent");
} finally {
setIsSubmitting(false);
}

View file

@ -1,6 +1,7 @@
import React, { useState, useEffect } from "react";
import { Card, Title, Text, Button as TremorButton, Tab, TabGroup, TabList, TabPanel, TabPanels} from "@tremor/react";
import { Form, Input, InputNumber, Button as AntButton, message, Spin, Descriptions, Divider } from "antd";
import { Form, Input, InputNumber, Button as AntButton, Spin, Descriptions, Divider } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { ArrowLeftIcon } from "@heroicons/react/outline";
import { getAgentInfo, patchAgentCall, getAgentCreateMetadata, AgentCreateInfo } from "../networking";
import { Agent } from "./types";
@ -72,7 +73,7 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({
}
} catch (error) {
console.error("Error fetching agent info:", error);
message.error("Failed to load agent information");
MessageManager.error("Failed to load agent information");
} finally {
setIsLoading(false);
}
@ -111,12 +112,12 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({
}
await patchAgentCall(accessToken, agentId, updateData);
message.success("Agent updated successfully");
MessageManager.success("Agent updated successfully");
setIsEditing(false);
fetchAgentInfo();
} catch (error) {
console.error("Error updating agent:", error);
message.error("Failed to update agent");
MessageManager.error("Failed to update agent");
} finally {
setIsSaving(false);
}

View file

@ -1,7 +1,8 @@
"use client";
import React, { useCallback, useEffect, useRef, useState, useLayoutEffect } from "react";
import { Tooltip, Skeleton, Popover, message } from "antd";
import { Tooltip, Skeleton, Popover } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import {
SettingOutlined,
PlusOutlined,
@ -212,7 +213,7 @@ const ChatPage: React.FC<ChatPageProps> = ({ accessToken, userRole, userId, user
localStorage.setItem(LOCALSTORAGE_MODEL_KEY, JSON.stringify([names[0]]));
}
})
.catch(() => message.error("Could not load models"))
.catch(() => MessageManager.error("Could not load models"))
.finally(() => setIsLoadingModels(false));
}, [accessToken]);

View file

@ -5,7 +5,7 @@ import { Spin, Input, Button, Skeleton } from "antd";
import { SearchOutlined, ArrowLeftOutlined, RightOutlined, ToolOutlined, CheckCircleOutlined } from "@ant-design/icons";
import { deleteMCPOAuthUserCredential, fetchMCPServers, getMCPOAuthUserCredentialStatus, listMCPTools } from "../networking";
import { AUTH_TYPE, MCPServer, MCPTool, handleTransport } from "../mcp_tools/types";
import { message } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { useUserMcpOAuthFlow } from "@/hooks/useUserMcpOAuthFlow";
// ── OAuth2 connect button ─────────────────────────────────────────────────────
@ -198,7 +198,7 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange
const idToFetch = serverId ?? serverName;
const result = await listMCPTools(accessToken, idToFetch);
if (result?.error) {
message.warning(`Could not load tools for ${serverName}`);
MessageManager.warning(`Could not load tools for ${serverName}`);
return;
}
// Use the ref so we read the most up-to-date list; guard against duplicates
@ -207,7 +207,7 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange
onChange([...selectedServersRef.current, serverName]);
}
} catch {
message.warning(`Could not load tools for ${serverName}`);
MessageManager.warning(`Could not load tools for ${serverName}`);
} finally {
setTogglingOn((prev) => {
const next = new Set(prev);

View file

@ -1,5 +1,6 @@
import React, { useEffect, useState } from "react";
import { Switch, Spin, message } from "antd";
import { Switch, Spin } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { fetchMCPServers, listMCPTools } from "../networking";
import { MCPServer } from "../mcp_tools/types";
@ -57,7 +58,7 @@ const MCPConnectPicker: React.FC<Props> = ({ accessToken, selectedServers, onCha
const result = await listMCPTools(accessToken, serverName);
// listMCPTools never throws; it returns { tools, error, message } on failure
if (result?.error) {
message.warning(
MessageManager.warning(
`Could not load tools for ${serverName} — it will be excluded from this message.`
);
// Do not add to selectedServers
@ -65,7 +66,7 @@ const MCPConnectPicker: React.FC<Props> = ({ accessToken, selectedServers, onCha
}
onChange([...selectedServers, serverName]);
} catch {
message.warning(
MessageManager.warning(
`Could not load tools for ${serverName} — it will be excluded from this message.`
);
// Do not add to selectedServers

View file

@ -8,7 +8,8 @@
*/
import React, { useCallback, useEffect, useState } from "react";
import { Spin, message } from "antd";
import { Spin } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { DeleteOutlined, LinkOutlined } from "@ant-design/icons";
import { Badge, Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react";
import {
@ -77,7 +78,7 @@ const MCPCredentialsTab: React.FC<Props> = ({ accessToken }) => {
await deleteMCPOAuthUserCredential(accessToken, serverId);
setCredentials((prev) => prev.filter((c) => c.server_id !== serverId));
} catch {
message.error("Failed to revoke connection. Please try again.");
MessageManager.error("Failed to revoke connection. Please try again.");
} finally {
setRevoking((prev) => { const n = new Set(prev); n.delete(serverId); return n; });
}

View file

@ -1,5 +1,6 @@
import React, { useState } from "react";
import { Modal, Form, Input, Select, message } from "antd";
import { Modal, Form, Input, Select } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import { Button } from "@tremor/react";
import { registerClaudeCodePlugin } from "../networking";
import {
@ -43,13 +44,13 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
const handleSubmit = async (values: any) => {
if (!accessToken) {
message.error("No access token available");
MessageManager.error("No access token available");
return;
}
// Validate plugin name
if (!validatePluginName(values.name)) {
message.error(
MessageManager.error(
"Plugin name must be kebab-case (lowercase letters, numbers, and hyphens only)"
);
return;
@ -57,7 +58,7 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
// Validate semantic version if provided
if (values.version && !isValidSemanticVersion(values.version)) {
message.error(
MessageManager.error(
"Version must be in semantic versioning format (e.g., 1.0.0)"
);
return;
@ -65,13 +66,13 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
// Validate email if provided
if (values.authorEmail && !isValidEmail(values.authorEmail)) {
message.error("Invalid email format");
MessageManager.error("Invalid email format");
return;
}
// Validate homepage URL if provided
if (values.homepage && !isValidUrl(values.homepage)) {
message.error("Invalid homepage URL format");
MessageManager.error("Invalid homepage URL format");
return;
}
@ -119,14 +120,14 @@ const AddPluginForm: React.FC<AddPluginFormProps> = ({
}
await registerClaudeCodePlugin(accessToken, pluginData);
message.success("Plugin registered successfully");
MessageManager.success("Plugin registered successfully");
form.resetFields();
setSourceType("github");
onSuccess();
onClose();
} catch (error) {
console.error("Error registering plugin:", error);
message.error("Failed to register plugin");
MessageManager.error("Failed to register plugin");
} finally {
setIsSubmitting(false);
}

View file

@ -6,6 +6,7 @@ import {
ChevronUpIcon,
ChevronDownIcon,
ExternalLinkIcon,
ClipboardCopyIcon,
} from "@heroicons/react/outline";
import { Tooltip } from "antd";
import BaseActionButton from "../BaseActionButton";
@ -32,6 +33,7 @@ export const TableIconActionButtonMap: Record<string, TableIconActionButtonBaseP
Up: { icon: ChevronUpIcon, className: "hover:text-blue-600" },
Down: { icon: ChevronDownIcon, className: "hover:text-blue-600" },
Open: { icon: ExternalLinkIcon, className: "hover:text-green-600" },
Copy: { icon: ClipboardCopyIcon, className: "hover:text-blue-600" },
};
export default function TableIconActionButton({

View file

@ -1,13 +1,16 @@
import React from "react";
import { Select } from "antd";
import { Select, Typography } from "antd";
import { Organization } from "../networking";
const { Text } = Typography;
interface OrganizationDropdownProps {
organizations?: Organization[] | null;
value?: string;
onChange?: (value: string) => void;
disabled?: boolean;
loading?: boolean;
style?: React.CSSProperties;
}
const OrganizationDropdown: React.FC<OrganizationDropdownProps> = ({
@ -16,16 +19,18 @@ const OrganizationDropdown: React.FC<OrganizationDropdownProps> = ({
onChange,
disabled,
loading,
style,
}) => {
return (
<Select
showSearch
placeholder="Search or select an organization"
placeholder="All Organizations"
value={value}
onChange={onChange}
disabled={disabled}
loading={loading}
allowClear
style={{ minWidth: 280, ...style }}
filterOption={(input, option) => {
if (!option) return false;
const org = organizations?.find((o) => o.organization_id === option.key);
@ -37,12 +42,11 @@ const OrganizationDropdown: React.FC<OrganizationDropdownProps> = ({
return orgAlias.includes(searchTerm) || orgId.includes(searchTerm);
}}
optionFilterProp="children"
>
{organizations?.map((org) => (
<Select.Option key={org.organization_id} value={org.organization_id}>
<span className="font-medium">{org.organization_alias}</span>{" "}
<span className="text-gray-500">({org.organization_id})</span>
<Text type="secondary">({org.organization_id})</Text>
</Select.Option>
))}
</Select>

View file

@ -0,0 +1,90 @@
import { describe, expect, it, vi } from "vitest";
import { fetchTeamFilterOptions } from "./filter_helpers";
const mockKeyListCall = vi.fn();
vi.mock("@/components/networking", () => ({
keyListCall: (...args: unknown[]) => mockKeyListCall(...args),
teamListCall: vi.fn(),
organizationListCall: vi.fn(),
}));
describe("fetchTeamFilterOptions", () => {
it("should return empty arrays when accessToken is null", async () => {
const result = await fetchTeamFilterOptions(null, "team-1");
expect(result).toEqual({ keyAliases: [], organizationIds: [], userIds: [] });
expect(mockKeyListCall).not.toHaveBeenCalled();
});
it("should return empty arrays when teamId is empty", async () => {
const result = await fetchTeamFilterOptions("tok-123", "");
expect(result).toEqual({ keyAliases: [], organizationIds: [], userIds: [] });
expect(mockKeyListCall).not.toHaveBeenCalled();
});
it("should return sorted key aliases from fetched keys", async () => {
mockKeyListCall.mockResolvedValue({
keys: [
{ key_alias: "zeta-key" },
{ key_alias: "alpha-key" },
{ key_alias: "mid-key" },
],
total_pages: 1,
});
const result = await fetchTeamFilterOptions("tok-123", "team-1");
expect(result.keyAliases).toEqual(["alpha-key", "mid-key", "zeta-key"]);
});
it("should deduplicate organization IDs across pages", async () => {
mockKeyListCall
.mockResolvedValueOnce({
keys: [
{ organization_id: "org-b" },
{ organization_id: "org-a" },
],
total_pages: 2,
})
.mockResolvedValueOnce({
keys: [
{ organization_id: "org-a" },
{ organization_id: "org-c" },
],
total_pages: 2,
});
const result = await fetchTeamFilterOptions("tok-123", "team-1");
expect(result.organizationIds).toEqual(["org-a", "org-b", "org-c"]);
});
it("should map user IDs with email addresses", async () => {
mockKeyListCall.mockResolvedValue({
keys: [
{ user_id: "u1", user: { user_email: "alice@example.com" } },
{ user_id: "u2", user: { user_email: "bob@example.com" } },
],
total_pages: 1,
});
const result = await fetchTeamFilterOptions("tok-123", "team-1");
expect(result.userIds).toEqual(
expect.arrayContaining([
{ id: "u1", email: "alice@example.com" },
{ id: "u2", email: "bob@example.com" },
]),
);
});
it("should handle API errors gracefully and return empty arrays", async () => {
mockKeyListCall.mockRejectedValue(new Error("Network error"));
const result = await fetchTeamFilterOptions("tok-123", "team-1");
expect(result).toEqual({ keyAliases: [], organizationIds: [], userIds: [] });
});
});

View file

@ -0,0 +1,62 @@
import { describe, it, expect } from "vitest";
import { transformKeyInfo } from "./transform_key_info";
describe("transformKeyInfo", () => {
it("should combine key and info fields into a single object", () => {
const apiResponse = {
key: "sk-abc123",
info: {
token_id: "tok_1",
key_name: "my-key",
spend: 10.5,
},
};
const result = transformKeyInfo(apiResponse);
expect(result).toEqual({
token: "sk-abc123",
token_id: "tok_1",
key_name: "my-key",
spend: 10.5,
});
});
it("should set the token field from the key property", () => {
const apiResponse = {
key: "sk-xyz789",
info: { key_name: "test" },
};
const result = transformKeyInfo(apiResponse);
expect(result.token).toBe("sk-xyz789");
});
it("should preserve all info fields in the result", () => {
const apiResponse = {
key: "sk-abc",
info: {
token_id: "tok_2",
key_name: "prod-key",
spend: 42,
models: ["gpt-4"],
team_id: "team-1",
metadata: { env: "production" },
},
};
const result = transformKeyInfo(apiResponse);
expect(result.token_id).toBe("tok_2");
expect(result.key_name).toBe("prod-key");
expect(result.spend).toBe(42);
expect(result.models).toEqual(["gpt-4"]);
expect(result.team_id).toBe("team-1");
expect(result.metadata).toEqual({ env: "production" });
});
it("should handle empty info object", () => {
const apiResponse = {
key: "sk-empty",
info: {},
};
const result = transformKeyInfo(apiResponse);
expect(result.token).toBe("sk-empty");
expect(Object.keys(result)).toContain("token");
});
});

View file

@ -232,8 +232,8 @@ const menuGroups: MenuGroup[] = [
groupLabel: "DEVELOPER TOOLS",
items: [
{
key: "api_ref",
page: "api_ref",
key: "api-reference",
page: "api-reference",
label: "API Reference",
icon: <ApiOutlined />,
},

View file

@ -1,7 +1,8 @@
"use client";
import React, { useState } from "react";
import { Modal, Input, Switch, message } from "antd";
import { Modal, Input, Switch } from "antd";
import MessageManager from "@/components/molecules/message_manager";
import {
KeyOutlined,
LockOutlined,
@ -46,7 +47,7 @@ export const ByokCredentialModal: React.FC<ByokCredentialModalProps> = ({
const handleAuthorize = async () => {
if (!apiKey.trim()) {
message.error("Please enter your API key");
MessageManager.error("Please enter your API key");
return;
}
setLoading(true);
@ -63,11 +64,11 @@ export const ByokCredentialModal: React.FC<ByokCredentialModalProps> = ({
const err = await response.json();
throw new Error(err?.detail?.error || "Failed to save credential");
}
message.success(`Connected to ${serverDisplayName}`);
MessageManager.success(`Connected to ${serverDisplayName}`);
onSuccess(server.server_id);
handleClose();
} catch (e: any) {
message.error(e.message || "Failed to connect");
MessageManager.error(e.message || "Failed to connect");
} finally {
setLoading(false);
}

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