Merge branch 'main' into ttl-prompt-caching-bedrock

This commit is contained in:
Sameer Kankute 2026-02-05 16:50:54 +05:30 • committed by GitHub
commit 7e8be5f542
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
54 changed files with 2244 additions and 378 deletions

View file

@ -215,6 +215,66 @@ The following parameters can be updated on a continuation of a trace by passing
Any other key value pairs passed into the metadata not listed in the above spec for a `litellm` completion will be added as a metadata key value pair for the generation.
#### Multiple Langfuse Projects (Per-Request Credentials)
You can send traces to different Langfuse projects per request by passing credentials directly to `completion()` or `acompletion()`. This works alongside (or instead of) the global env vars and is useful when different teams or business processes use different Langfuse projects.
Pass **`langfuse_public_key`**, **`langfuse_secret_key`** (or **`langfuse_secret`**), and optionally **`langfuse_host`** as keyword arguments:
```python
import litellm
from litellm import completion
# Optional: set a default via env for requests that don't pass credentials
# os.environ["LANGFUSE_PUBLIC_KEY"] = "pk-default..."
# os.environ["LANGFUSE_SECRET_KEY"] = "sk-default..."
litellm.success_callback = ["langfuse"]
litellm.failure_callback = ["langfuse"]
# Request 1 → Langfuse Project A
response_a = completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello from team A"}],
langfuse_public_key="pk-lf-project-a...",
langfuse_secret_key="sk-lf-project-a...",
langfuse_host="https://us.cloud.langfuse.com", # optional
)
# Request 2 → Langfuse Project B (different project)
response_b = completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello from team B"}],
langfuse_public_key="pk-lf-project-b...",
langfuse_secret_key="sk-lf-project-b...",
langfuse_host="https://eu.cloud.langfuse.com", # optional, can differ per project
)
```
Async usage with per-request credentials:
```python
import litellm
from litellm import acompletion
litellm.success_callback = ["langfuse"]
litellm.failure_callback = ["langfuse"]
response = await acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hi"}],
langfuse_public_key="pk-lf-...",
langfuse_secret_key="sk-lf-...",
langfuse_host="https://us.cloud.langfuse.com", # optional
)
```
- **`langfuse_public_key`** – Langfuse project public key (required for per-request override).
- **`langfuse_secret_key`** or **`langfuse_secret`** – Langfuse secret key (either name is accepted).
- **`langfuse_host`** – Langfuse host URL (e.g. `https://us.cloud.langfuse.com`); optional, defaults to env or Langfuse cloud.
When these are passed, that request uses this project (and host) for the Langfuse callback; when omitted, the callback uses the global Langfuse client (from env vars if set). LiteLLM caches a Langfuse client per credential set to avoid creating a new client on every request.
#### Disable Logging - Specific Calls
To disable logging for specific calls use the `no-log` flag.

View file

