mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'main' into docs/prompt-caching-gemini-support
This commit is contained in:
commit
c8a7d5d237
128 changed files with 7355 additions and 1222 deletions
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 [],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
14
litellm/types/proxy/control_plane_endpoints.py
Normal file
14
litellm/types/proxy/control_plane_endpoints.py
Normal 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
|
||||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
8
poetry.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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>"
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"] == []
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
167
tests/test_litellm/proxy/test_model_info_default_limits.py
Normal file
167
tests/test_litellm/proxy/test_model_info_default_limits.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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} />;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
},
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
>
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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" ? (
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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)`);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
},
|
||||
},
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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]);
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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; });
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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: [] });
|
||||
});
|
||||
});
|
||||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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 />,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue