diff --git a/docs/my-website/docs/observability/langfuse_integration.md b/docs/my-website/docs/observability/langfuse_integration.md index a81336c5bc6..d3c5a44d481 100644 --- a/docs/my-website/docs/observability/langfuse_integration.md +++ b/docs/my-website/docs/observability/langfuse_integration.md @@ -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. diff --git a/docs/my-website/docs/proxy/admin_ui_sso.md b/docs/my-website/docs/proxy/admin_ui_sso.md index 7b299429db7..37e45b50284 100644 --- a/docs/my-website/docs/proxy/admin_ui_sso.md +++ b/docs/my-website/docs/proxy/admin_ui_sso.md @@ -23,26 +23,75 @@ From v1.76.0, SSO is now Free for up to 5 users. -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:///sso/callback` +- **Sign-out redirect URI** (optional): `https://` + + + +After creating the app, copy your **Client ID** and **Client Secret** from the application's General tab: + + + +#### 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** + + + +2. Select the **default** authorization server (or your custom one) + + + +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 + + + +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 = "" -GENERIC_CLIENT_SECRET = "" -GENERIC_AUTHORIZATION_ENDPOINT = "/authorize" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/authorize -GENERIC_TOKEN_ENDPOINT = "/token" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/oauth/token -GENERIC_USERINFO_ENDPOINT = "/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="" +GENERIC_CLIENT_SECRET="" +GENERIC_AUTHORIZATION_ENDPOINT="https:///oauth2/default/v1/authorize" +GENERIC_TOKEN_ENDPOINT="https:///oauth2/default/v1/token" +GENERIC_USERINFO_ENDPOINT="https:///oauth2/default/v1/userinfo" +GENERIC_CLIENT_STATE="random-string" +PROXY_BASE_URL="https://" ``` -You can get your domain specific auth/token/userinfo endpoints at `/.well-known/openid-configuration` +:::tip +You can find all OAuth endpoints at `https:///.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 `/sso/callback` +1. Start your LiteLLM proxy +2. Navigate to `https:///ui` +3. Click the SSO login button +4. Authenticate with Okta and verify you're redirected back to LiteLLM +#### Troubleshooting - +| Error | Cause | Solution | +|-------|-------|----------| +| `redirect_uri` error | Redirect URI not configured | Add `/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) | diff --git a/docs/my-website/img/okta_access_policies.png b/docs/my-website/img/okta_access_policies.png new file mode 100644 index 00000000000..e09adc2ce7f Binary files /dev/null and b/docs/my-website/img/okta_access_policies.png differ diff --git a/docs/my-website/img/okta_authorization_server.png b/docs/my-website/img/okta_authorization_server.png new file mode 100644 index 00000000000..bddb3e07a4a Binary files /dev/null and b/docs/my-website/img/okta_authorization_server.png differ diff --git a/docs/my-website/img/okta_client_credentials.png b/docs/my-website/img/okta_client_credentials.png new file mode 100644 index 00000000000..a00a9f4657e Binary files /dev/null and b/docs/my-website/img/okta_client_credentials.png differ diff --git a/docs/my-website/img/okta_redirect_uri.png b/docs/my-website/img/okta_redirect_uri.png new file mode 100644 index 00000000000..a1e58560c72 Binary files /dev/null and b/docs/my-website/img/okta_redirect_uri.png differ diff --git a/docs/my-website/img/okta_security_api.png b/docs/my-website/img/okta_security_api.png new file mode 100644 index 00000000000..7f9e218074c Binary files /dev/null and b/docs/my-website/img/okta_security_api.png differ diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30-py3-none-any.whl b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30-py3-none-any.whl new file mode 100644 index 00000000000..383f9b7b43f Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30-py3-none-any.whl differ diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30.tar.gz b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30.tar.gz new file mode 100644 index 00000000000..484c28ba7b1 Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30.tar.gz differ diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260205091235_allow_team_guardrail_config/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260205091235_allow_team_guardrail_config/migration.sql new file mode 100644 index 00000000000..000b96b3b87 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260205091235_allow_team_guardrail_config/migration.sql @@ -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; + diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index fb6996b71db..d43b591686c 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -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==", diff --git a/litellm/constants.py b/litellm/constants.py index 444e78f8ed4..872ad899f84 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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) diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index 4f76a5bad03..435ae078a65 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -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]] ): diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 444f821c20a..169b138a5f7 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -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 diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index f6ea9c57f77..7ec32fecc46 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -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 diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 506798bb8d6..be8ad7d0877 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -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, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 6d399d25bf6..d3038d13a1e 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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", diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 8e4c705df21..cb277d44ee9 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -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. diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 429e56c805b..dc928921425 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b06071481eb..637893872d9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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. diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 39ef5fdd1d5..92fb88f7147 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 30ec0766dbf..a308d0b2703 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -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"], diff --git a/litellm/proxy_auth/credentials.py b/litellm/proxy_auth/credentials.py index cddaf1278f9..103b0088d80 100644 --- a/litellm/proxy_auth/credentials.py +++ b/litellm/proxy_auth/credentials.py @@ -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: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 6d399d25bf6..d3038d13a1e 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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", diff --git a/poetry.lock b/poetry.lock index 537367c5aa0..674725c2ce4 100644 --- a/poetry.lock +++ b/poetry.lock @@ -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" diff --git a/pyproject.toml b/pyproject.toml index 9832ca483dc..0d0968ee8ee 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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" ] diff --git a/requirements.txt b/requirements.txt index 69768b6c1f7..aca27bf284b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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 diff --git a/schema.prisma b/schema.prisma index b118400b620..03b910db85e 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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") diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py index b334966b441..2ebf9174e29 100644 --- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py @@ -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): """ diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index 17e72f29152..0a581fb512d 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -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 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 6a5c022ac7a..b228a51447b 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -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 diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 72403b0ba7b..6ccecf59eed 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -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" \ No newline at end of file + 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", + ) \ No newline at end of file diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index e8d89c0c7eb..1288a9b2c9f 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -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( diff --git a/tests/test_team.py b/tests/test_team.py index d67c5e670f4..275181590c5 100644 --- a/tests/test_team.py +++ b/tests/test_team.py @@ -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: diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings.ts new file mode 100644 index 00000000000..e91f5aa670b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings.ts @@ -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>({ + queryKey: mcpSemanticFilterSettingsKeys.list({}), + queryFn: async () => await getMCPSemanticFilterSettings(accessToken), + enabled: !!accessToken, + staleTime: 60 * 60 * 1000, // 1 hour + gcTime: 60 * 60 * 1000, // 1 hour + }); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings.ts new file mode 100644 index 00000000000..2062b4f4c29 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings.ts @@ -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) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return updateMCPSemanticFilterSettings(accessToken, settings); + }, + onSuccess: () => { + queryClient.invalidateQueries({ + queryKey: mcpSemanticFilterSettingsKeys.all, + }); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index 5f94db7e9f4..8e887d80fb3 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -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"; diff --git a/ui/litellm-dashboard/src/components/search_tools/create_search_tool.tsx b/ui/litellm-dashboard/src/components/SearchTools/CreateSearchTools.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/search_tools/create_search_tool.tsx rename to ui/litellm-dashboard/src/components/SearchTools/CreateSearchTools.tsx index 9dd14256eee..49b57c2a884 100644 --- a/ui/litellm-dashboard/src/components/search_tools/create_search_tool.tsx +++ b/ui/litellm-dashboard/src/components/SearchTools/CreateSearchTools.tsx @@ -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 = ({ }, search_tool_info: formValues.description ? { - description: formValues.description, - } + description: formValues.description, + } : undefined, }; @@ -130,7 +130,7 @@ const CreateSearchTool: React.FC = ({ 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 = ({ optionLabelProp="label" > {availableProviders.map((provider) => ( - void, + onEdit: (searchToolId: string) => void, + onDelete: (searchToolId: string) => void, + availableProviders: Array<{ provider_name: string; ui_friendly_name: string }>, +): ColumnsType => [ + { + title: "Search Tool ID", + dataIndex: "search_tool_id", + key: "search_tool_id", + render: (_, tool) => { + const isFromConfig = tool.is_from_config; + + if (isFromConfig) { + return -; + } + + return ( + + ); + }, + }, + { + title: "Name", + dataIndex: "search_tool_name", + key: "search_tool_name", + render: (name: string) => {name}, + }, + { + 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 {displayName}; + }, + }, + { + title: "Created At", + dataIndex: "created_at", + key: "created_at", + render: (_, tool) => { + return {tool.created_at ? new Date(tool.created_at).toLocaleDateString() : "-"}; + }, + }, + { + title: "Updated At", + dataIndex: "updated_at", + key: "updated_at", + render: (_, tool) => { + return {tool.updated_at ? new Date(tool.updated_at).toLocaleDateString() : "-"}; + }, + }, + { + title: "Source", + key: "source", + render: (_, tool) => { + const isFromConfig = tool.is_from_config ?? false; + + return ( + + {isFromConfig ? "Config" : "DB"} + + ); + }, + }, + { + title: "Actions", + key: "actions", + render: (_, tool) => { + const toolId = tool.search_tool_id; + const isFromConfig = tool.is_from_config ?? false; + + return ( +
+ { + if (toolId && !isFromConfig) { + onEdit(toolId); + } + }} + /> + { + if (toolId && !isFromConfig) { + onDelete(toolId); + } + }} + /> +
+ ); + }, + }, + ]; diff --git a/ui/litellm-dashboard/src/components/search_tools/search_tool_tester.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchToolTester.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/search_tools/search_tool_tester.tsx rename to ui/litellm-dashboard/src/components/SearchTools/SearchToolTester.tsx diff --git a/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.test.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.test.tsx new file mode 100644 index 00000000000..bb04a04a992 --- /dev/null +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.test.tsx @@ -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 }) => ( +
+ Search Tool Tester for {searchToolName} + Access Token: {accessToken} +
+ ), +})); + +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(); + expect(screen.getByText("Test Search Tool")).toBeInTheDocument(); + }); + + it("should display search tool name", () => { + render(); + expect(screen.getByText("Test Search Tool")).toBeInTheDocument(); + }); + + it("should display search tool ID", () => { + render(); + expect(screen.getByText("test-tool-id-123")).toBeInTheDocument(); + }); + + it("should display provider name using UI-friendly name when available", () => { + render(); + 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( + , + ); + expect(screen.getByText("unknown-provider")).toBeInTheDocument(); + }); + + it("should display masked API key when API key is set", () => { + render(); + 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( + , + ); + expect(screen.getByText("Not set")).toBeInTheDocument(); + }); + + it("should display formatted created_at date", () => { + render(); + 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( + , + ); + expect(screen.getByText("Unknown")).toBeInTheDocument(); + }); + + it("should display description when search_tool_info.description is provided", () => { + render(); + 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( + , + ); + 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(); + + 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(); + + 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(); + + 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(); + + 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(); + + 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(); + 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(); + expect(screen.queryByTestId("search-tool-tester")).not.toBeInTheDocument(); + }); + + it("should pass correct props to SearchToolTester", () => { + render(); + expect(screen.getByText("Access Token: test-token")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/search_tools/search_tool_view.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.tsx similarity index 86% rename from ui/litellm-dashboard/src/components/search_tools/search_tool_view.tsx rename to ui/litellm-dashboard/src/components/SearchTools/SearchToolView.tsx index a5cad4e8370..ad88acd127f 100644 --- a/ui/litellm-dashboard/src/components/search_tools/search_tool_view.tsx +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.tsx @@ -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 = ({ size="small" icon={copiedStates["search-tool-name"] ? : } 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" + }`} />
@@ -67,11 +66,10 @@ export const SearchToolView: React.FC = ({ size="small" icon={copiedStates["search-tool-id"] ? : } 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" + }`} />
diff --git a/ui/litellm-dashboard/src/components/SearchTools/SearchTools.test.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.test.tsx new file mode 100644 index 00000000000..f1f7bf8ab42 --- /dev/null +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.test.tsx @@ -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 }) => ( +
+
Search Tool View: {searchTool.search_tool_name}
+ +
+ ), +})); + +vi.mock("./CreateSearchTools", () => ({ + default: ({ + isModalVisible, + setModalVisible, + }: { + isModalVisible: boolean; + setModalVisible: (visible: boolean) => void; + }) => + isModalVisible ? ( +
+ +
+ ) : null, +})); + +vi.mock("../common_components/DeleteResourceModal", () => ({ + default: ({ + isOpen, + onOk, + onCancel, + }: { + isOpen: boolean; + onOk: () => void; + onCancel: () => void; + }) => + isOpen ? ( +
+ + +
+ ) : 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 }) => ( + {children} + ); +}; + +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(, { wrapper: createWrapper() }); + await waitFor(() => { + expect(screen.getByText("Search Tools")).toBeInTheDocument(); + }); + }); + + it("should display missing authentication parameters message when accessToken is missing", () => { + render(, { wrapper: createWrapper() }); + expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument(); + }); + + it("should display missing authentication parameters message when userRole is missing", () => { + render(, { wrapper: createWrapper() }); + expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument(); + }); + + it("should display missing authentication parameters message when userID is missing", () => { + render(, { wrapper: createWrapper() }); + expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument(); + }); + + it("should display search tools table with tools", async () => { + render(, { 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(, { 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(, { 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(, { 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(, { 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(, { 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(, { 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(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/search_tools/search_tools.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.tsx similarity index 78% rename from ui/litellm-dashboard/src/components/search_tools/search_tools.tsx rename to ui/litellm-dashboard/src/components/SearchTools/SearchTools.tsx index 2fbdd4d27c6..dd2033fc18f 100644 --- a/ui/litellm-dashboard/src/components/search_tools/search_tools.tsx +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.tsx @@ -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 ( - - - {title} - -

Are you sure you want to delete this search tool?

- -
-
- ); -}; const SearchTools: React.FC = ({ accessToken, userRole, userID }) => { const { @@ -72,6 +55,7 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID // State const [toolIdToDelete, setToolToDelete] = useState(null); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); + const [isDeleting, setIsDeleting] = useState(false); const [selectedToolId, setSelectedToolId] = useState(null); const [editTool, setEditTool] = useState(false); const [isCreateModalVisible, setCreateModalVisible] = useState(false); @@ -116,16 +100,19 @@ const SearchTools: React.FC = ({ 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 = ({ 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 = ({ accessToken, userRole, userID /> ) : (
-
- } size="large"> +
} - 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" /> - + + ); return (
- ([]); + const [loadingModels, setLoadingModels] = useState(true); + + // Test section state + const [testQuery, setTestQuery] = useState(""); + const [testModel, setTestModel] = useState("gpt-4o"); + const [testResult, setTestResult] = useState(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 ( +
+ Please log in to configure semantic filter settings. +
+ ); + } + + return ( +
+ {isLoading ? ( + + ) : isError ? ( + + ) : ( + <> + + + {saveSuccess && ( + } + showIcon + closable + style={{ marginBottom: 16 }} + /> + )} + + {updateError && ( + + )} + + + {/* Left Column - Settings */} +
+ { + setIsDirty(true); + }} + > + + + Enable Semantic Filtering + + + + + } + valuePropName="checked" + > + + + + + {schema?.properties?.enabled?.description} + + + + + + Embedding Model + + + + + } + > +