@ -23,26 +23,75 @@ From v1.76.0, SSO is now Free for up to 5 users.
<Tabs>
<TabItem value="okta" label="Okta SSO">
1. Add Okta credentials to your .env
#### Step 1: Create an OIDC Application in Okta
In your Okta Admin Console, create a new **OIDC Web Application**. See [Okta's guide on creating OIDC app integrations](https://help.okta.com/en-us/content/topics/apps/apps_app_integration_wizard_oidc.htm) for detailed instructions.
When configuring the application:
- **Sign-in redirect URI**: `https://<your-proxy-base-url>/sso/callback`
- **Sign-out redirect URI** (optional): `https://<your-proxy-base-url>`
<Image img={require('../../img/okta_redirect_uri.png')} />
After creating the app, copy your **Client ID** and **Client Secret** from the application's General tab:
<Image img={require('../../img/okta_client_credentials.png')} />
#### Step 2: Assign Users to the Application
Ensure users are assigned to the app in the **Assignments** tab. If Federation Broker Mode is enabled, you may need to disable it to assign users manually.
#### Step 3: Configure Authorization Server Access Policy
:::warning Important
This step is required. Without an Access Policy for your app, users will get a `no_matching_policy` error when attempting to log in.
:::
1. Go to **Security** → **API**
<Image img={require('../../img/okta_security_api.png')} />
2. Select the **default** authorization server (or your custom one)
<Image img={require('../../img/okta_authorization_server.png')} />
3. Click on **Access Policies** tab, create a new policy assigned to your LiteLLM app
4. Add a rule that allows the **Authorization Code** grant type
<Image img={require('../../img/okta_access_policies.png')} />
See [Okta's Access Policy documentation](https://help.okta.com/en-us/content/topics/security/api-access-management/access-policies.htm) for more details.
#### Step 4: Configure LiteLLM Environment Variables
```bash
GENERIC_CLIENT_ID = "<your-okta-client-id>"
GENERIC_CLIENT_SECRET = "<your-okta-client-secret>"
GENERIC_AUTHORIZATION_ENDPOINT = "<your-okta-domain>/authorize" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/authorize
GENERIC_TOKEN_ENDPOINT = "<your-okta-domain>/token" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/oauth/token
GENERIC_USERINFO_ENDPOINT = "<your-okta-domain>/userinfo" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/userinfo
GENERIC_CLIENT_STATE = "random-string" # [OPTIONAL] REQUIRED BY OKTA, if not set random state value is generated
GENERIC_SSO_HEADERS = "Content-Type=application/json, X-Custom-Header=custom-value" # [OPTIONAL] Comma-separated list of additional headers to add to the request - e.g. Content-Type=application/json, etc.
GENERIC_CLIENT_ID="<your-client-id>"
GENERIC_CLIENT_SECRET="<your-client-secret>"
GENERIC_AUTHORIZATION_ENDPOINT="https://<your-okta-domain>/oauth2/default/v1/authorize"
GENERIC_TOKEN_ENDPOINT="https://<your-okta-domain>/oauth2/default/v1/token"
GENERIC_USERINFO_ENDPOINT="https://<your-okta-domain>/oauth2/default/v1/userinfo"
GENERIC_CLIENT_STATE="random-string"
PROXY_BASE_URL="https://<your-proxy-base-url>"
```
You can get your domain specific auth/token/userinfo endpoints at `<YOUR-OKTA-DOMAIN>/.well-known/openid-configuration`
:::tip
You can find all OAuth endpoints at `https://<your-okta-domain>/.well-known/openid-configuration`
:::
2. Add proxy url as callback_url on Okta
#### Step 5: Test the SSO Flow
On Okta, add the 'callback_url' as `<proxy_base_url>/sso/callback`
1. Start your LiteLLM proxy
2. Navigate to `https://<your-proxy-base-url>/ui`
3. Click the SSO login button
4. Authenticate with Okta and verify you're redirected back to LiteLLM
#### Troubleshooting
<Image img={require('../../img/okta_callback_url.png')} />
| Error | Cause | Solution |
|-------|-------|----------|
| `redirect_uri` error | Redirect URI not configured | Add `<proxy_base_url>/sso/callback` to Sign-in redirect URIs in Okta |
| `access_denied` | User not assigned to app | Assign the user in the Assignments tab |
| `no_matching_policy` | Missing Access Policy | Create an Access Policy in the Authorization Server (see Step 3) |
</TabItem>
<TabItem value="google" label="Google SSO">

Binary file not shown.

After

Width:  |  Height:  |  Size: 82 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 52 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 64 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 60 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 38 KiB

Binary file not shown.

View file

@ -0,0 +1,6 @@
-- AlterTable
ALTER TABLE "LiteLLM_DeletedTeamTable" ADD COLUMN "allow_team_guardrail_config" BOOLEAN NOT NULL DEFAULT false;
-- AlterTable
ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN "allow_team_guardrail_config" BOOLEAN NOT NULL DEFAULT false;

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-proxy-extras"
version = "0.4.29"
version = "0.4.30"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.4.29"
version = "0.4.30"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-proxy-extras==",

View file

@ -81,6 +81,11 @@ MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int(
os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150)
)
LITELLM_UI_ALLOW_HEADERS = [
"x-litellm-semantic-filter",
"x-litellm-semantic-filter-tools",
]
# Gemini model-specific minimal thinking budget constants
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH = int(
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH", 1)

View file

@ -114,6 +114,27 @@ class LoggingCallbackManager:
for c in remove_list:
callback_list.remove(c)
def remove_callbacks_by_type(self, callback_list, callback_type):
"""
Remove all callbacks of a specific type from a callback list.
Args:
callback_list: The list to remove callbacks from (e.g., litellm.callbacks)
callback_type: The class type to match (e.g., SemanticToolFilterHook)
Example:
litellm.logging_callback_manager.remove_callbacks_by_type(
litellm.callbacks, SemanticToolFilterHook
)
"""
if not isinstance(callback_list, list):
return
remove_list = [c for c in callback_list if isinstance(c, callback_type)]
for c in remove_list:
callback_list.remove(c)
def _add_string_callback_to_list(
self, callback: str, parent_list: List[Union[CustomLogger, Callable, str]]
):

View file

@ -939,22 +939,8 @@ class LiteLLMAnthropicMessagesAdapter:
self,
choices: List[Choices],
tool_name_mapping: Optional[Dict[str, str]] = None,
) -> List[
Union[
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockToolUse,
AnthropicResponseContentBlockThinking,
AnthropicResponseContentBlockRedactedThinking,
]
]:
new_content: List[
Union[
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockToolUse,
AnthropicResponseContentBlockThinking,
AnthropicResponseContentBlockRedactedThinking,
]
] = []
) -> List[Dict[str, Any]]:
new_content: List[Dict[str, Any]] = []
for choice in choices:
# Handle thinking blocks first
if (
@ -978,7 +964,7 @@ class LiteLLMAnthropicMessagesAdapter:
if signature_value is not None
else None
),
)
).model_dump()
)
elif thinking_block.get("type") == "redacted_thinking":
data_value = thinking_block.get("data", "")
@ -986,7 +972,7 @@ class LiteLLMAnthropicMessagesAdapter:
AnthropicResponseContentBlockRedactedThinking(
type="redacted_thinking",
data=str(data_value) if data_value is not None else "",
)
).model_dump()
)
# Handle reasoning_content when thinking_blocks is not present
elif (
@ -998,7 +984,7 @@ class LiteLLMAnthropicMessagesAdapter:
type="thinking",
thinking=str(choice.message.reasoning_content),
signature=None,
)
).model_dump()
)
# Handle text content
@ -1006,7 +992,7 @@ class LiteLLMAnthropicMessagesAdapter:
new_content.append(
AnthropicResponseContentBlockText(
type="text", text=choice.message.content
)
).model_dump()
)
# Handle tool calls (in parallel to text content)
if (
@ -1044,7 +1030,7 @@ class LiteLLMAnthropicMessagesAdapter:
tool_use_block.provider_specific_fields = (
provider_specific_fields
)
new_content.append(tool_use_block)
new_content.append(tool_use_block.model_dump())
return new_content

View file

@ -239,7 +239,7 @@ class FireworksAIConfig(OpenAIGPTConfig):
# Remove fields not permitted by FireworksAI that may cause:
# "Not permitted, field: 'messages[n].provider_specific_fields'"
if isinstance(message, dict) and "provider_specific_fields" in message:
message.pop("provider_specific_fields", None)
cast(dict, message).pop("provider_specific_fields", None)
return messages

View file

@ -1,8 +1,5 @@
<<<<<<< ttl-prompt-caching-bedrock
from typing import List, Optional, Tuple
=======
from typing import Any, List, Optional, Tuple, cast
>>>>>>> main
from litellm.exceptions import AuthenticationError
from litellm.llms.openai.openai import OpenAIConfig
@ -51,22 +48,12 @@ class GithubCopilotConfig(OpenAIConfig):
):
import litellm
<<<<<<< ttl-prompt-caching-bedrock
disable_copilot_system_to_assistant = (
litellm.disable_copilot_system_to_assistant
)
if not disable_copilot_system_to_assistant:
for message in messages:
if "role" in message and message["role"] == "system":
message["role"] = "assistant"
return messages
=======
# Check if system-to-assistant conversion is disabled
if litellm.disable_copilot_system_to_assistant:
# GitHub Copilot API now supports system prompts for all models (Claude, GPT, etc.)
# No conversion needed - just return messages as-is
return messages
# Default behavior: convert system messages to assistant for compatibility
transformed_messages = []
for message in messages:
@ -77,9 +64,8 @@ class GithubCopilotConfig(OpenAIConfig):
transformed_messages.append(transformed_message)
else:
transformed_messages.append(message)
return transformed_messages
>>>>>>> main
def validate_environment(
self,

View file

@ -24419,6 +24419,31 @@
"supports_tool_choice": true,
"supports_function_calling": true
},
"openrouter/qwen/qwen3-235b-a22b-2507": {
"input_cost_per_token": 7.1e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 1e-07,
"source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507",
"supports_function_calling": true,
"supports_tool_choice": true
},
"openrouter/qwen/qwen3-235b-a22b-thinking-2507": {
"input_cost_per_token": 1.1e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"source": "https://openrouter.ai/qwen/qwen3-235b-a22b-thinking-2507",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"openrouter/switchpoint/router": {
"input_cost_per_token": 8.5e-07,
"litellm_provider": "openrouter",

View file

@ -11,7 +11,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
async def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]:
def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]:
"""
Route A2A agent requests directly to litellm with injected API base.

View file

@ -1187,119 +1187,130 @@ class DBSpendUpdateWriter:
)
break
async with prisma_client.db.batch_() as batcher:
for _, transaction in transactions_to_process.items():
entity_id = transaction.get(entity_id_field)
try:
async with prisma_client.db.batch_() as batcher:
for _, transaction in transactions_to_process.items():
entity_id = transaction.get(entity_id_field)
# Construct the where clause dynamically
where_clause = {
unique_constraint_name: {
# Construct the where clause dynamically
where_clause = {
unique_constraint_name: {
entity_id_field: entity_id,
"date": transaction["date"],
"api_key": transaction["api_key"],
"model": transaction["model"],
"custom_llm_provider": transaction.get(
"custom_llm_provider"
)
or "",
"mcp_namespaced_tool_name": transaction.get(
"mcp_namespaced_tool_name"
)
or "",
"endpoint": transaction.get("endpoint") or "",
}
}
# Get the table dynamically
table = getattr(batcher, table_name)
# Common data structure for both create and update
common_data = {
entity_id_field: entity_id,
"date": transaction["date"],
"api_key": transaction["api_key"],
"model": transaction["model"],
"custom_llm_provider": transaction.get(
"custom_llm_provider"
)
or "",
"model": transaction.get("model"),
"model_group": transaction.get("model_group"),
"mcp_namespaced_tool_name": transaction.get(
"mcp_namespaced_tool_name"
)
or "",
"custom_llm_provider": transaction.get(
"custom_llm_provider"
),
"endpoint": transaction.get("endpoint") or "",
"prompt_tokens": transaction["prompt_tokens"],
"completion_tokens": transaction["completion_tokens"],
"spend": transaction["spend"],
"api_requests": transaction["api_requests"],
"successful_requests": transaction[
"successful_requests"
],
"failed_requests": transaction["failed_requests"],
}
}
# Get the table dynamically
table = getattr(batcher, table_name)
# Common data structure for both create and update
common_data = {
entity_id_field: entity_id,
"date": transaction["date"],
"api_key": transaction["api_key"],
"model": transaction.get("model"),
"model_group": transaction.get("model_group"),
"mcp_namespaced_tool_name": transaction.get(
"mcp_namespaced_tool_name"
)
or "",
"custom_llm_provider": transaction.get(
"custom_llm_provider"
),
"endpoint": transaction.get("endpoint"),
"prompt_tokens": transaction["prompt_tokens"],
"completion_tokens": transaction["completion_tokens"],
"spend": transaction["spend"],
"api_requests": transaction["api_requests"],
"successful_requests": transaction[
"successful_requests"
],
"failed_requests": transaction["failed_requests"],
}
# Add cache-related fields if they exist
if "cache_read_input_tokens" in transaction:
common_data["cache_read_input_tokens"] = (
transaction.get("cache_read_input_tokens", 0)
)
if "cache_creation_input_tokens" in transaction:
common_data["cache_creation_input_tokens"] = (
transaction.get("cache_creation_input_tokens", 0)
)
if entity_type == "tag" and "request_id" in transaction:
common_data["request_id"] = transaction.get(
"request_id"
)
# Create update data structure
update_data = {
"prompt_tokens": {
"increment": transaction["prompt_tokens"]
},
"completion_tokens": {
"increment": transaction["completion_tokens"]
},
"spend": {"increment": transaction["spend"]},
"api_requests": {
"increment": transaction["api_requests"]
},
"successful_requests": {
"increment": transaction["successful_requests"]
},
"failed_requests": {
"increment": transaction["failed_requests"]
},
}
# Add cache-related fields to update if they exist
if "cache_read_input_tokens" in transaction:
update_data["cache_read_input_tokens"] = {
"increment": transaction.get(
"cache_read_input_tokens", 0
# Add cache-related fields if they exist
if "cache_read_input_tokens" in transaction:
common_data["cache_read_input_tokens"] = (
transaction.get("cache_read_input_tokens", 0)
)
}
if "cache_creation_input_tokens" in transaction:
update_data["cache_creation_input_tokens"] = {
"increment": transaction.get(
"cache_creation_input_tokens", 0
if "cache_creation_input_tokens" in transaction:
common_data["cache_creation_input_tokens"] = (
transaction.get("cache_creation_input_tokens", 0)
)
if entity_type == "tag" and "request_id" in transaction:
common_data["request_id"] = transaction.get(
"request_id"
)
# Create update data structure
update_data = {
"prompt_tokens": {
"increment": transaction["prompt_tokens"]
},
"completion_tokens": {
"increment": transaction["completion_tokens"]
},
"spend": {"increment": transaction["spend"]},
"api_requests": {
"increment": transaction["api_requests"]
},
"successful_requests": {
"increment": transaction["successful_requests"]
},
"failed_requests": {
"increment": transaction["failed_requests"]
},
}
if entity_type == "tag" and "request_id" in transaction:
update_data["request_id"] = transaction.get("request_id")
# Add cache-related fields to update if they exist
if "cache_read_input_tokens" in transaction:
update_data["cache_read_input_tokens"] = {
"increment": transaction.get(
"cache_read_input_tokens", 0
)
}
if "cache_creation_input_tokens" in transaction:
update_data["cache_creation_input_tokens"] = {
"increment": transaction.get(
"cache_creation_input_tokens", 0
)
}
# Add endpoint to update_data so existing rows get their endpoint field updated
update_data["endpoint"] = transaction.get("endpoint") or ""
if entity_type == "tag" and "request_id" in transaction:
update_data["request_id"] = transaction.get("request_id")
table.upsert(
where=where_clause,
data={
"create": common_data,
"update": update_data,
},
)
# Add endpoint to update_data so existing rows get their endpoint field updated
update_data["endpoint"] = transaction.get("endpoint") or ""
table.upsert(
where=where_clause,
data={
"create": common_data,
"update": update_data,
},
)
except Exception as batch_error:
# Log detailed error information for debugging batch upsert failures
# This helps diagnose issues like unique constraint violations
verbose_proxy_logger.exception(
f"Daily {entity_type} spend batch upsert failed. "
f"Table: {table_name}, Constraint: {unique_constraint_name}, "
f"Batch size: {len(transactions_to_process)}, "
f"Error: {str(batch_error)}"
)
raise
verbose_proxy_logger.debug(
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"

View file

@ -47,6 +47,7 @@ from litellm.constants import (
DEFAULT_SLACK_ALERTING_THRESHOLD,
LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS,
LITELLM_SETTINGS_SAFE_DB_OVERRIDES,
LITELLM_UI_ALLOW_HEADERS,
)
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
@ -1214,6 +1215,7 @@ app.add_middleware(
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
expose_headers=LITELLM_UI_ALLOW_HEADERS,
)
app.add_middleware(PrometheusAuthMiddleware)
@ -1862,6 +1864,7 @@ class ProxyConfig:
def __init__(self) -> None:
self.config: Dict[str, Any] = {}
self._last_semantic_filter_config: Optional[Dict[str, Any]] = None
def is_yaml(self, config_file_path: str) -> bool:
if not os.path.isfile(config_file_path):
@ -3920,6 +3923,93 @@ class ProxyConfig:
prisma_client=prisma_client, proxy_config=self
)
if self._should_load_db_object(object_type="semantic_filter_settings"):
await self._init_semantic_filter_settings_in_db(
prisma_client=prisma_client
)
async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient):
"""
Initialize MCP semantic filter settings from database.
Called periodically (approximately every 10 seconds) by background task to hot-reload settings across all pods.
"""
import json
import litellm
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
try:
# Load litellm_settings from DB
config_record = await prisma_client.db.litellm_config.find_unique(
where={"param_name": "litellm_settings"}
)
if config_record is None or config_record.param_value is None:
return
litellm_settings = config_record.param_value
if isinstance(litellm_settings, str):
litellm_settings = json.loads(litellm_settings)
mcp_semantic_filter_config = litellm_settings.get(
"mcp_semantic_tool_filter", None
)
if mcp_semantic_filter_config is None:
return
# Check if settings have changed (compare with in-memory state)
if hasattr(self, "_last_semantic_filter_config"):
if self._last_semantic_filter_config == mcp_semantic_filter_config:
# If hook is missing or router isn't built yet, reinitialize anyway
active_hooks = (
litellm.logging_callback_manager.get_custom_loggers_for_type(
SemanticToolFilterHook
)
)
if active_hooks:
for active_hook in active_hooks:
if isinstance(active_hook, SemanticToolFilterHook):
if (
active_hook.filter is not None
and active_hook.filter.tool_router is not None
):
verbose_proxy_logger.debug(
"Semantic filter settings unchanged, skipping reinitialization"
)
return
verbose_proxy_logger.info(
"Semantic filter settings unchanged, but hook is missing or uninitialized. Reinitializing."
)
# Remove old hooks using logging callback manager
litellm.logging_callback_manager.remove_callbacks_by_type(
litellm.callbacks, SemanticToolFilterHook
)
# Initialize new hook if enabled
if mcp_semantic_filter_config.get("enabled", False):
global llm_router
hook = await SemanticToolFilterHook.initialize_from_config(
config=mcp_semantic_filter_config,
llm_router=llm_router,
)
if hook:
litellm.logging_callback_manager.add_litellm_callback(hook)
verbose_proxy_logger.info(
"MCP Semantic Filter reinitialized from DB"
)
else:
verbose_proxy_logger.info("MCP Semantic Filter disabled")
# Store current config for comparison next time
self._last_semantic_filter_config = mcp_semantic_filter_config.copy()
except Exception as e:
verbose_proxy_logger.exception(
f"Error initializing semantic filter settings from DB: {e}"
)
async def _init_sso_settings_in_db(self, prisma_client: PrismaClient):
"""
Initialize SSO settings from database into the router on startup.

View file

@ -332,7 +332,10 @@ async def route_request(
route_a2a_agent_request,
)
return await route_a2a_agent_request(data, route_type)
result = route_a2a_agent_request(data, route_type)
if result is not None:
return result
# Fall through to raise exception below if result is None
elif user_model is not None:
return getattr(litellm, f"{route_type}")(**data)

View file

@ -98,6 +98,40 @@ ALLOWED_UI_SETTINGS_FIELDS = {
}
class MCPSemanticFilterSettings(BaseModel):
"""Configuration for MCP Semantic Tool Filter"""
enabled: bool = Field(
default=False,
description="Enable semantic filtering of MCP tools based on query relevance",
)
embedding_model: str = Field(
default="text-embedding-3-small",
description="Embedding model to use for semantic similarity (e.g., 'text-embedding-3-small', 'text-embedding-ada-002')",
)
top_k: int = Field(
default=10,
description="Number of most relevant tools to return",
ge=1,
le=100,
)
similarity_threshold: float = Field(
default=0.3,
description="Minimum similarity score for tool inclusion (0.0 to 1.0, where 1.0 = exact match)",
ge=0.0,
le=1.0,
)
class MCPSemanticFilterSettingsResponse(SettingsResponse):
"""Response model for MCP semantic filter settings"""
pass
@router.get(
"/get/allowed_ips",
tags=["Budget & Spend Tracking"],
@ -325,7 +359,7 @@ async def update_default_team_member_budget(
async def _update_litellm_setting(
settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams],
settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings],
settings_key: str,
in_memory_var: Any,
success_message: str,
@ -769,6 +803,70 @@ async def update_ui_theme_settings(theme_config: UIThemeConfig):
}
@router.get(
"/get/mcp_semantic_filter_settings",
tags=["Settings"],
dependencies=[Depends(user_api_key_auth)],
response_model=MCPSemanticFilterSettingsResponse,
)
async def get_mcp_semantic_filter_settings(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Get MCP semantic filter configuration.
Returns current settings for semantic tool filtering.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_config
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": "Database not connected. Please connect a database."},
)
config = await proxy_config.get_config()
return await _get_settings_with_schema(
settings_key="mcp_semantic_tool_filter",
settings_class=MCPSemanticFilterSettings,
config=config,
)
@router.patch(
"/update/mcp_semantic_filter_settings",
tags=["Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_mcp_semantic_filter_settings(
settings: MCPSemanticFilterSettings,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Update MCP semantic filter settings in database.
Settings will be picked up by all pods within approximately 10 seconds via background polling.
"""
result = await _update_litellm_setting(
settings=settings,
settings_key="mcp_semantic_tool_filter",
in_memory_var=None,
success_message="MCP Semantic Filter settings updated successfully. Changes will be applied across all pods within 10 seconds.",
)
try:
from litellm.proxy.proxy_server import prisma_client, proxy_config
if prisma_client is not None:
await proxy_config._init_semantic_filter_settings_in_db(
prisma_client=prisma_client
)
except Exception as e:
verbose_proxy_logger.warning(
f"Failed to reinitialize MCP semantic filter settings immediately: {e}"
)
return result
@router.get(
"/in_product_nudges",
tags=["UI Settings"],

View file

@ -79,7 +79,7 @@ class AzureADCredential:
credential: An azure-identity credential object. If None,
DefaultAzureCredential will be used on first token request.
"""
self._credential = credential
self._credential: Any = credential
self._initialized = credential is not None
def get_token(self, scope: str) -> AccessToken:

View file

@ -24419,6 +24419,31 @@
"supports_tool_choice": true,
"supports_function_calling": true
},
"openrouter/qwen/qwen3-235b-a22b-2507": {
"input_cost_per_token": 7.1e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 1e-07,
"source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507",
"supports_function_calling": true,
"supports_tool_choice": true
},
"openrouter/qwen/qwen3-235b-a22b-thinking-2507": {
"input_cost_per_token": 1.1e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6e-07,
"source": "https://openrouter.ai/qwen/qwen3-235b-a22b-thinking-2507",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"openrouter/switchpoint/router": {
"input_cost_per_token": 8.5e-07,
"litellm_provider": "openrouter",

52
poetry.lock generated
View file

@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand.
# This file is automatically @generated by Poetry 2.3.2 and should not be changed by hand.
[[package]]
name = "a2a-sdk"
@ -398,6 +398,7 @@ files = [
{file = "azure_core-1.36.0-py3-none-any.whl", hash = "sha256:fee9923a3a753e94a259563429f3644aaf05c486d45b1215d098115102d91d3b"},
{file = "azure_core-1.36.0.tar.gz", hash = "sha256:22e5605e6d0bf1d229726af56d9e92bc37b6e726b141a18be0b4d424131741b7"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
requests = ">=2.21.0"
@ -418,6 +419,7 @@ files = [
{file = "azure_identity-1.25.1-py3-none-any.whl", hash = "sha256:e9edd720af03dff020223cd269fa3a61e8f345ea75443858273bcb44844ab651"},
{file = "azure_identity-1.25.1.tar.gz", hash = "sha256:87ca8328883de6036443e1c37b40e8dc8fb74898240f61071e09d2e369361456"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
azure-core = ">=1.31.0"
@ -718,7 +720,7 @@ files = [
{file = "cffi-2.0.0-cp39-cp39-win_amd64.whl", hash = "sha256:b882b3df248017dba09d6b16defe9b5c407fe32fc7c65a9c69798e6175601be9"},
{file = "cffi-2.0.0.tar.gz", hash = "sha256:44d1b5909021139fe36001ae048dbdde8214afa20200eda0f64c068cac5d5529"},
]
markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"proxy\"", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""}
markers = {main = "(platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""}
[package.dependencies]
pycparser = {version = "*", markers = "implementation_name != \"PyPy\""}
@ -1137,7 +1139,6 @@ description = "cryptography is a package which provides cryptographic recipes an
optional = false
python-versions = ">=3.7"
groups = ["main", "dev", "proxy-dev"]
markers = "python_version == \"3.9\""
files = [
{file = "cryptography-43.0.3-cp37-abi3-macosx_10_9_universal2.whl", hash = "sha256:bf7a1932ac4176486eab36a19ed4c0492da5d97123f1406cf15e41b05e787d2e"},
{file = "cryptography-43.0.3-cp37-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:63efa177ff54aec6e1c0aefaa1a241232dcd37413835a9b674b6e3f0ae2bfd3e"},
@ -1167,6 +1168,7 @@ files = [
{file = "cryptography-43.0.3-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:2ce6fae5bdad59577b44e4dfed356944fbf1d925269114c28be377692643b4ff"},
{file = "cryptography-43.0.3.tar.gz", hash = "sha256:315b9001266a492a6ff443b61238f956b214dbec9910a081ba5b6646a055a805"},
]
markers = {main = "python_version == \"3.9\" and (extra == \"proxy\" or extra == \"extra-proxy\")", dev = "python_version == \"3.9\"", proxy-dev = "python_version == \"3.9\""}
[package.dependencies]
cffi = {version = ">=1.12", markers = "platform_python_implementation != \"PyPy\""}
@ -1188,7 +1190,6 @@ description = "cryptography is a package which provides cryptographic recipes an
optional = false
python-versions = "!=3.9.0,!=3.9.1,>=3.8"
groups = ["main", "dev", "proxy-dev"]
markers = "python_version >= \"3.10\""
files = [
{file = "cryptography-46.0.3-cp311-abi3-macosx_10_9_universal2.whl", hash = "sha256:109d4ddfadf17e8e7779c39f9b18111a09efb969a301a31e987416a0191ed93a"},
{file = "cryptography-46.0.3-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:09859af8466b69bc3c27bdf4f5d84a665e0f7ab5088412e9e2ec49758eca5cbc"},
@ -1245,6 +1246,7 @@ files = [
{file = "cryptography-46.0.3-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:6b5063083824e5509fdba180721d55909ffacccc8adbec85268b48439423d78c"},
{file = "cryptography-46.0.3.tar.gz", hash = "sha256:a8b17438104fed022ce745b362294d9ce35b4c2e45c1d958ad4a4b019285f4a1"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
cffi = {version = ">=2.0.0", markers = "python_full_version >= \"3.9.0\" and platform_python_implementation != \"PyPy\""}
@ -2242,11 +2244,11 @@ files = [
]
[package.dependencies]
google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0dev", extras = ["grpc"]}
google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0dev"
grpc-google-iam-v1 = ">=0.12.4,<1.0.0dev"
proto-plus = ">=1.22.3,<2.0.0dev"
protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0dev"
google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0.dev0", extras = ["grpc"]}
google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0.dev0"
grpc-google-iam-v1 = ">=0.12.4,<1.0.0.dev0"
proto-plus = ">=1.22.3,<2.0.0.dev0"
protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0.dev0"
[[package]]
name = "google-cloud-resource-manager"
@ -3249,7 +3251,7 @@ files = [
[package.dependencies]
attrs = ">=22.2.0"
jsonschema-specifications = ">=2023.03.6"
jsonschema-specifications = ">=2023.3.6"
referencing = ">=0.28.4"
rpds-py = ">=0.7.1"
@ -3426,15 +3428,15 @@ files = [
[[package]]
name = "litellm-proxy-extras"
version = "0.4.29"
version = "0.4.30"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
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_proxy_extras-0.4.29-py3-none-any.whl", hash = "sha256:c36c1b69675c61acccc6b61dd610eb37daeb72c6fd819461cefb5b0cc7e0550f"},
{file = "litellm_proxy_extras-0.4.29.tar.gz", hash = "sha256:1a8266911e0546f1e17e6714ca20b72e9fef47c1683f9c16399cf2d1786437a0"},
{file = "litellm_proxy_extras-0.4.30-py3-none-any.whl", hash = "sha256:0b7df68f0968eb817462b847eaee81bba23d935adb2e84d2e342a77711887051"},
{file = "litellm_proxy_extras-0.4.30.tar.gz", hash = "sha256:5d32f8dc3d37d36fb15ab6995fea706dd8a453ff7f12e70b47cba35e5368da10"},
]
[[package]]
@ -3913,6 +3915,7 @@ files = [
{file = "msal-1.34.0-py3-none-any.whl", hash = "sha256:f669b1644e4950115da7a176441b0e13ec2975c29528d8b9e81316023676d6e1"},
{file = "msal-1.34.0.tar.gz", hash = "sha256:76ba83b716ea5a6d75b0279c0ac353a0e05b820ca1f6682c0eb7f45190c43c2f"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
cryptography = ">=2.5,<49"
@ -3933,6 +3936,7 @@ files = [
{file = "msal_extensions-1.3.1-py3-none-any.whl", hash = "sha256:96d3de4d034504e969ac5e85bae8106c8373b5c6568e4c8fa7af2eca9dbe6bca"},
{file = "msal_extensions-1.3.1.tar.gz", hash = "sha256:c5b0fd10f65ef62b5f1d62f4251d51cbcaf003fcedae8c91b040a488614be1a4"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
msal = ">=1.29,<2"
@ -4183,6 +4187,7 @@ files = [
{file = "nodeenv-1.9.1-py2.py3-none-any.whl", hash = "sha256:ba11c9782d29c27c70ffbdda2d7415098754709be8a7056d79a737cd901155c9"},
{file = "nodeenv-1.9.1.tar.gz", hash = "sha256:6ec12890a2dab7946721edbfbcd91f3319c6ccc9aec47be7c7e6b7011ee6645f"},
]
markers = {main = "extra == \"extra-proxy\""}
[[package]]
name = "numpy"
@ -4390,7 +4395,7 @@ files = [
{file = "opentelemetry_api-1.39.1-py3-none-any.whl", hash = "sha256:2edd8463432a7f8443edce90972169b195e7d6a05500cd29e6d13898187c9950"},
{file = "opentelemetry_api-1.39.1.tar.gz", hash = "sha256:fbde8c80e1b937a2c61f20347e91c0c18a1940cecf012d62e65a7caf08967c9c"},
]
markers = {main = "python_version >= \"3.10\""}
markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
[package.dependencies]
importlib-metadata = ">=6.0,<8.8.0"
@ -4505,7 +4510,7 @@ files = [
{file = "opentelemetry_sdk-1.39.1-py3-none-any.whl", hash = "sha256:4d5482c478513ecb0a5d938dcc61394e647066e0cc2676bee9f3af3f3f45f01c"},
{file = "opentelemetry_sdk-1.39.1.tar.gz", hash = "sha256:cf4d4563caf7bff906c9f7967e2be22d0d6b349b908be0d90fb21c8e9c995cc6"},
]
markers = {main = "python_version >= \"3.10\""}
markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
[package.dependencies]
opentelemetry-api = "1.39.1"
@ -4523,7 +4528,7 @@ files = [
{file = "opentelemetry_semantic_conventions-0.60b1-py3-none-any.whl", hash = "sha256:9fa8c8b0c110da289809292b0591220d3a7b53c1526a23021e977d68597893fb"},
{file = "opentelemetry_semantic_conventions-0.60b1.tar.gz", hash = "sha256:87c228b5a0669b748c76d76df6c364c369c28f1c465e50f661e39737e84bc953"},
]
markers = {main = "python_version >= \"3.10\""}
markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
[package.dependencies]
opentelemetry-api = "1.39.1"
@ -5000,6 +5005,7 @@ files = [
{file = "prisma-0.11.0-py3-none-any.whl", hash = "sha256:22bb869e59a2968b99f3483bb417717273ffbc569fd1e9ceed95e5614cbaf53a"},
{file = "prisma-0.11.0.tar.gz", hash = "sha256:3f2f2fd2361e1ec5ff655f2a04c7860c2f2a5bc4c91f78ca9c5c6349735bf693"},
]
markers = {main = "extra == \"extra-proxy\""}
[package.dependencies]
click = ">=7.1.2"
@ -5316,7 +5322,7 @@ files = [
{file = "pycparser-2.23-py3-none-any.whl", hash = "sha256:e5c6e8d3fbad53479cab09ac03729e0a9faf2bee3db8208a550daf5af81a5934"},
{file = "pycparser-2.23.tar.gz", hash = "sha256:78816d4f24add8f10a06d6f05b4d424ad9e96cfebf68a4ddc99c65c0720d00c2"},
]
markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""}
markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""}
[[package]]
name = "pydantic"
@ -5539,6 +5545,7 @@ files = [
{file = "PyJWT-2.10.1-py3-none-any.whl", hash = "sha256:dcdd193e30abefd5debf142f9adfcdd2b58004e644f25406ffaebd50bd98dacb"},
{file = "pyjwt-2.10.1.tar.gz", hash = "sha256:3cc5772eb20009233caf06e9d8a0577824723b44e6648ee0a2aedb6cf9381953"},
]
markers = {main = "extra == \"extra-proxy\" or extra == \"proxy\""}
[package.dependencies]
cryptography = {version = ">=3.4.0", optional = true, markers = "extra == \"crypto\""}
@ -6640,10 +6647,10 @@ files = [
]
[package.dependencies]
botocore = ">=1.37.4,<2.0a.0"
botocore = ">=1.37.4,<2.0a0"
[package.extras]
crt = ["botocore[crt] (>=1.37.4,<2.0a.0)"]
crt = ["botocore[crt] (>=1.37.4,<2.0a0)"]
[[package]]
name = "scikit-learn"
@ -6876,9 +6883,9 @@ tornado = ">=6.4.2,<7"
urllib3 = ">=1.26,<3"
[package.extras]
all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.00)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.0)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
bedrock = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)"]
cohere = ["cohere (>=5.9.4,<6.00)"]
cohere = ["cohere (>=5.9.4,<6.0)"]
dev = ["dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "ipykernel (>=6.25.0,<7)", "mypy (>=1.7.1,<2)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
docs = ["pydoc-markdown (>=4.8.2) ; python_version < \"3.12\""]
fastembed = ["fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\""]
@ -7722,6 +7729,7 @@ files = [
{file = "tomlkit-0.13.3-py3-none-any.whl", hash = "sha256:c89c649d79ee40629a9fda55f8ace8c6a1b42deb912b2a8fd8d942ddadb606b0"},
{file = "tomlkit-0.13.3.tar.gz", hash = "sha256:430cf247ee57df2b94ee3fbe588e71d362a941ebb545dec29b53961d61add2a1"},
]
markers = {main = "extra == \"extra-proxy\""}
[[package]]
name = "tornado"
@ -8490,4 +8498,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "95fd27dc139d0e52e70093220c50582f16c78e5977ec77f4297f50a30df964c6"
content-hash = "e5447e14dd37e324ac07a8fc6286d27e9a0d355ed93ebb24fc11e3f5df12fd3e"

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.81.7"
version = "1.81.8"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -61,7 +61,7 @@ boto3 = { version = "1.40.76", optional = true }
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
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.29", optional = true}
litellm-proxy-extras = {version = "0.4.30", optional = true}
rich = {version = "13.7.1", optional = true}
litellm-enterprise = {version = "0.1.27", optional = true}
diskcache = {version = "^5.6.1", optional = true}
@ -174,7 +174,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.81.7"
version = "1.81.8"
version_files = [
"pyproject.toml:^version"
]

View file

@ -50,7 +50,7 @@ sentry_sdk==2.21.0 # for sentry error handling
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
cryptography==44.0.1
tzdata==2025.1 # IANA time zone database
litellm-proxy-extras==0.4.29 # for proxy extras - e.g. prisma migrations
litellm-proxy-extras==0.4.30 # for proxy extras - e.g. prisma migrations
llm-sandbox==0.3.31 # for skill execution in sandbox
### LITELLM PACKAGE DEPENDENCIES
python-dotenv==1.0.1 # for env

View file

@ -129,6 +129,7 @@ model LiteLLM_TeamTable {
team_member_permissions String[] @default([])
policies String[] @default([])
model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team
litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id])
litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id])
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
@ -160,6 +161,7 @@ model LiteLLM_DeletedTeamTable {
team_member_permissions String[] @default([])
policies String[] @default([])
model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases
allow_team_guardrail_config Boolean @default(false)
// Original timestamps from team creation/updates
created_at DateTime? @map("created_at")

View file

@ -130,6 +130,67 @@ class BaseAnthropicMessagesTest:
return collected_chunks
@pytest.mark.asyncio
async def test_response_format_consistency(self):
"""
Test that response content blocks are consistently dicts (not Pydantic objects).
This ensures that code like response["content"][0]["type"] works
regardless of the target provider.
Issue: https://github.com/BerriAI/litellm/issues/20342
"""
litellm._turn_on_debug()
request_params = self.model_config
# Set up test parameters
messages = [{"role": "user", "content": "Say hi"}]
# Prepare call arguments
call_args = {
"messages": messages,
"max_tokens": 100,
}
# Add any additional config from subclass
call_args.update(request_params)
# Call the handler
response = await litellm.anthropic.messages.acreate(**call_args)
print(f"Response for {request_params['model']}: {json.dumps(response, indent=2, default=str)}")
# Verify response structure
assert "content" in response, "Response should have 'content' field"
assert len(response["content"]) > 0, "Response content should not be empty"
# Get the first content block
block = response["content"][0]
# Check that the block is a dict, not a Pydantic object
assert isinstance(block, dict), (
f"Content block should be a dict, but got {type(block)}. "
f"This means response format is inconsistent across providers."
)
# Verify we can access fields using dict syntax (not object attributes)
try:
block_type = block["type"]
print(f"✓ Successfully accessed block['type']: {block_type}")
except TypeError as e:
pytest.fail(
f"Cannot access content block using dict syntax: {e}. "
f"Block type: {type(block)}"
)
# Verify the block has expected structure
assert "type" in block, "Content block should have 'type' field"
if block["type"] == "text":
assert "text" in block, "Text content block should have 'text' field"
print(f"✓ Response format consistency test passed for {request_params['model']}")
@pytest.mark.asyncio
async def test_anthropic_messages_litellm_router_streaming_with_logging(self):
"""

View file

@ -876,4 +876,4 @@ def test_sync_openai_messages():
assert response is not None
assert isinstance(response, dict)
assert response["content"][0].text is not None
assert response["content"][0]["text"] is not None

View file

@ -341,10 +341,10 @@ def test_translate_openai_content_to_anthropic_empty_function_arguments():
result = adapter._translate_openai_content_to_anthropic(choices=openai_choices)
assert len(result) == 1
assert result[0].type == "tool_use"
assert result[0].id == "call_empty_args"
assert result[0].name == "test_function"
assert result[0].input == {}, "Empty function arguments should result in empty dict"
assert result[0]["type"] == "tool_use"
assert result[0]["id"] == "call_empty_args"
assert result[0]["name"] == "test_function"
assert result[0]["input"] == {}, "Empty function arguments should result in empty dict"
def test_translate_openai_content_to_anthropic_text_and_tool_calls():
@ -372,12 +372,12 @@ def test_translate_openai_content_to_anthropic_text_and_tool_calls():
result = adapter._translate_openai_content_to_anthropic(choices=openai_choices)
assert len(result) == 2
assert result[0].type == "text"
assert result[0].text == "Calling get_weather now."
assert result[1].type == "tool_use"
assert result[1].id == "call_weather"
assert result[1].name == "get_weather"
assert result[1].input == {"location": "Boston"}
assert result[0]["type"] == "text"
assert result[0]["text"] == "Calling get_weather now."
assert result[1]["type"] == "tool_use"
assert result[1]["id"] == "call_weather"
assert result[1]["name"] == "get_weather"
assert result[1]["input"] == {"location": "Boston"}
def test_translate_openai_response_to_anthropic_text_and_tool_calls():
@ -414,11 +414,11 @@ def test_translate_openai_response_to_anthropic_text_and_tool_calls():
anthropic_content = anthropic_response.get("content")
assert anthropic_content is not None
assert len(anthropic_content) == 2
assert cast(Any, anthropic_content[0]).type == "text"
assert cast(Any, anthropic_content[0]).text == "Let me grab the current weather."
assert cast(Any, anthropic_content[1]).type == "tool_use"
assert cast(Any, anthropic_content[1]).id == "call_tool_combo"
assert cast(Any, anthropic_content[1]).input == {"location": "Paris"}
assert anthropic_content[0]["type"] == "text"
assert anthropic_content[0]["text"] == "Let me grab the current weather."
assert anthropic_content[1]["type"] == "tool_use"
assert anthropic_content[1]["id"] == "call_tool_combo"
assert anthropic_content[1]["input"] == {"location": "Paris"}
assert anthropic_response.get("stop_reason") == "tool_use"
@ -484,11 +484,11 @@ def test_translate_openai_content_to_anthropic_thinking_and_redacted_thinking():
result = adapter._translate_openai_content_to_anthropic(choices=openai_choices)
assert len(result) == 2
assert result[0].type == "thinking"
assert result[0].thinking == "I need to summar"
assert result[0].signature == "sigsig"
assert result[1].type == "redacted_thinking"
assert result[1].data == "REDACTED"
assert result[0]["type"] == "thinking"
assert result[0]["thinking"] == "I need to summar"
assert result[0]["signature"] == "sigsig"
assert result[1]["type"] == "redacted_thinking"
assert result[1]["data"] == "REDACTED"
def test_translate_streaming_openai_chunk_to_anthropic_with_thinking():
@ -1443,13 +1443,13 @@ def test_translate_openai_content_to_anthropic_reasoning_content_without_thinkin
assert len(result) == 2
# First block should be thinking block with reasoning_content
assert result[0].type == "thinking"
assert "Considering Letter Frequency" in result[0].thinking
assert "Calculating the Count" in result[0].thinking
assert result[0].signature is None
assert result[0]["type"] == "thinking"
assert "Considering Letter Frequency" in result[0]["thinking"]
assert "Calculating the Count" in result[0]["thinking"]
assert result[0]["signature"] is None
# Second block should be text block with content
assert result[1].type == "text"
assert result[1].text == "There are **3** \"r\"s in the word strawberry."
assert result[1]["type"] == "text"
assert result[1]["text"] == "There are **3** \"r\"s in the word strawberry."
def test_translate_streaming_openai_chunk_to_anthropic_reasoning_content_without_thinking_blocks():
@ -1522,13 +1522,13 @@ def test_translate_openai_response_to_anthropic_with_reasoning_content_only():
assert len(anthropic_content) == 2
# First block should be thinking
assert cast(Any, anthropic_content[0]).type == "thinking"
assert "Considering Letter Frequency" in cast(Any, anthropic_content[0]).thinking
assert cast(Any, anthropic_content[0]).signature is None
assert anthropic_content[0]["type"] == "thinking"
assert "Considering Letter Frequency" in anthropic_content[0]["thinking"]
assert anthropic_content[0].get("signature") is None
# Second block should be text
assert cast(Any, anthropic_content[1]).type == "text"
assert cast(Any, anthropic_content[1]).text == "There are **3** \"r\"s in the word strawberry."
assert anthropic_content[1]["type"] == "text"
assert anthropic_content[1]["text"] == "There are **3** \"r\"s in the word strawberry."
assert anthropic_response.get("stop_reason") == "end_turn"
@ -1702,7 +1702,7 @@ def test_translate_openai_response_restores_tool_names():
)
# Find the tool_use block in the response
tool_use_blocks = [c for c in result["content"] if getattr(c, "type", None) == "tool_use"]
tool_use_blocks = [c for c in result["content"] if c.get("type") == "tool_use"]
assert len(tool_use_blocks) == 1
# Name should be restored to original
assert getattr(tool_use_blocks[0], "name", None) == original_name
assert tool_use_blocks[0]["name"] == original_name

View file

@ -132,7 +132,7 @@ async def test_update_daily_spend_with_null_entity_id():
assert create_data["model"] == "gpt-4"
assert create_data["custom_llm_provider"] == "openai"
assert create_data["mcp_namespaced_tool_name"] == ""
assert create_data["endpoint"] is None
assert create_data["endpoint"] == ""
assert create_data["prompt_tokens"] == 10
assert create_data["completion_tokens"] == 20
assert create_data["spend"] == 0.1
@ -194,7 +194,7 @@ async def test_update_daily_spend_sorting():
"model_group": None,
"mcp_namespaced_tool_name": "",
"custom_llm_provider": "openai",
"endpoint": None,
"endpoint": "",
"prompt_tokens": 10,
"completion_tokens": 20,
"spend": 0.1,
@ -838,4 +838,126 @@ async def test_endpoint_field_is_correctly_mapped_from_call_type():
assert transaction["date"] == "2024-01-01"
assert transaction["api_key"] == "test-key"
assert transaction["model"] == "gpt-4"
assert transaction["custom_llm_provider"] == "openai"
assert transaction["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure():
"""
Test that when batch upsert fails, detailed error information is logged.
This ensures proper debugging information is available for issues like unique constraint violations.
"""
from litellm._logging import verbose_proxy_logger
# Setup
mock_prisma_client = MagicMock()
mock_batcher = MagicMock()
mock_table = MagicMock()
mock_batch_context = MagicMock()
mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher)
mock_batcher.litellm_dailyuserspend = mock_table
# Make the batch context manager's exit raise an exception
# This simulates a batch commit failure (e.g., unique constraint violation)
test_exception = Exception("Unique constraint violation")
mock_batch_context.__aexit__ = AsyncMock(side_effect=test_exception)
mock_prisma_client.db.batch_.return_value = mock_batch_context
# Create a transaction
daily_spend_transactions = {
"test_key": {
"user_id": "test-user",
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
"custom_llm_provider": "openai",
"prompt_tokens": 10,
"completion_tokens": 20,
"spend": 0.1,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
}
}
# Create a mock proxy_logging_obj with failure_handler as AsyncMock
mock_proxy_logging = MagicMock()
mock_proxy_logging.failure_handler = AsyncMock()
# Mock the logger to capture exception calls
with patch.object(verbose_proxy_logger, 'exception') as mock_exception_logger:
# Call the method and expect it to raise the exception
with pytest.raises(Exception, match="Unique constraint violation"):
await DBSpendUpdateWriter._update_daily_spend(
n_retry_times=0, # No retries to make test faster
prisma_client=mock_prisma_client,
proxy_logging_obj=mock_proxy_logging,
daily_spend_transactions=daily_spend_transactions,
entity_type="user",
entity_id_field="user_id",
table_name="litellm_dailyuserspend",
unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
# Verify that exception was logged with detailed information
assert mock_exception_logger.called
call_args = mock_exception_logger.call_args[0][0]
assert "Daily user spend batch upsert failed" in call_args
assert "Table: litellm_dailyuserspend" in call_args
assert "Constraint: user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint" in call_args
assert "Batch size: 1" in call_args
assert "Unique constraint violation" in call_args
@pytest.mark.asyncio
async def test_update_daily_spend_re_raises_exception_after_logging():
"""
Test that when batch upsert fails, the exception is properly re-raised after logging.
This ensures that error handling continues to work correctly upstream.
"""
# Setup
mock_prisma_client = MagicMock()
mock_batcher = MagicMock()
mock_table = MagicMock()
mock_batch_context = MagicMock()
mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher)
mock_batcher.litellm_dailyuserspend = mock_table
# Create a transaction
daily_spend_transactions = {
"test_key": {
"user_id": "test-user",
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
"custom_llm_provider": "openai",
"prompt_tokens": 10,
"completion_tokens": 20,
"spend": 0.1,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
}
}
# Create a custom exception to verify it's re-raised
custom_exception = ValueError("Database connection lost")
mock_batch_context.__aexit__ = AsyncMock(side_effect=custom_exception)
mock_prisma_client.db.batch_.return_value = mock_batch_context
# Create a mock proxy_logging_obj with failure_handler as AsyncMock
mock_proxy_logging = MagicMock()
mock_proxy_logging.failure_handler = AsyncMock()
# Verify the exception is re-raised
with pytest.raises(ValueError, match="Database connection lost"):
await DBSpendUpdateWriter._update_daily_spend(
n_retry_times=0, # No retries to make test faster
prisma_client=mock_prisma_client,
proxy_logging_obj=mock_proxy_logging,
daily_spend_transactions=daily_spend_transactions,
entity_type="user",
entity_id_field="user_id",
table_name="litellm_dailyuserspend",
unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)

View file

@ -55,7 +55,7 @@ async def test_route_a2a_model_bypasses_router():
with patch("litellm.acompletion", mock_acompletion):
with patch(
"litellm.proxy.agent_endpoints.a2a_routing.global_agent_registry",
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
mock_registry,
):
result = await route_request(

View file

@ -532,6 +532,7 @@ async def test_team_update_sc_2():
or k == "object_permission"
or k == "litellm_model_table"
or k == "policies"
or k == "allow_team_guardrail_config"
):
pass
else:

View file

@ -0,0 +1,19 @@
import { getMCPSemanticFilterSettings } from "@/components/networking";
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import useAuthorized from "../useAuthorized";
const mcpSemanticFilterSettingsKeys = createQueryKeys(
"mcpSemanticFilterSettings"
);
export const useMCPSemanticFilterSettings = () => {
const { accessToken } = useAuthorized();
return useQuery<Record<string, any>>({
queryKey: mcpSemanticFilterSettingsKeys.list({}),
queryFn: async () => await getMCPSemanticFilterSettings(accessToken),
enabled: !!accessToken,
staleTime: 60 * 60 * 1000, // 1 hour
gcTime: 60 * 60 * 1000, // 1 hour
});
};

View file

@ -0,0 +1,25 @@
import { updateMCPSemanticFilterSettings } from "@/components/networking";
import { useMutation, useQueryClient } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
const mcpSemanticFilterSettingsKeys = createQueryKeys(
"mcpSemanticFilterSettings"
);
export const useUpdateMCPSemanticFilterSettings = (accessToken: string) => {
const queryClient = useQueryClient();
return useMutation({
mutationFn: async (settings: Record<string, any>) => {
if (!accessToken) {
throw new Error("Access token is required");
}
return updateMCPSemanticFilterSettings(accessToken, settings);
},
onSuccess: () => {
queryClient.invalidateQueries({
queryKey: mcpSemanticFilterSettingsKeys.all,
});
},
});
};

View file

@ -27,7 +27,7 @@ import Organizations, { fetchOrganizations } from "@/components/organizations";
import PassThroughSettings from "@/components/pass_through_settings";
import PromptsPanel from "@/components/prompts";
import PublicModelHub from "@/components/public_model_hub";
import { SearchTools } from "@/components/search_tools";
import { SearchTools } from "@/components/SearchTools";
import Settings from "@/components/settings";
import { SurveyPrompt, SurveyModal, ClaudeCodePrompt, ClaudeCodeModal } from "@/components/survey";
import TagManagement from "@/components/tag_management";

View file

@ -1,14 +1,14 @@
import React, { useState } from "react";
import { Modal, Tooltip, Form, Select, Input, Typography } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { Button, TextInput } from "@tremor/react";
import { createSearchTool, fetchAvailableSearchProviders } from "../networking";
import { SearchTool, AvailableSearchProvider } from "./types";
import { isAdminRole } from "@/utils/roles";
import NotificationsManager from "../molecules/notifications_manager";
import { InfoCircleOutlined } from "@ant-design/icons";
import { useQuery } from "@tanstack/react-query";
import SearchConnectionTest from "./search_connection_test";
import { Button, TextInput } from "@tremor/react";
import { Form, Input, Modal, Select, Tooltip, Typography } from "antd";
import Image from "next/image";
import React, { useState } from "react";
import NotificationsManager from "../molecules/notifications_manager";
import { createSearchTool, fetchAvailableSearchProviders } from "../networking";
import SearchConnectionTest from "./SearchConnectionTest";
import { AvailableSearchProvider, SearchTool } from "./types";
const { TextArea } = Input;
@ -97,8 +97,8 @@ const CreateSearchTool: React.FC<CreateSearchToolProps> = ({
},
search_tool_info: formValues.description
? {
description: formValues.description,
}
description: formValues.description,
}
: undefined,
};
@ -130,7 +130,7 @@ const CreateSearchTool: React.FC<CreateSearchToolProps> = ({
try {
// Validate required fields for testing
await form.validateFields(["search_provider", "api_key"]);
setIsTestingConnection(true);
// Generate a new test ID (using timestamp for uniqueness)
setConnectionTestId(`test-${Date.now()}`);
@ -225,8 +225,8 @@ const CreateSearchTool: React.FC<CreateSearchToolProps> = ({
optionLabelProp="label"
>
{availableProviders.map((provider) => (
<Select.Option
key={provider.provider_name}
<Select.Option
key={provider.provider_name}
value={provider.provider_name}
label={
<SearchProviderLabel

View file

@ -1,8 +1,8 @@
import { InfoCircleOutlined, WarningOutlined } from "@ant-design/icons";
import { Button, Divider, Typography } from "antd";
import React, { useEffect, useState } from "react";
import { testSearchToolConnection } from "../networking";
import { Button, Typography, Divider } from "antd";
import { WarningOutlined, InfoCircleOutlined } from "@ant-design/icons";
import NotificationsManager from "../molecules/notifications_manager";
import { testSearchToolConnection } from "../networking";
const { Text } = Typography;

View file

@ -0,0 +1,114 @@
import { Tag } from "antd";
import { ColumnsType } from "antd/es/table";
import TableIconActionButton from "../common_components/IconActionButton/TableIconActionButtons/TableIconActionButton";
import { SearchTool } from "./types";
export const searchToolColumns = (
onView: (searchToolId: string) => void,
onEdit: (searchToolId: string) => void,
onDelete: (searchToolId: string) => void,
availableProviders: Array<{ provider_name: string; ui_friendly_name: string }>,
): ColumnsType<SearchTool> => [
{
title: "Search Tool ID",
dataIndex: "search_tool_id",
key: "search_tool_id",
render: (_, tool) => {
const isFromConfig = tool.is_from_config;
if (isFromConfig) {
return <span className="text-xs">-</span>;
}
return (
<button
onClick={() => onView(tool.search_tool_id!)}
className="font-mono text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs font-normal px-2 py-0.5 text-left cursor-pointer max-w-40"
>
<span className="truncate block">{tool.search_tool_id}</span>
</button>
);
},
},
{
title: "Name",
dataIndex: "search_tool_name",
key: "search_tool_name",
render: (name: string) => <span className="font-medium">{name}</span>,
},
{
title: "Provider",
key: "provider",
render: (_, tool) => {
const provider = tool.litellm_params.search_provider;
const providerInfo = availableProviders.find((p) => p.provider_name === provider);
const displayName = providerInfo?.ui_friendly_name || provider;
return <span className="text-sm">{displayName}</span>;
},
},
{
title: "Created At",
dataIndex: "created_at",
key: "created_at",
render: (_, tool) => {
return <span className="text-xs">{tool.created_at ? new Date(tool.created_at).toLocaleDateString() : "-"}</span>;
},
},
{
title: "Updated At",
dataIndex: "updated_at",
key: "updated_at",
render: (_, tool) => {
return <span className="text-xs">{tool.updated_at ? new Date(tool.updated_at).toLocaleDateString() : "-"}</span>;
},
},
{
title: "Source",
key: "source",
render: (_, tool) => {
const isFromConfig = tool.is_from_config ?? false;
return (
<Tag color={isFromConfig ? "default" : "blue"}>
{isFromConfig ? "Config" : "DB"}
</Tag>
);
},
},
{
title: "Actions",
key: "actions",
render: (_, tool) => {
const toolId = tool.search_tool_id;
const isFromConfig = tool.is_from_config ?? false;
return (
<div className="flex items-center gap-2">
<TableIconActionButton
variant="Edit"
tooltipText="Edit search tool"
disabled={isFromConfig}
disabledTooltipText="Config search tool cannot be edited on the dashboard. Please edit it from the config file."
onClick={() => {
if (toolId && !isFromConfig) {
onEdit(toolId);
}
}}
/>
<TableIconActionButton
variant="Delete"
tooltipText="Delete search tool"
disabled={isFromConfig}
disabledTooltipText="Config search tool cannot be deleted on the dashboard. Please delete it from the config file."
onClick={() => {
if (toolId && !isFromConfig) {
onDelete(toolId);
}
}}
/>
</div>
);
},
},
];

View file

@ -0,0 +1,278 @@
import { render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { SearchToolView } from "./SearchToolView";
import { AvailableSearchProvider, SearchTool } from "./types";
vi.mock("@/utils/dataUtils", () => ({
copyToClipboard: vi.fn().mockResolvedValue(true),
}));
vi.mock("./SearchToolTester", () => ({
SearchToolTester: ({ searchToolName, accessToken }: { searchToolName: string; accessToken: string }) => (
<div data-testid="search-tool-tester">
<span>Search Tool Tester for {searchToolName}</span>
<span>Access Token: {accessToken}</span>
</div>
),
}));
describe("SearchToolView", () => {
const mockSearchTool: SearchTool = {
search_tool_id: "test-tool-id-123",
search_tool_name: "Test Search Tool",
litellm_params: {
search_provider: "perplexity",
api_key: "sk-test-key",
},
search_tool_info: {
description: "Test description",
},
created_at: "2024-01-15T10:30:00Z",
};
const mockAvailableProviders: AvailableSearchProvider[] = [
{
provider_name: "perplexity",
ui_friendly_name: "Perplexity AI",
},
{
provider_name: "tavily",
ui_friendly_name: "Tavily Search",
},
];
const defaultProps = {
searchTool: mockSearchTool,
onBack: vi.fn(),
isEditing: false,
accessToken: "test-token",
availableProviders: mockAvailableProviders,
};
beforeEach(async () => {
vi.clearAllMocks();
const { copyToClipboard } = await import("@/utils/dataUtils");
vi.mocked(copyToClipboard).mockResolvedValue(true);
});
it("should render", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("Test Search Tool")).toBeInTheDocument();
});
it("should display search tool name", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("Test Search Tool")).toBeInTheDocument();
});
it("should display search tool ID", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("test-tool-id-123")).toBeInTheDocument();
});
it("should display provider name using UI-friendly name when available", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("Perplexity AI")).toBeInTheDocument();
});
it("should display provider name using provider_name when UI-friendly name is not available", () => {
const searchToolWithoutProvider: SearchTool = {
...mockSearchTool,
litellm_params: {
search_provider: "unknown-provider",
},
};
render(
<SearchToolView
{...defaultProps}
searchTool={searchToolWithoutProvider}
/>,
);
expect(screen.getByText("unknown-provider")).toBeInTheDocument();
});
it("should display masked API key when API key is set", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("****")).toBeInTheDocument();
});
it("should display 'Not set' when API key is not set", () => {
const searchToolWithoutApiKey: SearchTool = {
...mockSearchTool,
litellm_params: {
search_provider: "perplexity",
},
};
render(
<SearchToolView
{...defaultProps}
searchTool={searchToolWithoutApiKey}
/>,
);
expect(screen.getByText("Not set")).toBeInTheDocument();
});
it("should display formatted created_at date", () => {
render(<SearchToolView {...defaultProps} />);
const dateText = screen.getByText(/2024-01-15/);
expect(dateText).toBeInTheDocument();
});
it("should display 'Unknown' when created_at is not set", () => {
const searchToolWithoutDate: SearchTool = {
...mockSearchTool,
created_at: undefined,
};
render(
<SearchToolView
{...defaultProps}
searchTool={searchToolWithoutDate}
/>,
);
expect(screen.getByText("Unknown")).toBeInTheDocument();
});
it("should display description when search_tool_info.description is provided", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("Test description")).toBeInTheDocument();
});
it("should not display description card when search_tool_info.description is not provided", () => {
const searchToolWithoutDescription: SearchTool = {
...mockSearchTool,
search_tool_info: {},
};
render(
<SearchToolView
{...defaultProps}
searchTool={searchToolWithoutDescription}
/>,
);
expect(screen.queryByText("Description")).not.toBeInTheDocument();
});
it("should call onBack when back button is clicked", async () => {
const user = userEvent.setup({ delay: null });
const onBack = vi.fn();
render(<SearchToolView {...defaultProps} onBack={onBack} />);
const backButton = screen.getByRole("button", { name: /back to all search tools/i });
await user.click(backButton);
expect(onBack).toHaveBeenCalledTimes(1);
});
it("should copy search tool name to clipboard when copy button is clicked", async () => {
const user = userEvent.setup({ delay: null });
const { copyToClipboard } = await import("@/utils/dataUtils");
render(<SearchToolView {...defaultProps} />);
const toolNameContainer = screen.getByText("Test Search Tool").closest("div");
expect(toolNameContainer).toBeInTheDocument();
const copyButtons = within(toolNameContainer!).getAllByRole("button");
const nameCopyButton = copyButtons.find((button) => {
return button.querySelector("svg") !== null;
});
expect(nameCopyButton).toBeInTheDocument();
await user.click(nameCopyButton!);
await waitFor(() => {
expect(copyToClipboard).toHaveBeenCalledWith("Test Search Tool");
});
});
it("should copy search tool ID to clipboard when copy button is clicked", async () => {
const user = userEvent.setup({ delay: null });
const { copyToClipboard } = await import("@/utils/dataUtils");
render(<SearchToolView {...defaultProps} />);
const toolIdContainer = screen.getByText("test-tool-id-123").closest("div");
expect(toolIdContainer).toBeInTheDocument();
const copyButtons = within(toolIdContainer!).getAllByRole("button");
const idCopyButton = copyButtons.find((button) => {
return button.querySelector("svg") !== null;
});
expect(idCopyButton).toBeInTheDocument();
await user.click(idCopyButton!);
await waitFor(() => {
expect(copyToClipboard).toHaveBeenCalledWith("test-tool-id-123");
});
});
it("should show check icon after copying search tool name", async () => {
const user = userEvent.setup({ delay: null });
const { copyToClipboard } = await import("@/utils/dataUtils");
vi.mocked(copyToClipboard).mockResolvedValue(true);
render(<SearchToolView {...defaultProps} />);
const toolNameContainer = screen.getByText("Test Search Tool").closest("div");
const copyButtons = within(toolNameContainer!).getAllByRole("button");
const nameCopyButton = copyButtons.find((button) => {
return button.querySelector("svg") !== null;
});
expect(nameCopyButton).toBeInTheDocument();
const initialSvg = nameCopyButton!.querySelector("svg");
expect(initialSvg).toBeInTheDocument();
await user.click(nameCopyButton!);
await waitFor(() => {
const updatedSvg = nameCopyButton!.querySelector("svg");
expect(updatedSvg).toBeInTheDocument();
expect(nameCopyButton).toHaveClass("text-green-600");
});
});
it("should not show check icon when copy fails", async () => {
const user = userEvent.setup({ delay: null });
const { copyToClipboard } = await import("@/utils/dataUtils");
vi.mocked(copyToClipboard).mockResolvedValue(false);
render(<SearchToolView {...defaultProps} />);
const toolNameContainer = screen.getByText("Test Search Tool").closest("div");
const copyButtons = within(toolNameContainer!).getAllByRole("button");
const nameCopyButton = copyButtons.find((button) => {
return button.querySelector("svg") !== null;
});
expect(nameCopyButton).toBeInTheDocument();
await user.click(nameCopyButton!);
await waitFor(() => {
expect(copyToClipboard).toHaveBeenCalledWith("Test Search Tool");
}, { timeout: 3000 });
expect(nameCopyButton).not.toHaveClass("text-green-600");
});
it("should render SearchToolTester when accessToken is provided", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByTestId("search-tool-tester")).toBeInTheDocument();
expect(screen.getByText(/Search Tool Tester for Test Search Tool/)).toBeInTheDocument();
});
it("should not render SearchToolTester when accessToken is null", () => {
render(<SearchToolView {...defaultProps} accessToken={null} />);
expect(screen.queryByTestId("search-tool-tester")).not.toBeInTheDocument();
});
it("should pass correct props to SearchToolTester", () => {
render(<SearchToolView {...defaultProps} />);
expect(screen.getByText("Access Token: test-token")).toBeInTheDocument();
});
});

View file

@ -1,11 +1,11 @@
import React, { useState } from "react";
import { ArrowLeftIcon } from "@heroicons/react/outline";
import { Title, Card, Button, Text, Grid } from "@tremor/react";
import { SearchTool, AvailableSearchProvider } from "./types";
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils";
import { CheckIcon, CopyIcon } from "lucide-react";
import { ArrowLeftIcon } from "@heroicons/react/outline";
import { Button, Card, Grid, Text, Title } from "@tremor/react";
import { Button as AntdButton } from "antd";
import { SearchToolTester } from "./search_tool_tester";
import { CheckIcon, CopyIcon } from "lucide-react";
import React, { useState } from "react";
import { SearchToolTester } from "./SearchToolTester";
import { AvailableSearchProvider, SearchTool } from "./types";
interface SearchToolViewProps {
searchTool: SearchTool;
@ -53,11 +53,10 @@ export const SearchToolView: React.FC<SearchToolViewProps> = ({
size="small"
icon={copiedStates["search-tool-name"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
onClick={() => copyToClipboard(searchTool.search_tool_name, "search-tool-name")}
className={`left-2 z-10 transition-all duration-200 ${
copiedStates["search-tool-name"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
className={`left-2 z-10 transition-all duration-200 ${copiedStates["search-tool-name"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
/>
</div>
<div className="flex items-center cursor-pointer">
@ -67,11 +66,10 @@ export const SearchToolView: React.FC<SearchToolViewProps> = ({
size="small"
icon={copiedStates["search-tool-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
onClick={() => copyToClipboard(searchTool.search_tool_id, "search-tool-id")}
className={`left-2 z-10 transition-all duration-200 ${
copiedStates["search-tool-id"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
className={`left-2 z-10 transition-all duration-200 ${copiedStates["search-tool-id"]
? "text-green-600 bg-green-50 border-green-200"
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
}`}
/>
</div>
</div>

View file

@ -0,0 +1,234 @@
import * as roles from "@/utils/roles";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import * as networking from "../networking";
import SearchTools from "./SearchTools";
import { AvailableSearchProvider, SearchTool } from "./types";
vi.mock("../networking", () => ({
fetchSearchTools: vi.fn(),
updateSearchTool: vi.fn(),
deleteSearchTool: vi.fn(),
fetchAvailableSearchProviders: vi.fn(),
}));
vi.mock("@/utils/roles", () => ({
isAdminRole: vi.fn(),
}));
vi.mock("./SearchToolView", () => ({
SearchToolView: ({ searchTool, onBack }: { searchTool: SearchTool; onBack: () => void }) => (
<div data-testid="search-tool-view">
<div>Search Tool View: {searchTool.search_tool_name}</div>
<button onClick={onBack}>Back</button>
</div>
),
}));
vi.mock("./CreateSearchTools", () => ({
default: ({
isModalVisible,
setModalVisible,
}: {
isModalVisible: boolean;
setModalVisible: (visible: boolean) => void;
}) =>
isModalVisible ? (
<div data-testid="create-search-tool-modal">
<button onClick={() => setModalVisible(false)}>Close Create Modal</button>
</div>
) : null,
}));
vi.mock("../common_components/DeleteResourceModal", () => ({
default: ({
isOpen,
onOk,
onCancel,
}: {
isOpen: boolean;
onOk: () => void;
onCancel: () => void;
}) =>
isOpen ? (
<div data-testid="delete-resource-modal">
<button onClick={onOk}>Confirm Delete</button>
<button onClick={onCancel}>Cancel Delete</button>
</div>
) : null,
}));
const mockSearchTools: SearchTool[] = [
{
search_tool_id: "tool-1",
search_tool_name: "Perplexity Search",
litellm_params: {
search_provider: "perplexity",
api_key: "sk-test-key",
},
search_tool_info: {
description: "Test description",
},
created_at: "2024-01-15T10:30:00Z",
},
{
search_tool_id: "tool-2",
search_tool_name: "Tavily Search",
litellm_params: {
search_provider: "tavily",
},
created_at: "2024-01-16T10:30:00Z",
},
];
const mockAvailableProviders: AvailableSearchProvider[] = [
{
provider_name: "perplexity",
ui_friendly_name: "Perplexity AI",
},
{
provider_name: "tavily",
ui_friendly_name: "Tavily Search",
},
];
const createWrapper = () => {
const queryClient = new QueryClient({
defaultOptions: {
queries: {
retry: false,
},
},
});
return ({ children }: { children: React.ReactNode }) => (
<QueryClientProvider client={queryClient}>{children}</QueryClientProvider>
);
};
describe("SearchTools", () => {
const defaultProps = {
accessToken: "test-token",
userRole: "Admin",
userID: "user-1",
};
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(networking.fetchSearchTools).mockResolvedValue({ search_tools: mockSearchTools });
vi.mocked(networking.fetchAvailableSearchProviders).mockResolvedValue({ providers: mockAvailableProviders });
vi.mocked(roles.isAdminRole).mockReturnValue(true);
});
it("should render", async () => {
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByText("Search Tools")).toBeInTheDocument();
});
});
it("should display missing authentication parameters message when accessToken is missing", () => {
render(<SearchTools {...defaultProps} accessToken={null} />, { wrapper: createWrapper() });
expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument();
});
it("should display missing authentication parameters message when userRole is missing", () => {
render(<SearchTools {...defaultProps} userRole={null} />, { wrapper: createWrapper() });
expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument();
});
it("should display missing authentication parameters message when userID is missing", () => {
render(<SearchTools {...defaultProps} userID={null} />, { wrapper: createWrapper() });
expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument();
});
it("should display search tools table with tools", async () => {
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByText("Perplexity Search")).toBeInTheDocument();
});
expect(screen.getAllByText("Tavily Search").length).toBeGreaterThan(0);
});
it("should display empty state when no search tools are available", async () => {
vi.mocked(networking.fetchSearchTools).mockResolvedValue({ search_tools: [] });
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByText("No search tools configured")).toBeInTheDocument();
});
});
it("should show Add New Search Tool button when user is admin", async () => {
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByRole("button", { name: /add new search tool/i })).toBeInTheDocument();
});
});
it("should not show Add New Search Tool button when user is not admin", async () => {
vi.mocked(roles.isAdminRole).mockReturnValue(false);
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByText("Search Tools")).toBeInTheDocument();
});
expect(screen.queryByRole("button", { name: /add new search tool/i })).not.toBeInTheDocument();
});
it("should open create modal when Add New Search Tool button is clicked", async () => {
const user = userEvent.setup({ delay: null });
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByRole("button", { name: /add new search tool/i })).toBeInTheDocument();
});
const addButton = screen.getByRole("button", { name: /add new search tool/i });
await user.click(addButton);
expect(screen.getByTestId("create-search-tool-modal")).toBeInTheDocument();
});
it("should navigate to tool view when tool ID is clicked", async () => {
const user = userEvent.setup({ delay: null });
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByText("Perplexity Search")).toBeInTheDocument();
});
const toolIdButton = screen.getByRole("button", { name: /tool-1/i });
await user.click(toolIdButton);
await waitFor(() => {
expect(screen.getByTestId("search-tool-view")).toBeInTheDocument();
});
expect(screen.getByText(/Search Tool View: Perplexity Search/i)).toBeInTheDocument();
});
it("should navigate back from tool view to table", async () => {
const user = userEvent.setup({ delay: null });
render(<SearchTools {...defaultProps} />, { wrapper: createWrapper() });
await waitFor(() => {
expect(screen.getByText("Perplexity Search")).toBeInTheDocument();
});
const toolIdButton = screen.getByRole("button", { name: /tool-1/i });
await user.click(toolIdButton);
await waitFor(() => {
expect(screen.getByTestId("search-tool-view")).toBeInTheDocument();
});
const backButton = screen.getByRole("button", { name: /back/i });
await user.click(backButton);
await waitFor(() => {
expect(screen.queryByTestId("search-tool-view")).not.toBeInTheDocument();
expect(screen.getByText("Perplexity Search")).toBeInTheDocument();
});
});
});

View file

@ -1,20 +1,21 @@
import React, { useState } from "react";
import { isAdminRole } from "@/utils/roles";
import { LoadingOutlined } from "@ant-design/icons";
import { useQuery } from "@tanstack/react-query";
import { Modal, Form, Input, Select } from "antd";
import { Button, Title, Text, Grid, Col } from "@tremor/react";
import { DataTable } from "../view_logs/table";
import { searchToolColumns } from "./search_tool_columns";
import { Button, Text, Title } from "@tremor/react";
import { Form, Input, Modal, Select, Spin, Table } from "antd";
import React, { useState } from "react";
import DeleteResourceModal from "../common_components/DeleteResourceModal";
import NotificationsManager from "../molecules/notifications_manager";
import {
fetchSearchTools,
updateSearchTool,
deleteSearchTool,
fetchAvailableSearchProviders,
fetchSearchTools,
updateSearchTool,
} from "../networking";
import { SearchTool, AvailableSearchProvider } from "./types";
import { isAdminRole } from "@/utils/roles";
import NotificationsManager from "../molecules/notifications_manager";
import { SearchToolView } from "./search_tool_view";
import CreateSearchTool from "./create_search_tool";
import CreateSearchTool from "./CreateSearchTools";
import { searchToolColumns } from "./SearchToolColumn";
import { SearchToolView } from "./SearchToolView";
import { AvailableSearchProvider, SearchTool } from "./types";
interface SearchToolsProps {
accessToken: string | null;
@ -22,24 +23,6 @@ interface SearchToolsProps {
userID: string | null;
}
const DeleteModal: React.FC<{
isModalOpen: boolean;
title: string;
confirmDelete: () => void;
cancelDelete: () => void;
}> = ({ isModalOpen, title, confirmDelete, cancelDelete }) => {
if (!isModalOpen) return null;
return (
<Modal open={isModalOpen} onOk={confirmDelete} okType="danger" onCancel={cancelDelete}>
<Grid numItems={1} className="gap-2 w-full">
<Title>{title}</Title>
<Col numColSpan={1}>
<p>Are you sure you want to delete this search tool?</p>
</Col>
</Grid>
</Modal>
);
};
const SearchTools: React.FC<SearchToolsProps> = ({ accessToken, userRole, userID }) => {
const {
@ -72,6 +55,7 @@ const SearchTools: React.FC<SearchToolsProps> = ({ accessToken, userRole, userID
// State
const [toolIdToDelete, setToolToDelete] = useState<string | null>(null);
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
const [isDeleting, setIsDeleting] = useState(false);
const [selectedToolId, setSelectedToolId] = useState<string | null>(null);
const [editTool, setEditTool] = useState(false);
const [isCreateModalVisible, setCreateModalVisible] = useState(false);
@ -116,16 +100,19 @@ const SearchTools: React.FC<SearchToolsProps> = ({ accessToken, userRole, userID
if (toolIdToDelete == null || accessToken == null) {
return;
}
setIsDeleting(true);
try {
await deleteSearchTool(accessToken, toolIdToDelete);
NotificationsManager.success("Deleted search tool successfully");
setIsDeleteModalOpen(false);
setToolToDelete(null);
refetch();
} catch (error) {
console.error("Error deleting the search tool:", error);
NotificationsManager.error("Failed to delete search tool");
} finally {
setIsDeleting(false);
}
setIsDeleteModalOpen(false);
setToolToDelete(null);
};
const cancelDelete = () => {
@ -133,6 +120,11 @@ const SearchTools: React.FC<SearchToolsProps> = ({ accessToken, userRole, userID
setToolToDelete(null);
};
const toolToDelete = searchTools?.find((t) => t.search_tool_id === toolIdToDelete);
const providerInfo = toolToDelete
? availableProviders.find((p) => p.provider_name === toolToDelete.litellm_params.search_provider)
: null;
const handleCreateSuccess = (newSearchTool: SearchTool) => {
setCreateModalVisible(false);
refetch();
@ -231,26 +223,46 @@ const SearchTools: React.FC<SearchToolsProps> = ({ accessToken, userRole, userID
/>
) : (
<div className="w-full h-full">
<div className="w-full px-6 mt-6">
<DataTable
data={searchTools || []}
<Spin spinning={isLoadingTools} indicator={<LoadingOutlined spin />} size="large">
<Table
bordered
dataSource={searchTools || []}
columns={columns}
renderSubComponent={() => <div></div>}
getRowCanExpand={() => false}
isLoading={isLoadingTools}
noDataMessage="No search tools configured"
rowKey={(record) => record.search_tool_id || record.search_tool_name}
pagination={false}
locale={{
emptyText: "No search tools configured",
}}
size="small"
/>
</div>
</Spin>
</div>
);
return (
<div className="w-full h-full p-6">
<DeleteModal
isModalOpen={isDeleteModalOpen}
<DeleteResourceModal
isOpen={isDeleteModalOpen}
title="Delete Search Tool"
confirmDelete={confirmDelete}
cancelDelete={cancelDelete}
message="Are you sure you want to delete this search tool? This action cannot be undone."
resourceInformationTitle="Search Tool Information"
resourceInformation={
toolToDelete
? [
{ label: "Name", value: toolToDelete.search_tool_name },
{ label: "ID", value: toolToDelete.search_tool_id, code: true },
{
label: "Provider",
value: providerInfo?.ui_friendly_name || toolToDelete.litellm_params.search_provider,
},
{ label: "Description", value: toolToDelete.search_tool_info?.description || "-" },
]
: []
}
onCancel={cancelDelete}
onOk={confirmDelete}
confirmLoading={isDeleting}
/>
<CreateSearchTool

View file

@ -0,0 +1,6 @@
export { default as SearchTools } from './SearchTools';
export { SearchToolView } from './SearchToolView';
export { default as SearchConnectionTest } from './SearchConnectionTest';
export { SearchToolTester } from './SearchToolTester';
export * from './types';

View file

@ -19,6 +19,7 @@ export interface SearchTool {
search_tool_info?: SearchToolInfo;
created_at?: string;
updated_at?: string;
is_from_config?: boolean;
}
export interface SearchToolsResponse {

View file

@ -0,0 +1,309 @@
"use client";
import { useMCPSemanticFilterSettings } from "@/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings";
import { useUpdateMCPSemanticFilterSettings } from "@/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings";
import NotificationManager from "@/components/molecules/notifications_manager";
import {
Alert,
Button,
Card,
Col,
Form,
InputNumber,
Row,
Select,
Skeleton,
Slider,
Space,
Switch,
Typography,
Tooltip,
} from "antd";
import { QuestionCircleOutlined, CheckCircleOutlined, SaveOutlined } from "@ant-design/icons";
import { useEffect, useState } from "react";
import { fetchAvailableModels, ModelGroup } from "@/components/playground/llm_calls/fetch_models";
import MCPSemanticFilterTestPanel from "./MCPSemanticFilterTestPanel";
import { getCurlCommand, runSemanticFilterTest, TestResult } from "./semanticFilterTestUtils";
interface MCPSemanticFilterSettingsProps {
accessToken: string | null;
}
export default function MCPSemanticFilterSettings({ accessToken }: MCPSemanticFilterSettingsProps) {
const { data, isLoading, isError, error } = useMCPSemanticFilterSettings();
const {
mutate: updateSettings,
isPending: isUpdating,
error: updateError,
} = useUpdateMCPSemanticFilterSettings(accessToken || "");
const [form] = Form.useForm();
const [saveSuccess, setSaveSuccess] = useState(false);
const [isDirty, setIsDirty] = useState(false);
const [embeddingModels, setEmbeddingModels] = useState<ModelGroup[]>([]);
const [loadingModels, setLoadingModels] = useState(true);
// Test section state
const [testQuery, setTestQuery] = useState("");
const [testModel, setTestModel] = useState<string>("gpt-4o");
const [testResult, setTestResult] = useState<TestResult | null>(null);
const [isTesting, setIsTesting] = useState(false);
const schema = data?.field_schema;
const values = data?.values ?? {};
useEffect(() => {
const loadEmbeddingModels = async () => {
if (!accessToken) return;
try {
setLoadingModels(true);
const models = await fetchAvailableModels(accessToken);
const embeddingOnly = models.filter((model) => model.mode === "embedding");
setEmbeddingModels(embeddingOnly);
} catch (error) {
console.error("Error fetching embedding models:", error);
} finally {
setLoadingModels(false);
}
};
loadEmbeddingModels();
}, [accessToken]);
useEffect(() => {
if (values) {
form.setFieldsValue({
enabled: values.enabled ?? false,
embedding_model: values.embedding_model ?? "text-embedding-3-small",
top_k: values.top_k ?? 10,
similarity_threshold: values.similarity_threshold ?? 0.3,
});
setIsDirty(false);
}
}, [values, form]);
const handleSave = async () => {
try {
const formValues = await form.validateFields();
updateSettings(formValues, {
onSuccess: () => {
setIsDirty(false);
setSaveSuccess(true);
setTimeout(() => setSaveSuccess(false), 3000);
NotificationManager.success(
"Settings updated successfully. Changes will be applied across all pods within 10 seconds."
);
},
onError: (error) => {
NotificationManager.fromBackend(error);
},
});
} catch (error) {
console.error("Form validation failed:", error);
}
};
const handleTest = async () => {
if (!accessToken) {
return;
}
await runSemanticFilterTest({
accessToken,
testModel,
testQuery,
setIsTesting,
setTestResult,
});
};
if (!accessToken) {
return (
<div className="p-6 text-center text-gray-500">
Please log in to configure semantic filter settings.
</div>
);
}
return (
<div style={{ width: "100%" }}>
{isLoading ? (
<Skeleton active />
) : isError ? (
<Alert
type="error"
message="Could not load MCP Semantic Filter settings"
description={error instanceof Error ? error.message : undefined}
style={{ marginBottom: 24 }}
/>
) : (
<>
<Alert
type="info"
message="Semantic Tool Filtering"
description="Filter MCP tools semantically based on query relevance. This reduces context window size and improves tool selection accuracy. Click 'Save Settings' to apply changes across all pods (takes effect within 10 seconds)."
showIcon
style={{ marginBottom: 24 }}
/>
{saveSuccess && (
<Alert
type="success"
message="Settings saved successfully"
icon={<CheckCircleOutlined />}
showIcon
closable
style={{ marginBottom: 16 }}
/>
)}
{updateError && (
<Alert
type="error"
message="Could not update settings"
description={
updateError instanceof Error ? updateError.message : undefined
}
style={{ marginBottom: 16 }}
/>
)}
<Row gutter={24}>
{/* Left Column - Settings */}
<Col xs={24} lg={12}>
<Form
form={form}
layout="vertical"
disabled={isUpdating}
onValuesChange={() => {
setIsDirty(true);
}}
>
<Card style={{ marginBottom: 16 }}>
<Form.Item
name="enabled"
label={
<Space>
<Typography.Text strong>Enable Semantic Filtering</Typography.Text>
<Tooltip title="When enabled, only the most relevant MCP tools will be included in requests based on semantic similarity">
<QuestionCircleOutlined style={{ color: "#8c8c8c" }} />
</Tooltip>
</Space>
}
valuePropName="checked"
>
<Switch disabled={isUpdating} />
</Form.Item>
<Typography.Text type="secondary" style={{ display: "block", marginTop: -16, marginBottom: 16 }}>
{schema?.properties?.enabled?.description}
</Typography.Text>
</Card>
<Card title="Configuration" style={{ marginBottom: 16 }}>
<Form.Item
name="embedding_model"
label={
<Space>
<Typography.Text strong>Embedding Model</Typography.Text>
<Tooltip title="The model used to generate embeddings for semantic matching">
<QuestionCircleOutlined style={{ color: "#8c8c8c" }} />
</Tooltip>
</Space>
}
>
<Select
options={embeddingModels.map((model) => ({
label: model.model_group,
value: model.model_group,
}))}
placeholder={loadingModels ? "Loading models..." : "Select embedding model"}
showSearch
disabled={isUpdating || loadingModels}
loading={loadingModels}
notFoundContent={
loadingModels ? "Loading..." : "No embedding models available"
}
/>
</Form.Item>
<Form.Item
name="top_k"
label={
<Space>
<Typography.Text strong>Top K Results</Typography.Text>
<Tooltip title="Maximum number of tools to return after filtering">
<QuestionCircleOutlined style={{ color: "#8c8c8c" }} />
</Tooltip>
</Space>
}
>
<InputNumber
min={1}
max={100}
style={{ width: "100%" }}
disabled={isUpdating}
/>
</Form.Item>
<Form.Item
name="similarity_threshold"
label={
<Space>
<Typography.Text strong>Similarity Threshold</Typography.Text>
<Tooltip title="Minimum similarity score (0-1) for a tool to be included">
<QuestionCircleOutlined style={{ color: "#8c8c8c" }} />
</Tooltip>
</Space>
}
>
<Slider
min={0}
max={1}
step={0.05}
marks={{
0: "0.0",
0.3: "0.3",
0.5: "0.5",
0.7: "0.7",
1: "1.0",
}}
disabled={isUpdating}
/>
</Form.Item>
</Card>
<div style={{ display: "flex", justifyContent: "flex-end", gap: 8 }}>
<Button
type="primary"
icon={<SaveOutlined />}
onClick={handleSave}
loading={isUpdating}
disabled={!isDirty}
>
Save Settings
</Button>
</div>
</Form>
</Col>
{/* Right Column - Test Configuration */}
<Col xs={24} lg={12}>
<MCPSemanticFilterTestPanel
accessToken={accessToken}
testQuery={testQuery}
setTestQuery={setTestQuery}
testModel={testModel}
setTestModel={setTestModel}
isTesting={isTesting}
onTest={handleTest}
filterEnabled={!!values.enabled}
testResult={testResult}
curlCommand={getCurlCommand(testModel, testQuery)}
/>
</Col>
</Row>
</>
)}
</div>
);
}

View file

@ -0,0 +1,164 @@
import { CodeOutlined, PlayCircleOutlined } from "@ant-design/icons";
import { Alert, Button, Card, Input, Space, Tabs, Typography } from "antd";
import ModelSelector from "@/components/common_components/ModelSelector";
import { TestResult } from "./semanticFilterTestUtils";
interface MCPSemanticFilterTestPanelProps {
accessToken: string | null;
testQuery: string;
setTestQuery: (value: string) => void;
testModel: string;
setTestModel: (value: string) => void;
isTesting: boolean;
onTest: () => void;
filterEnabled: boolean;
testResult: TestResult | null;
curlCommand: string;
}
export default function MCPSemanticFilterTestPanel({
accessToken,
testQuery,
setTestQuery,
testModel,
setTestModel,
isTesting,
onTest,
filterEnabled,
testResult,
curlCommand,
}: MCPSemanticFilterTestPanelProps) {
return (
<Card title="Test Configuration" style={{ marginBottom: 16 }}>
<Tabs
defaultActiveKey="test"
items={[
{
key: "test",
label: "Test",
children: (
<Space direction="vertical" style={{ width: "100%" }} size="large">
<div>
<Typography.Text strong style={{ display: "block", marginBottom: 8 }}>
<PlayCircleOutlined /> Test Query
</Typography.Text>
<Input.TextArea
placeholder="Enter a test query to see which tools would be selected..."
value={testQuery}
onChange={(e) => setTestQuery(e.target.value)}
rows={4}
disabled={isTesting}
/>
</div>
<div>
<ModelSelector
accessToken={accessToken || ""}
value={testModel}
onChange={setTestModel}
disabled={isTesting}
showLabel={true}
labelText="Select Model"
/>
</div>
<Button
type="primary"
icon={<PlayCircleOutlined />}
onClick={onTest}
loading={isTesting}
disabled={!testQuery || !testModel || !filterEnabled}
block
>
Test Filter
</Button>
{!filterEnabled && (
<Alert
type="warning"
message="Semantic filtering is disabled"
description="Enable semantic filtering and save settings to test the filter."
showIcon
/>
)}
{testResult && (
<div>
<Typography.Title level={5}>Results</Typography.Title>
<Alert
type="success"
message={`${testResult.selectedTools} tools selected`}
description={`Filtered from ${testResult.totalTools} available tools`}
showIcon
style={{ marginBottom: 16 }}
/>
<div>
<Typography.Text strong style={{ display: "block", marginBottom: 8 }}>
Selected Tools:
</Typography.Text>
<ul style={{ paddingLeft: 20, margin: 0 }}>
{testResult.tools.map((tool, index) => (
<li key={index} style={{ marginBottom: 4 }}>
<Typography.Text>{tool}</Typography.Text>
</li>
))}
</ul>
</div>
</div>
)}
</Space>
),
},
{
key: "api",
label: "API Usage",
children: (
<div>
<Space style={{ marginBottom: 8 }}>
<CodeOutlined />
<Typography.Text strong>API Usage</Typography.Text>
</Space>
<Typography.Text type="secondary" style={{ display: "block", marginBottom: 8 }}>
Use this curl command to test the semantic filter with your current configuration.
</Typography.Text>
<Typography.Text strong style={{ display: "block", marginBottom: 8 }}>
Response headers to check:
</Typography.Text>
<ul style={{ paddingLeft: 20, margin: "0 0 12px 0" }}>
<li>
<Typography.Text>
x-litellm-semantic-filter: shows total tools → selected tools
</Typography.Text>
<Typography.Text type="secondary" style={{ display: "block" }}>
Example: 10→3
</Typography.Text>
</li>
<li>
<Typography.Text>
x-litellm-semantic-filter-tools: CSV of selected tool names
</Typography.Text>
<Typography.Text type="secondary" style={{ display: "block" }}>
Example: wikipedia-fetch,github-search,slack-post
</Typography.Text>
</li>
</ul>
<pre
style={{
background: "#f5f5f5",
padding: 12,
borderRadius: 4,
overflow: "auto",
fontSize: 12,
margin: 0,
}}
>
{curlCommand}
</pre>
</div>
),
},
]}
/>
</Card>
);
}

View file

@ -0,0 +1,95 @@
import NotificationManager from "@/components/molecules/notifications_manager";
import { testMCPSemanticFilter } from "@/components/networking";
export interface TestResult {
totalTools: number;
selectedTools: number;
tools: string[];
}
interface FilterHeaders {
filter: string | null;
tools: string | null;
}
const parseFilterHeaders = (headers: FilterHeaders): TestResult | null => {
if (!headers.filter) {
return null;
}
const [total, selected] = headers.filter.split("->").map(Number);
const tools = headers.tools
? headers.tools.split(",").map((name) => name.trim())
: [];
return { totalTools: total, selectedTools: selected, tools };
};
export const runSemanticFilterTest = async ({
accessToken,
testModel,
testQuery,
setIsTesting,
setTestResult,
}: {
accessToken: string;
testModel: string;
testQuery: string;
setIsTesting: (value: boolean) => void;
setTestResult: (result: TestResult | null) => void;
}) => {
if (!testQuery || !testModel || !accessToken) {
NotificationManager.error("Please enter a query and select a model");
return;
}
setIsTesting(true);
setTestResult(null);
try {
const { headers } = await testMCPSemanticFilter(
accessToken,
testModel,
testQuery
);
const parsedResult = parseFilterHeaders(headers);
if (!parsedResult) {
NotificationManager.warning(
"Semantic filter is not enabled or no tools were filtered"
);
return;
}
setTestResult(parsedResult);
NotificationManager.success("Semantic filter test completed successfully");
} catch (error) {
console.error("Test failed:", error);
NotificationManager.error("Failed to test semantic filter");
} finally {
setIsTesting(false);
}
};
export const getCurlCommand = (testModel: string, testQuery: string) =>
`curl --location 'http://localhost:4000/v1/responses' \\
--header 'Content-Type: application/json' \\
--header 'Authorization: Bearer sk-1234' \\
--data '{
"model": "${testModel}",
"input": [
{
"role": "user",
"content": "${testQuery || "Your query here"}",
"type": "message"
}
],
"tools": [
{
"type": "mcp",
"server_url": "litellm_proxy",
"require_approval": "never"
}
],
"tool_choice": "required"
}'`;

View file

@ -13,6 +13,7 @@ import MCPConnect from "./mcp_connect";
import { mcpServerColumns } from "./mcp_server_columns";
import { MCPServerView } from "./mcp_server_view";
import { MCPServer, MCPServerProps, Team } from "./types";
import MCPSemanticFilterSettings from "../Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings";
const { Text: AntdText, Title: AntdTitle } = Typography;
const EDIT_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-edit-state";
@ -302,6 +303,7 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
<div className="flex">
<Tab>All Servers</Tab>
<Tab>Connect</Tab>
<Tab>Semantic Filter</Tab>
</div>
</TabList>
<TabPanels>
@ -390,6 +392,9 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID })
<TabPanel>
<MCPConnect />
</TabPanel>
<TabPanel>
<MCPSemanticFilterSettings accessToken={accessToken} />
</TabPanel>
</TabPanels>
</TabGroup>
</div>

View file

@ -5398,6 +5398,137 @@ export const updateUISettings = async (accessToken: string, settings: any) => {
}
};
export const getMCPSemanticFilterSettings = async (accessToken: string) => {
/**
* Get MCP semantic filter configuration
*/
try {
const url = proxyBaseUrl
? `${proxyBaseUrl}/get/mcp_semantic_filter_settings`
: `/get/mcp_semantic_filter_settings`;
const response = await fetch(url, {
method: "GET",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const data = await response.json();
return data;
} catch (error) {
console.error("Failed to get MCP semantic filter settings:", error);
throw error;
}
};
export const updateMCPSemanticFilterSettings = async (
accessToken: string,
settings: Record<string, any>
) => {
/**
* Update MCP semantic filter settings
* Settings will be applied across all pods within 10 seconds
*/
try {
const url = proxyBaseUrl
? `${proxyBaseUrl}/update/mcp_semantic_filter_settings`
: `/update/mcp_semantic_filter_settings`;
const response = await fetch(url, {
method: "PATCH",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(settings),
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const data = await response.json();
return data;
} catch (error) {
console.error("Failed to update MCP semantic filter settings:", error);
throw error;
}
};
export const testMCPSemanticFilter = async (
accessToken: string,
model: string,
query: string
) => {
/**
* Test MCP semantic filter by making a responses API call
* Returns both the response data and headers containing filter information
*/
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/responses` : `/v1/responses`;
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify({
model: model,
input: [
{
role: "user",
content: query,
type: "message",
},
],
tools: [
{
type: "mcp",
server_url: "litellm_proxy",
require_approval: "never",
},
],
tool_choice: "required",
}),
});
// Extract headers before checking response status
const filterHeader = response.headers.get("x-litellm-semantic-filter");
const toolsHeader = response.headers.get("x-litellm-semantic-filter-tools");
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const data = await response.json();
// Return both data and headers
return {
data,
headers: {
filter: filterHeader,
tools: toolsHeader,
},
};
} catch (error) {
console.error("Failed to test MCP semantic filter:", error);
throw error;
}
};
export const getGuardrailsList = async (accessToken: string) => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/v2/guardrails/list` : `/v2/guardrails/list`;

View file

@ -1,6 +0,0 @@
export { default as SearchTools } from './search_tools';
export { SearchToolView } from './search_tool_view';
export { default as SearchConnectionTest } from './search_connection_test';
export { SearchToolTester } from './search_tool_tester';
export * from './types';

View file

@ -1,78 +0,0 @@
import { ColumnDef } from "@tanstack/react-table";
import { SearchTool } from "./types";
import { Icon } from "@tremor/react";
import { PencilAltIcon, TrashIcon } from "@heroicons/react/outline";
export const searchToolColumns = (
onView: (searchToolId: string) => void,
onEdit: (searchToolId: string) => void,
onDelete: (searchToolId: string) => void,
availableProviders: Array<{ provider_name: string; ui_friendly_name: string }>,
): ColumnDef<SearchTool>[] => [
{
accessorKey: "search_tool_id",
header: "Search Tool ID",
cell: ({ row }) => (
<button
onClick={() => onView(row.original.search_tool_id!)}
className="font-mono text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs font-normal px-2 py-0.5 text-left w-full truncate whitespace-nowrap cursor-pointer max-w-[15ch]"
>
{row.original.search_tool_id?.slice(0, 7)}...
</button>
),
},
{
accessorKey: "search_tool_name",
header: "Name",
cell: ({ getValue }) => <span className="font-medium">{getValue() as string}</span>,
},
{
id: "provider",
header: "Provider",
cell: ({ row }) => {
const provider = row.original.litellm_params.search_provider;
const providerInfo = availableProviders.find((p) => p.provider_name === provider);
const displayName = providerInfo?.ui_friendly_name || provider;
return <span className="text-sm">{displayName}</span>;
},
},
{
header: "Created At",
accessorKey: "created_at",
sortingFn: "datetime",
cell: ({ row }) => {
const tool = row.original;
return <span className="text-xs">{tool.created_at ? new Date(tool.created_at).toLocaleDateString() : "-"}</span>;
},
},
{
header: "Updated At",
accessorKey: "updated_at",
sortingFn: "datetime",
cell: ({ row }) => {
const tool = row.original;
return <span className="text-xs">{tool.updated_at ? new Date(tool.updated_at).toLocaleDateString() : "-"}</span>;
},
},
{
id: "actions",
header: "Actions",
cell: ({ row }) => (
<div className="flex items-center gap-2">
<Icon
icon={PencilAltIcon}
size="sm"
onClick={() => onEdit(row.original.search_tool_id!)}
className="cursor-pointer"
/>
<Icon
icon={TrashIcon}
size="sm"
onClick={() => onDelete(row.original.search_tool_id!)}
className="cursor-pointer"
/>
</div>
),
},
];