diff --git a/.circleci/requirements.txt b/.circleci/requirements.txt index 8c44dc18305..a5ec74424fe 100644 --- a/.circleci/requirements.txt +++ b/.circleci/requirements.txt @@ -16,4 +16,5 @@ uvloop==0.21.0 mcp==1.25.0 # for MCP server semantic_router==0.1.10 # for auto-routing with litellm fastuuid==0.12.0 -responses==0.25.7 # for proxy client tests \ No newline at end of file +responses==0.25.7 # for proxy client tests +pytest-retry==1.6.3 # for automatic test retries \ No newline at end of file diff --git a/AGENTS.md b/AGENTS.md index 61afbd035fe..5a48049ef45 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -51,12 +51,14 @@ LiteLLM is a unified interface for 100+ LLMs that: ### MAKING CODE CHANGES FOR THE UI (IGNORE FOR BACKEND) -1. **Use Common Components as much as possible**: +1. **Tremor is DEPRECATED, do not use Tremor components in new features/changes** + - The only exception is the Tremor Table component and its required Tremor Table sub components. + +2. **Use Common Components as much as possible**: - These are usually defined in the `common_components` directory - Use these components as much as possible and avoid building new components unless needed - - Tremor components are deprecated; prefer using Ant Design (AntD) as much as possible -2. **Testing**: +3. **Testing**: - The codebase uses **Vitest** and **React Testing Library** - **Query Priority Order**: Use query methods in this order: `getByRole`, `getByLabelText`, `getByPlaceholderText`, `getByText`, `getByTestId` - **Always use `screen`** instead of destructuring from `render()` (e.g., use `screen.getByText()` not `getByText`) diff --git a/Dockerfile b/Dockerfile index 0e7a8412bbc..2c54e2dec28 100644 --- a/Dockerfile +++ b/Dockerfile @@ -69,8 +69,8 @@ RUN find /usr/lib -type f -path "*/tornado/test/*" -delete && \ # Convert Windows line endings to Unix and make executable RUN sed -i 's/\r$//' docker/install_auto_router.sh && chmod +x docker/install_auto_router.sh && ./docker/install_auto_router.sh -# Generate prisma client -RUN prisma generate +# Generate prisma client using the correct schema +RUN prisma generate --schema=./litellm/proxy/schema.prisma # Convert Windows line endings to Unix for entrypoint scripts RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh diff --git a/README.md b/README.md index 914fda384b0..77adddf8978 100644 --- a/README.md +++ b/README.md @@ -267,6 +267,7 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature Greptile OpenHands

Netflix

+ OpenAI Agents SDK diff --git a/docs/my-website/docs/a2a.md b/docs/my-website/docs/a2a.md index d7145e4b83c..a7e8b52d99a 100644 --- a/docs/my-website/docs/a2a.md +++ b/docs/my-website/docs/a2a.md @@ -68,7 +68,7 @@ Follow [this guide, to add your pydantic ai agent to LiteLLM Agent Gateway](./pr ## Invoking your Agents -Use the [A2A Python SDK](https://pypi.org/project/a2a/) to invoke agents through LiteLLM. +Use the [A2A Python SDK](https://pypi.org/project/a2a-sdk) to invoke agents through LiteLLM. This example shows how to: 1. **List available agents** - Query `/v1/agents` to see which agents your key can access @@ -193,6 +193,120 @@ The logs show: style={{width: '100%', display: 'block', margin: '2rem auto'}} /> + +## Forwarding LiteLLM Context Headers + +When LiteLLM invokes your A2A agent, it sends special headers that enable: +- **Trace Grouping**: All LLM calls from the same agent execution appear under one trace +- **Agent Spend Tracking**: Costs are attributed to the specific agent + +| Header | Purpose | +|--------|---------| +| `X-LiteLLM-Trace-Id` | Links all LLM calls to the same execution flow | +| `X-LiteLLM-Agent-Id` | Attributes spend to the correct agent | + + +To enable these features, your A2A server must **forward these headers** to any LLM calls it makes back to LiteLLM. + +### Implementation Steps + +**Step 1: Extract headers from incoming A2A request** +```python def get_litellm_headers(request) -> dict: + """Extract X-LiteLLM-* headers from incoming A2A request.""" + all_headers = request.call_context.state.get('headers', {}) + return { + k: v for k, v in all_headers.items() + if k.lower().startswith('x-litellm-') + } +``` + +**Step 2: Forward headers to your LLM calls** +Pass the extracted headers when making calls back to LiteLLM: + + + +```python from openai import OpenAI + +headers = get_litellm_headers(request) + +client = OpenAI( + api_key="sk-your-litellm-key", + base_url="http://localhost:4000", + default_headers=headers, # Forward headers +) + +response = client.chat.completions.create( + model="gpt-4o", + messages=[{"role": "user", "content": "Hello"}] +) +``` + + + + +```python +from langchain_openai import ChatOpenAI + +headers = get_litellm_headers(request) + +llm = ChatOpenAI( + model="gpt-4o", + openai_api_key="sk-your-litellm-key", + base_url="http://localhost:4000", + default_headers=headers, # Forward headers +) +``` + + + +```python +import litellm + +headers = get_litellm_headers(request) + +response = litellm.completion( + model="gpt-4o", + messages=[{"role": "user", "content": "Hello"}], + api_base="http://localhost:4000", + extra_headers=headers, # Forward headers +) +``` + + + +```python +import httpx + +headers = get_litellm_headers(request) +headers["Authorization"] = "Bearer sk-your-litellm-key" + +response = httpx.post( + "http://localhost:4000/v1/chat/completions", + headers=headers, + json={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]} +) +``` + + + +### Result + +With header forwarding enabled, you'll see: + +**Trace Grouping in Langfuse:** + + + +**Agent Spend Attribution:** + + + ## API Reference ### Endpoint diff --git a/docs/my-website/docs/observability/datadog.md b/docs/my-website/docs/observability/datadog.md index 7cf91ced34c..6f785be1013 100644 --- a/docs/my-website/docs/observability/datadog.md +++ b/docs/my-website/docs/observability/datadog.md @@ -7,6 +7,7 @@ import TabItem from '@theme/TabItem'; LiteLLM Supports logging to the following Datdog Integrations: - `datadog` [Datadog Logs](https://docs.datadoghq.com/logs/) - `datadog_llm_observability` [Datadog LLM Observability](https://www.datadoghq.com/product/llm-observability/) +- `datadog_cost_management` [Datadog Cloud Cost Management](#datadog-cloud-cost-management) - `ddtrace-run` [Datadog Tracing](#datadog-tracing) ## Datadog Logs @@ -73,7 +74,7 @@ Send logs through a local DataDog agent (useful for containerized environments): ```shell LITELLM_DD_AGENT_HOST="localhost" # hostname or IP of DataDog agent LITELLM_DD_AGENT_PORT="10518" # [OPTIONAL] port of DataDog agent (default: 10518) -DD_API_KEY="5f2d0f310***********" # [OPTIONAL] your datadog API Key (agent handles auth) +DD_API_KEY="5f2d0f310***********" # [OPTIONAL] your datadog API Key (Agent handles auth for Logs. REQUIRED for LLM Observability) DD_SOURCE="litellm_dev" # [OPTIONAL] your datadog source ``` @@ -84,6 +85,9 @@ When `LITELLM_DD_AGENT_HOST` is set, logs are sent to the agent instead of direc **Note:** We use `LITELLM_DD_AGENT_HOST` instead of `DD_AGENT_HOST` to avoid conflicts with `ddtrace` which automatically sets `DD_AGENT_HOST` for APM tracing. +> [!IMPORTANT] +> **Datadog LLM Observability**: `DD_API_KEY` is **REQUIRED** even when using the Datadog Agent (`LITELLM_DD_AGENT_HOST`). The agent acts as a proxy but the API key header is mandatory for the LLM Observability endpoint. + **Step 3**: Start the proxy, make a test request Start proxy @@ -161,6 +165,50 @@ On the Datadog LLM Observability page, you should see that both input messages a + + + +## Datadog Cloud Cost Management + +| Feature | Details | +|---------|---------| +| **What is logged** | Aggregated LLM Costs (FOCUS format) | +| **Events** | Periodic Uploads of Aggregated Cost Data | +| **Product Link** | [Datadog Cloud Cost Management](https://docs.datadoghq.com/cost_management/) | + +We will use the `--config` to set `litellm.callbacks = ["datadog_cost_management"]`. This will periodically upload aggregated LLM cost data to Datadog. + +**Step 1**: Create a `config.yaml` file and set `litellm_settings`: `success_callback` + +```yaml +model_list: + - model_name: gpt-3.5-turbo + litellm_params: + model: gpt-3.5-turbo +litellm_settings: + callbacks: ["datadog_cost_management"] +``` + +**Step 2**: Set Required env variables + +```shell +DD_API_KEY="your-api-key" +DD_APP_KEY="your-app-key" # REQUIRED for Cost Management +DD_SITE="us5.datadoghq.com" +``` + +**Step 3**: Start the proxy + +```shell +litellm --config config.yaml +``` + +**How it works** +* LiteLLM aggregates costs in-memory by Provider, Model, Date, and Tags. +* Requires `DD_APP_KEY` for the Custom Costs API. +* Costs are uploaded periodically (flushed). + + ### Datadog Tracing Use `ddtrace-run` to enable [Datadog Tracing](https://ddtrace.readthedocs.io/en/stable/installation_quickstart.html) on litellm proxy @@ -203,5 +251,5 @@ LiteLLM supports customizing the following Datadog environment variables | `POD_NAME` | Pod name tag (useful for Kubernetes deployments) | "unknown" | ❌ No | \* **Required when using Direct API** (default): `DD_API_KEY` and `DD_SITE` are required -\* **Optional when using DataDog Agent**: Set `LITELLM_DD_AGENT_HOST` to use agent mode; `DD_API_KEY` and `DD_SITE` are not required +\* **Optional when using DataDog Agent**: Set `LITELLM_DD_AGENT_HOST` to use agent mode; `DD_API_KEY` and `DD_SITE` are not required for **Datadog Logs**. (**Note: `DD_API_KEY` IS REQUIRED for Datadog LLM Observability**) diff --git a/docs/my-website/docs/providers/vercel_ai_gateway.md b/docs/my-website/docs/providers/vercel_ai_gateway.md index 91f0a18ea1c..3ff007171ed 100644 --- a/docs/my-website/docs/providers/vercel_ai_gateway.md +++ b/docs/my-website/docs/providers/vercel_ai_gateway.md @@ -11,7 +11,7 @@ import TabItem from '@theme/TabItem'; | Provider Route on LiteLLM | `vercel_ai_gateway/` | | Link to Provider Doc | [Vercel AI Gateway Documentation ↗](https://vercel.com/docs/ai-gateway) | | Base URL | `https://ai-gateway.vercel.sh/v1` | -| Supported Operations | `/chat/completions`, `/models` | +| Supported Operations | `/chat/completions`, `/embeddings`, `/models` |

@@ -73,7 +73,7 @@ messages = [{"content": "Hello, how are you?", "role": "user"}] # Vercel AI Gateway call with streaming response = completion( - model="vercel_ai_gateway/openai/gpt-4o", + model="vercel_ai_gateway/openai/gpt-4o", messages=messages, stream=True ) @@ -82,6 +82,33 @@ for chunk in response: print(chunk) ``` +### Embeddings + +```python showLineNumbers title="Vercel AI Gateway Embeddings" +import os +from litellm import embedding + +os.environ["VERCEL_AI_GATEWAY_API_KEY"] = "your-api-key" + +# Vercel AI Gateway embedding call +response = embedding( + model="vercel_ai_gateway/openai/text-embedding-3-small", + input="Hello world" +) + +print(response.data[0]["embedding"][:5]) # Print first 5 dimensions +``` + +You can also specify the `dimensions` parameter: + +```python showLineNumbers title="Vercel AI Gateway Embeddings with Dimensions" +response = embedding( + model="vercel_ai_gateway/openai/text-embedding-3-small", + input=["Hello world", "Goodbye world"], + dimensions=768 +) +``` + ## Usage - LiteLLM Proxy Add the following to your LiteLLM Proxy configuration file: @@ -97,6 +124,11 @@ model_list: litellm_params: model: vercel_ai_gateway/anthropic/claude-4-sonnet api_key: os.environ/VERCEL_AI_GATEWAY_API_KEY + + - model_name: text-embedding-3-small-gateway + litellm_params: + model: vercel_ai_gateway/openai/text-embedding-3-small + api_key: os.environ/VERCEL_AI_GATEWAY_API_KEY ``` Start your LiteLLM Proxy server: diff --git a/docs/my-website/docs/proxy/cli_sso.md b/docs/my-website/docs/proxy/cli_sso.md index cde6bf266d4..ad0f033f802 100644 --- a/docs/my-website/docs/proxy/cli_sso.md +++ b/docs/my-website/docs/proxy/cli_sso.md @@ -28,6 +28,37 @@ EXPERIMENTAL_UI_LOGIN="True" litellm --config config.yaml ::: +### Configuration + +#### JWT Token Expiration + +By default, CLI authentication tokens expire after **24 hours**. You can customize this expiration time by setting the `LITELLM_CLI_JWT_EXPIRATION_HOURS` environment variable when starting your LiteLLM Proxy: + +```bash +# Set CLI JWT tokens to expire after 48 hours +export LITELLM_CLI_JWT_EXPIRATION_HOURS=48 +export EXPERIMENTAL_UI_LOGIN="True" +litellm --config config.yaml +``` + +Or in a single command: + +```bash +LITELLM_CLI_JWT_EXPIRATION_HOURS=48 EXPERIMENTAL_UI_LOGIN="True" litellm --config config.yaml +``` + +**Examples:** +- `LITELLM_CLI_JWT_EXPIRATION_HOURS=12` - Tokens expire after 12 hours +- `LITELLM_CLI_JWT_EXPIRATION_HOURS=168` - Tokens expire after 7 days (168 hours) +- `LITELLM_CLI_JWT_EXPIRATION_HOURS=720` - Tokens expire after 30 days (720 hours) + +:::tip +You can check your current token's age and expiration status using: +```bash +litellm-proxy whoami +``` +::: + ### Steps 1. **Install the CLI** diff --git a/docs/my-website/img/a2a_agent_spend.png b/docs/my-website/img/a2a_agent_spend.png new file mode 100644 index 00000000000..15ec769392a Binary files /dev/null and b/docs/my-website/img/a2a_agent_spend.png differ diff --git a/docs/my-website/img/a2a_trace_grouping.png b/docs/my-website/img/a2a_trace_grouping.png new file mode 100644 index 00000000000..05130420aae Binary files /dev/null and b/docs/my-website/img/a2a_trace_grouping.png differ diff --git a/docs/my-website/release_notes/v1.81.3-stable/index.md b/docs/my-website/release_notes/v1.81.3-stable/index.md new file mode 100644 index 00000000000..22b6f43deef --- /dev/null +++ b/docs/my-website/release_notes/v1.81.3-stable/index.md @@ -0,0 +1,423 @@ +--- +title: "v1.81.3-stable - Performance - 25% CPU Usage Reduction" +slug: "v1-81-3" +date: 2026-01-26T10:00:00 +authors: + - name: Krrish Dholakia + title: CEO, LiteLLM + url: https://www.linkedin.com/in/krish-d/ + image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg + - name: Ishaan Jaff + title: CTO, LiteLLM + url: https://www.linkedin.com/in/reffajnaahsi/ + image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg +hide_table_of_contents: false +--- + +import Image from '@theme/IdealImage'; +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +## Deploy this version + + + + +``` showLineNumbers title="docker run litellm" +docker run \ +-e STORE_MODEL_IN_DB=True \ +-p 4000:4000 \ +docker.litellm.ai/berriai/litellm:v1.81.3.rc.2 +``` + + + + + +``` showLineNumbers title="pip install litellm" +pip install litellm==1.81.3.rc.2 +``` + + + + +--- + +## New Models / Updated Models + +### New Model Support + +| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) | Deprecation Date | +| -------- | ----- | -------------- | ------------------- | -------------------- | ---------------- | +| OpenAI | `gpt-audio`, `gpt-audio-2025-08-28` | 128K | $32/1M audio tokens, $2.5/1M text tokens | $64/1M audio tokens, $10/1M text tokens | - | +| OpenAI | `gpt-audio-mini`, `gpt-audio-mini-2025-08-28` | 128K | $10/1M audio tokens, $0.6/1M text tokens | $20/1M audio tokens, $2.4/1M text tokens | - | +| Deepinfra, Vertex AI, Google AI Studio, OpenRouter, Vercel AI Gateway | `gemini-2.0-flash-001`, `gemini-2.0-flash` | - | - | - | 2026-03-31 | +| Groq | `openai/gpt-oss-120b` | 131K | 0.075/1M cache read | 0.6/1M output tokens | - | +| Groq | `groq/openai/gpt-oss-20b` | 131K | 0.0375/1M cache read, $0.075/1M text tokens | 0.3/1M output tokens | - | +| Vertex AI | `gemini-2.5-computer-use-preview-10-2025` | 128K | $1.25 | $10 | - | +| Azure AI | `claude-haiku-4-5` | $1.25/1M cache read, $2/1M cache read above 1 hr, $0.1/1M text tokens | $5/1M output tokens | - | +| Azure AI | `claude-sonnet-4-5` | $3.75/1M cache read, $6/1M cache read above 1 hr, $3/1M text tokens | $15/1M output tokens | - | +| Azure AI | `claude-opus-4-5` | $6.25/1M cache read, $10/1M cache read above 1 hr, $0.5/1M text tokens | $25/1M output tokens | - | +| Azure AI | `claude-opus-4-1` | $18.75/1M cache read, $30/1M cache read above 1 hr, $1.5/1M text tokens | $75/1M output tokens | - | + +### Features + +- **[OpenAI](../../docs/providers/openai)** + - Add gpt-audio and gpt-audio-mini models to pricing - [PR #19509](https://github.com/BerriAI/litellm/pull/19509) + - correct audio token costs for gpt-4o-audio-preview models - [PR #19500](https://github.com/BerriAI/litellm/pull/19500) + - Limit stop sequence as per openai spec (ensures JetBrains IDE compatibility) - [PR #19562](https://github.com/BerriAI/litellm/pull/19562) + +- **[VertexAI](../../docs/providers/vertex)** + - Docs - Google Workload Identity Federation (WIF) support - [PR #19320](https://github.com/BerriAI/litellm/pull/19320) + +- **[Agentcore](../../docs/providers/bedrock_agentcore)** + - Fixes streaming issues with AWS Bedrock AgentCore where responses would stop after the first chunk, particularly affecting OAuth-enabled agents - [PR #17141](https://github.com/BerriAI/litellm/pull/17141) + +- **[Chatgpt](../../docs/providers/chatgpt)** + - Adds support for calling chatgpt subscription via LiteLLM - [PR #19030](https://github.com/BerriAI/litellm/pull/19030) + - Adds responses API bridge support for chatgpt subscription provider - [PR #19030](https://github.com/BerriAI/litellm/pull/19030) + +- **[Bedrock](../../docs/providers/bedrock)** + - support for output format for bedrock invoke via v1/messages - [PR #19560](https://github.com/BerriAI/litellm/pull/19560) + +- **[Azure](../../docs/providers/azure/azure)** + - Add support for Azure OpenAI v1 API - [PR #19313](https://github.com/BerriAI/litellm/pull/19313) + - preserve content_policy_violation details for images (#19328) - [PR #19372](https://github.com/BerriAI/litellm/pull/19372) + - Support OpenAI-format nested tool definitions for Responses API - [PR #19526](https://github.com/BerriAI/litellm/pull/19526) + +- **Gemini([Vertex AI](../../docs/providers/vertex), [Google AI Studio](../../docs/providers/gemini))** + - use responseJsonSchema for Gemini 2.0+ models - [PR #19314](https://github.com/BerriAI/litellm/pull/19314) + +- **[Volcengine](../../docs/providers/volcano)** + - Support Volcengine responses api - [PR #18508](https://github.com/BerriAI/litellm/pull/18508) + +- **[Anthropic](../../docs/providers/anthropic)** + - Add Support for calling Claude Code Max subscriptions via LiteLLM - [PR #19453](https://github.com/BerriAI/litellm/pull/19453) + - Add Structured output for /v1/messages with Anthropic API, Azure Anthropic API, Bedrock Converse - [PR #19545](https://github.com/BerriAI/litellm/pull/19545) + +- **[Brave Search](../../docs/search/brave)** + - New Search provider - [PR #19433](https://github.com/BerriAI/litellm/pull/19433) + +- **Sarvam ai** + - Add support for new sarvam models - [PR #19479](https://github.com/BerriAI/litellm/pull/19479) + +- **[GMI](../../docs/providers/gmi)** + - add GMI Cloud provider support - [PR #19376](https://github.com/BerriAI/litellm/pull/19376) + + +### Bug Fixes + +- **[Anthropic](../../docs/providers/anthropic)** + - Fix anthropic-beta sent client side being overridden instead of appended to - [PR #19343](https://github.com/BerriAI/litellm/pull/19343) + - Filter out unsupported fields from JSON schema for Anthropic's output_format API - [PR #19482](https://github.com/BerriAI/litellm/pull/19482) + +- **[Bedrock](../../docs/providers/bedrock)** + - Expose stability models via /image_edits endpoint and ensure proper request transformation - [PR #19323](https://github.com/BerriAI/litellm/pull/19323) + - Claude Code x Bedrock Invoke fails with advanced-tool-use-2025-11-20 - [PR #19373](https://github.com/BerriAI/litellm/pull/19373) + - deduplicate tool calls in assistant history - [PR #19324](https://github.com/BerriAI/litellm/pull/19324) + - fix: correct us.anthropic.claude-opus-4-5 In-region pricing - [PR #19310](https://github.com/BerriAI/litellm/pull/19310) + - Fix request validation errors when using Claude 4 via bedrock invoke - [PR #19381](https://github.com/BerriAI/litellm/pull/19381) + - Handle thinking with tool calls for Claude 4 models - [PR #19506](https://github.com/BerriAI/litellm/pull/19506) + - correct streaming choice index for tool calls - [PR #19506](https://github.com/BerriAI/litellm/pull/19506) + +- **[Ollama](../../docs/providers/ollama)** + - Fix tool call errors due with improved message extraction - [PR #19369](https://github.com/BerriAI/litellm/pull/19369) + +- **[VertexAI](../../docs/providers/vertex)** + - Removed optional vertex_count_tokens_location param before request is sent to vertex - [PR #19359](https://github.com/BerriAI/litellm/pull/19359) + +- **Gemini([Vertex AI](../../docs/providers/vertex), [Google AI Studio](../../docs/providers/gemini))** + - Supports setting media_resolution and fps parameters on each video file, when using Gemini video understanding - [PR #19273](https://github.com/BerriAI/litellm/pull/19273) + - handle reasoning_effort as dict from OpenAI Agents SDK - [PR #19419](https://github.com/BerriAI/litellm/pull/19419) + - add file content support in tool results - [PR #19416](https://github.com/BerriAI/litellm/pull/19416) + +- **[Azure](../../docs/providers/azure_ai)** + - Fix Azure AI costs for Anthropic models - [PR #19530](https://github.com/BerriAI/litellm/pull/19530) + +- **[Giga Chat](../../docs/providers/gigachat)** + - Add tool choice mapping - [PR #19645](https://github.com/BerriAI/litellm/pull/19645) +--- + +## AI API Endpoints (LLMs, MCP, Agents) + +### Features + +- **[Files API](../../docs/files_endpoints)** + - Add managed files support when load_balancing is True - [PR #19338](https://github.com/BerriAI/litellm/pull/19338) + +- **[Claude Plugin Marketplace](../../docs/tutorials/claude_code_plugin_marketplace)** + - Add self hosted Claude Code Plugin Marketplace - [PR #19378](https://github.com/BerriAI/litellm/pull/19378) + +- **[MCP](../../docs/mcp)** + - Add MCP Protocol version 2025-11-25 support - [PR #19379](https://github.com/BerriAI/litellm/pull/19379) + - Log MCP tool calls and list tools in the LiteLLM Spend Logs table for easier debugging - [PR #19469](https://github.com/BerriAI/litellm/pull/19469) + +- **[Vertex AI](../../docs/providers/vertex)** + - Ensure only anthropic betas are forwarded down to LLM API (by default) - [PR #19542](https://github.com/BerriAI/litellm/pull/19542) + - Allow overriding to support forwarding incoming headers are forwarded down to target - [PR #19524](https://github.com/BerriAI/litellm/pull/19524) + +- **[Chat/Completions](../../docs/completion/input)** + - Add MCP tools response to chat completions - [PR #19552](https://github.com/BerriAI/litellm/pull/19552) + - Add custom vertex ai finish reasons to the output - [PR #19558](https://github.com/BerriAI/litellm/pull/19558) + - Return MCP execution in /chat/completions before model output during streaming - [PR #19623](https://github.com/BerriAI/litellm/pull/19623) + +### Bugs + +- **[Responses API](../../docs/response_api)** + - Fix duplicate messages during MCP streaming tool execution - [PR #19317](https://github.com/BerriAI/litellm/pull/19317) + - Fix pickle error when using OpenAI's Responses API with stream=True and tool_choice of type allowed_tools (an OpenAI-native parameter) - [PR #17205](https://github.com/BerriAI/litellm/pull/17205) + - stream tool call events for non-openai models - [PR #19368](https://github.com/BerriAI/litellm/pull/19368) + - preserve tool output ordering for gemini in responses bridge - [PR #19360](https://github.com/BerriAI/litellm/pull/19360) + - Add ID caching to prevent ID mismatch text-start and text-delta - [PR #19390](https://github.com/BerriAI/litellm/pull/19390) + - Include output_item, reasoning_summary_Text_done and reasoning_summary_part_done events for non-openai models - [PR #19472](https://github.com/BerriAI/litellm/pull/19472) + +- **[Chat/Completions](../../docs/completion/input)** + - fix: drop_params not dropping prompt_cache_key for non-OpenAI providers - [PR #19346](https://github.com/BerriAI/litellm/pull/19346) + +- **[Realtime API](../../docs/realtime)** + - disable SSL for ws:// WebSocket connections - [PR #19345](https://github.com/BerriAI/litellm/pull/19345) + +- **[Generate Content](../../docs/generateContent)** + - Log actual user input when google genai/vertex endpoints are called client-side - [PR #19156](https://github.com/BerriAI/litellm/pull/19156) + +- **[/messages/count_tokens Anthropic Token Counting](../../docs/anthropic_count_tokens)** + - ensure it works for Anthropic, Azure AI Anthropic on AI Gateway - [PR #19432](https://github.com/BerriAI/litellm/pull/19432) + +- **[MCP](../../docs/mcp)** + - forward static_headers to MCP servers - [PR #19366](https://github.com/BerriAI/litellm/pull/19366) + +- **[Batch API](../../docs/batches)** + - Fix: generation config empty for batch - [PR #19556](https://github.com/BerriAI/litellm/pull/19556) + +- **[Pass Through Endpoints](../../docs/proxy/pass_through)** + - Always reupdate registry - [PR #19420](https://github.com/BerriAI/litellm/pull/19420) +--- + +## Management Endpoints / UI + +### Features + +- **Cost Estimator** + - Fix model dropdown - [PR #19529](https://github.com/BerriAI/litellm/pull/19529) + +- **Claude Code Plugins** + - Allow Adding Claude Code Plugins via UI - [PR #19387](https://github.com/BerriAI/litellm/pull/19387) + +- **Guardrails** + - New Policy management UI - [PR #19668](https://github.com/BerriAI/litellm/pull/19668) + - Allow adding policies on Keys/Teams + Viewing on Info panels - [PR #19688](https://github.com/BerriAI/litellm/pull/19688) + +- **General** + - respects custom authentication header override - [PR #19276](https://github.com/BerriAI/litellm/pull/19276) + +- **Playground** + - Button to Fill Custom API Base - [PR #19440](https://github.com/BerriAI/litellm/pull/19440) + - display mcp output on the play ground - [PR #19553](https://github.com/BerriAI/litellm/pull/19553) + +- **Models** + - Paginate /v2/models/info - [PR #19521](https://github.com/BerriAI/litellm/pull/19521) + - All Model Tab Pagination - [PR #19525](https://github.com/BerriAI/litellm/pull/19525) + - Adding Optional scope Param to /models - [PR #19539](https://github.com/BerriAI/litellm/pull/19539) + - Model Search - [PR #19622](https://github.com/BerriAI/litellm/pull/19622) + - Filter by Model ID and Team ID - [PR #19713](https://github.com/BerriAI/litellm/pull/19713) + +- **MCP Servers** + - MCP Tools Tab Resetting to Overview - [PR #19468](https://github.com/BerriAI/litellm/pull/19468) + +- **Organizations** + - Prevent org admin from creating a new user with proxy_admin permissions - [PR #19296](https://github.com/BerriAI/litellm/pull/19296) + - Edit Page: Reusable Model Select - [PR #19601](https://github.com/BerriAI/litellm/pull/19601) + +- **Teams** + - Reusable Model Select - [PR #19543](https://github.com/BerriAI/litellm/pull/19543) + - [Fix] Team Update with Organization having All Proxy Models - [PR #19604](https://github.com/BerriAI/litellm/pull/19604) + +- **Logs** + - Include tool arguments in spend logs table - [PR #19640](https://github.com/BerriAI/litellm/pull/19640) + +- **Fallbacks / Loadbalancing** + - New fallbacks modal - [PR #19673](https://github.com/BerriAI/litellm/pull/19673) + - Set fallbacks/loadbalancing by team/key - [PR #19686](https://github.com/BerriAI/litellm/pull/19686) + +### Bugs + +- **Playground** + - increase model selector width in playground Compare view - [PR #19423](https://github.com/BerriAI/litellm/pull/19423) + +- **Virtual Keys** + - Sorting Shows Incorrect Entries - [PR #19534](https://github.com/BerriAI/litellm/pull/19534) + +- **General** + - UI 404 error when SERVER_ROOT_PATH is set - [PR #19467](https://github.com/BerriAI/litellm/pull/19467) + - Redirect to ui/login on expired JWT - [PR #19687](https://github.com/BerriAI/litellm/pull/19687) + +- **SSO** + - Fix SSO user roles not updating for existing users - [PR #19621](https://github.com/BerriAI/litellm/pull/19621) + +- **Guardrails** + - ensure guardrail patterns persist on edit and mode toggle - [PR #19265](https://github.com/BerriAI/litellm/pull/19265) +--- + +## AI Integrations + +### Logging + +- **General Logging** + - prevent printing duplicate StandardLoggingPayload logs - [PR #19325](https://github.com/BerriAI/litellm/pull/19325) + - Fix: log duplication when json_logs is enabled - [PR #19705](https://github.com/BerriAI/litellm/pull/19705) +- **Langfuse OTEL** + - ignore service logs and fix callback shadowing - [PR #19298](https://github.com/BerriAI/litellm/pull/19298) +- **Langfuse** + - Send litellm_trace_id - [PR #19528](https://github.com/BerriAI/litellm/pull/19528) + - Add Langfuse mock mode for testing without API calls - [PR #19676](https://github.com/BerriAI/litellm/pull/19676) +- **GCS Bucket** + - prevent unbounded queue growth due to slow API calls - [PR #19297](https://github.com/BerriAI/litellm/pull/19297) + - Add GCS mock mode for testing without API calls - [PR #19683](https://github.com/BerriAI/litellm/pull/19683) +- **Responses API Logging** + - Fix pydantic serialization error - [PR #19486](https://github.com/BerriAI/litellm/pull/19486) +- **Arize Phoenix** + - add openinference span kinds to arize phoenix - [PR #19267](https://github.com/BerriAI/litellm/pull/19267) +- **Prometheus** + - Added new prometheus metrics for user count and team count - [PR #19520](https://github.com/BerriAI/litellm/pull/19520) + +### Guardrails + +- **Bedrock Guardrails** + - Ensure post_call guardrail checks input+output - [PR #19151](https://github.com/BerriAI/litellm/pull/19151) +- **Prompt Security** + - fixing prompt-security's guardrail implementation - [PR #19374](https://github.com/BerriAI/litellm/pull/19374) +- **Presidio** + - Fixes crash in Presidio Guardrail when running in background threads (logging_hook) - [PR #19714](https://github.com/BerriAI/litellm/pull/19714) +- **Pillar Security** + - Migrate Pillar Security to Generic Guardrail API - [PR #19364](https://github.com/BerriAI/litellm/pull/19364) +- **Policy Engine** + - New LiteLLM Policy engine - create policies to manage guardrails, conditions - permissions per Key, Team - [PR #19612](https://github.com/BerriAI/litellm/pull/19612) +- **General** + - add case-insensitive support for guardrail mode and actions - [PR #19480](https://github.com/BerriAI/litellm/pull/19480) + +### Prompt Management + +- **General** + - fix prompt info lookup and delete using correct IDs - [PR #19358](https://github.com/BerriAI/litellm/pull/19358) + +### Secret Manager + +- **AWS Secret Manager** + - ensure auto-rotation updates existing AWS secret instead of creating new one - [PR #19455](https://github.com/BerriAI/litellm/pull/19455) +- **Hashicorp Vault** + - Ensure key rotations work with Vault - [PR #19634](https://github.com/BerriAI/litellm/pull/19634) + +--- + +## Spend Tracking, Budgets and Rate Limiting + +- **Pricing Updates** + - Add openai/dall-e base pricing entries - [PR #19133](https://github.com/BerriAI/litellm/pull/19133) + - Add `input_cost_per_video_per_second` in ModelInfoBase - [PR #19398](https://github.com/BerriAI/litellm/pull/19398) + +--- + +## Performance / Loadbalancing / Reliability improvements + + +- **General** + - Fix date overflow/division by zero in proxy utils - [PR #19527](https://github.com/BerriAI/litellm/pull/19527) + - Fix in-flight request termination on SIGTERM when health-check runs in a separate process - [PR #19427](https://github.com/BerriAI/litellm/pull/19427) + - Fix Pass through routes to work with server root path - [PR #19383](https://github.com/BerriAI/litellm/pull/19383) + - Fix logging error for stop iteration - [PR #19649](https://github.com/BerriAI/litellm/pull/19649) + - prevent retrying 4xx client errors - [PR #19275](https://github.com/BerriAI/litellm/pull/19275) + - add better error handling for misconfig on health check - [PR #19441](https://github.com/BerriAI/litellm/pull/19441) + +- **Router** + - Fix Azure RPM calculation formula - [PR #19513](https://github.com/BerriAI/litellm/pull/19513) + - Persist scheduler request queue to redis - [PR #19304](https://github.com/BerriAI/litellm/pull/19304) + - pass search_tools to Router during DB-triggered initialization - [PR #19388](https://github.com/BerriAI/litellm/pull/19388) + - Fixed PromptCachingCache to correctly handle messages where cache_control is a sibling key of string content - [PR #19266](https://github.com/BerriAI/litellm/pull/19266) + +- **Memory Leaks/OOM** + - prevent OOM with nested $defs in tool schemas - [PR #19112](https://github.com/BerriAI/litellm/pull/19112) + - fix: HTTP client memory leaks in Presidio, OpenAI, and Gemini - [PR #19190](https://github.com/BerriAI/litellm/pull/19190) + +- **Non root** + - fix logfile and pidfile of supervisor for non root environment - [PR #17267](https://github.com/BerriAI/litellm/pull/17267) + - resolve Read-only file system error in non-root images - [PR #19449](https://github.com/BerriAI/litellm/pull/19449) + +- **Dockerfile** + - Redis Semantic Caching - add missing redisvl dependency to requirements.txt - [PR #19417](https://github.com/BerriAI/litellm/pull/19417) + - Bump OTEL versions to support a2a dependency - resolves modulenotfounderror for Microsoft Agents by @Harshit28j in #18991 + +- **DB** + - Handle PostgreSQL cached plan errors during rolling deployments - [PR #19424](https://github.com/BerriAI/litellm/pull/19424) + +- **Timeouts** + - Fix: total timeout is not respected - [PR #19389](https://github.com/BerriAI/litellm/pull/19389) + +- **SDK** + - Field-Existence Checks to Type Classes to Prevent Attribute Errors - [PR #18321](https://github.com/BerriAI/litellm/pull/18321) + - add google-cloud-aiplatform as optional dependency with clear error message - [PR #19437](https://github.com/BerriAI/litellm/pull/19437) + - Make grpc dependency optional - [PR #19447](https://github.com/BerriAI/litellm/pull/19447) + - Add support for retry policies - [PR #19645](https://github.com/BerriAI/litellm/pull/19645) + +- **Performance** + - Cut chat_completion latency by ~21% by reducing pre-call processing time - [PR #19535](https://github.com/BerriAI/litellm/pull/19535) + - Optimize strip_trailing_slash with O(1) index check - [PR #19679](https://github.com/BerriAI/litellm/pull/19679) + - Optimize use_custom_pricing_for_model with set intersection - [PR #19677](https://github.com/BerriAI/litellm/pull/19677) + - perf: skip pattern_router.route() for non-wildcard models - [PR #19664](https://github.com/BerriAI/litellm/pull/19664) + - perf: Add LRU caching to get_model_info for faster cost lookups - [PR #19606](https://github.com/BerriAI/litellm/pull/19606) + +--- + +## General Proxy Improvements + +### Doc Improvements + - new tutorial for adding MCPs to Cursor via LiteLLM - [PR #19317](https://github.com/BerriAI/litellm/pull/19317) + - fix vertex_region to vertex_location in Vertex AI pass-through docs - [PR #19380](https://github.com/BerriAI/litellm/pull/19380) + - clarify Gemini and Vertex AI model prefix in json file - [PR #19443](https://github.com/BerriAI/litellm/pull/19443) + - update Claude Code integration guides - [PR #19415](https://github.com/BerriAI/litellm/pull/19415) + - adjust opencode tutorial - [PR #19605](https://github.com/BerriAI/litellm/pull/19605) + - add spend-queue-troubleshooting docs - [PR #19659](https://github.com/BerriAI/litellm/pull/19659) + - docs: add litellm-enterprise requirement for managed files - [PR #19689](https://github.com/BerriAI/litellm/pull/19689) + +### Helm + - Add support for keda in helm chart - [PR #19337](https://github.com/BerriAI/litellm/pull/19337) + - sync Helm chart version with LiteLLM release version - [PR #19438](https://github.com/BerriAI/litellm/pull/19438) + - Enable PreStop hook configuration in values.yaml - [PR #19613](https://github.com/BerriAI/litellm/pull/19613) + +### General + - Add health check scripts and parallel execution support - [PR #19295](https://github.com/BerriAI/litellm/pull/19295) + + +--- + +## New Contributors + + +* @dushyantzz made their first contribution in [PR #19158](https://github.com/BerriAI/litellm/pull/19158) +* @obod-mpw made their first contribution in [PR #19133](https://github.com/BerriAI/litellm/pull/19133) +* @msexxeta made their first contribution in [PR #19030](https://github.com/BerriAI/litellm/pull/19030) +* @rsicart made their first contribution in [PR #19337](https://github.com/BerriAI/litellm/pull/19337) +* @cluebbehusen made their first contribution in [PR #19311](https://github.com/BerriAI/litellm/pull/19311) +* @Lucky-Lodhi2004 made their first contribution in [PR #19315](https://github.com/BerriAI/litellm/pull/19315) +* @binbandit made their first contribution in [PR #19324](https://github.com/BerriAI/litellm/pull/19324) +* @flex-myeonghyeon made their first contribution in [PR #19381](https://github.com/BerriAI/litellm/pull/19381) +* @Lrakotoson made their first contribution in [PR #18321](https://github.com/BerriAI/litellm/pull/18321) +* @bensi94 made their first contribution in [PR #18787](https://github.com/BerriAI/litellm/pull/18787) +* @victorigualada made their first contribution in [PR #19368](https://github.com/BerriAI/litellm/pull/19368) +* @VedantMadane made their first contribution in #19266 +* @stiyyagura0901 made their first contribution in #19276 +* @kamilio made their first contribution in [PR #19447](https://github.com/BerriAI/litellm/pull/19447) +* @jonathansampson made their first contribution in [PR #19433](https://github.com/BerriAI/litellm/pull/19433) +* @rynecarbone made their first contribution in [PR #19416](https://github.com/BerriAI/litellm/pull/19416) +* @jayy-77 made their first contribution in #19366 +* @davida-ps made their first contribution in [PR #19374](https://github.com/BerriAI/litellm/pull/19374) +* @joaodinissf made their first contribution in [PR #19506](https://github.com/BerriAI/litellm/pull/19506) +* @ecao310 made their first contribution in [PR #19520](https://github.com/BerriAI/litellm/pull/19520) +* @mpcusack-altos made their first contribution in [PR #19577](https://github.com/BerriAI/litellm/pull/19577) +* @milan-berri made their first contribution in [PR #19602](https://github.com/BerriAI/litellm/pull/19602) +* @xqe2011 made their first contribution in #19621 + +--- + +## Full Changelog + +**[View complete changelog on GitHub](https://github.com/BerriAI/litellm/releases/tag/v1.81.3.rc)** diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 2d36dbeacda..ae95faede22 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -9,7 +9,7 @@ import datetime from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union import litellm -from litellm._logging import verbose_logger +from litellm._logging import verbose_logger, verbose_proxy_logger from litellm.a2a_protocol.streaming_iterator import A2AStreamingIterator from litellm.a2a_protocol.utils import A2ARequestUtils from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT @@ -20,6 +20,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.types.agents import LiteLLMSendMessageResponse from litellm.utils import client +import uuid if TYPE_CHECKING: from a2a.client import A2AClient as A2AClientType @@ -225,7 +226,11 @@ async def asend_message( raise ValueError( "Either a2a_client or api_base is required for standard A2A flow" ) - a2a_client = await create_a2a_client(base_url=api_base) + trace_id = str(uuid.uuid4()) + extra_headers = {"X-LiteLLM-Trace-Id": trace_id} + if agent_id: + extra_headers["X-LiteLLM-Agent-Id"] = agent_id + a2a_client = await create_a2a_client(base_url=api_base, extra_headers=extra_headers) # Type assertion: a2a_client is guaranteed to be non-None here assert a2a_client is not None @@ -490,6 +495,10 @@ async def create_a2a_client( ) httpx_client = http_handler.client + if extra_headers: + httpx_client.headers.update(extra_headers) + verbose_proxy_logger.debug(f"A2A client created with extra_headers={extra_headers}") + # Resolve agent card resolver = A2ACardResolver( httpx_client=httpx_client, diff --git a/litellm/constants.py b/litellm/constants.py index 49ca3a509b1..0525bf3843e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1165,6 +1165,7 @@ LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli" LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token" CLI_SSO_SESSION_CACHE_KEY_PREFIX = "cli_sso_session" CLI_JWT_TOKEN_NAME = "cli-jwt-token" +CLI_JWT_EXPIRATION_HOURS = int(os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", 24)) ########################### DB CRON JOB NAMES ########################### DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 50ec5cab429..e2de3cd5021 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -286,7 +286,9 @@ class MCPClient: return [] async def call_tool( - self, call_tool_request_params: MCPCallToolRequestParams + self, + call_tool_request_params: MCPCallToolRequestParams, + host_progress_callback: Optional[Callable] = None ) -> MCPCallToolResult: """ Call an MCP Tool. @@ -295,13 +297,28 @@ class MCPClient: f"MCP client calling tool '{call_tool_request_params.name}' with arguments: {call_tool_request_params.arguments}" ) + async def on_progress(progress: float, total: float | None, message: str | None): + percentage = (progress / total * 100) if total else 0 + verbose_logger.info( + f"MCP Tool '{call_tool_request_params.name}' progress: " + f"{progress}/{total} ({percentage:.0f}%) - {message or ''}" + ) + + # Forward to Host if callback provided + if host_progress_callback: + try: + await host_progress_callback(progress, total) + except Exception as e: + verbose_logger.warning(f"Failed to forward to Host: {e}") + async def _call_tool_operation(session: ClientSession): verbose_logger.debug("MCP client sending tool call to session") return await session.call_tool( name=call_tool_request_params.name, arguments=call_tool_request_params.arguments, - ) + progress_callback=on_progress, + ) try: tool_result = await self.run_with_session(_call_tool_operation) verbose_logger.info( diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 6b30b6b736e..6a003b8c499 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -83,6 +83,33 @@ }, "description": "Datadog Logging Integration" }, + { + "id": "datadog_cost_management", + "displayName": "Datadog Cost Management", + "logo": "datadog.png", + "supports_key_team_logging": false, + "dynamic_params": { + "dd_api_key": { + "type": "password", + "ui_name": "API Key", + "description": "Datadog API key for authentication", + "required": true + }, + "dd_app_key": { + "type": "password", + "ui_name": "App Key", + "description": "Datadog Application Key for Cloud Cost Management", + "required": true + }, + "dd_site": { + "type": "text", + "ui_name": "Site", + "description": "Datadog site URL (e.g., us5.datadoghq.com)", + "required": true + } + }, + "description": "Datadog Cloud Cost Management Integration" + }, { "id": "lago", "displayName": "Lago", @@ -407,4 +434,4 @@ }, "description": "SQS Queue (AWS) Logging Integration" } -] +] \ No newline at end of file diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 6a76b57e7f7..a5bb530fc56 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -516,7 +516,9 @@ class CustomGuardrail(CustomLogger): from litellm.types.utils import GuardrailMode # Use event_type if provided, otherwise fall back to self.event_hook - guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]] + guardrail_mode: Union[ + GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks] + ] if event_type is not None: guardrail_mode = event_type elif isinstance(self.event_hook, Mode): @@ -524,11 +526,21 @@ class CustomGuardrail(CustomLogger): else: guardrail_mode = self.event_hook # type: ignore[assignment] + from litellm.litellm_core_utils.core_helpers import ( + filter_exceptions_from_params, + ) + + # Sanitize the response to ensure it's JSON serializable and free of circular refs + # This prevents RecursionErrors in downstream loggers (Langfuse, Datadog, etc.) + clean_guardrail_response = filter_exceptions_from_params( + guardrail_json_response + ) + slg = StandardLoggingGuardrailInformation( guardrail_name=self.guardrail_name, guardrail_provider=guardrail_provider, guardrail_mode=guardrail_mode, - guardrail_response=guardrail_json_response, + guardrail_response=clean_guardrail_response, guardrail_status=guardrail_status, start_time=start_time, end_time=end_time, diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index 503e8d8c87a..735d1005d2c 100644 --- a/litellm/integrations/datadog/datadog.py +++ b/litellm/integrations/datadog/datadog.py @@ -32,6 +32,7 @@ from litellm.integrations.datadog.datadog_handler import ( get_datadog_service, get_datadog_source, get_datadog_tags, + get_datadog_base_url_from_env, ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.llms.custom_httpx.http_handler import ( @@ -100,7 +101,9 @@ class DataDogLogger( self._configure_dd_direct_api() # Optional override for testing - self._apply_dd_base_url_override() + dd_base_url = get_datadog_base_url_from_env() + if dd_base_url: + self.intake_url = f"{dd_base_url}/api/v2/logs" self.sync_client = _get_httpx_client() asyncio.create_task(self.periodic_flush()) self.flush_lock = asyncio.Lock() @@ -159,18 +162,6 @@ class DataDogLogger( self.DD_API_KEY = os.getenv("DD_API_KEY") self.intake_url = f"https://http-intake.logs.{os.getenv('DD_SITE')}/api/v2/logs" - def _apply_dd_base_url_override(self) -> None: - """ - Apply base URL override for testing purposes - """ - dd_base_url: Optional[str] = ( - os.getenv("_DATADOG_BASE_URL") - or os.getenv("DATADOG_BASE_URL") - or os.getenv("DD_BASE_URL") - ) - if dd_base_url is not None: - self.intake_url = f"{dd_base_url}/api/v2/logs" - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): """ Async Log success events to Datadog diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py new file mode 100644 index 00000000000..2eb94b59dd8 --- /dev/null +++ b/litellm/integrations/datadog/datadog_cost_management.py @@ -0,0 +1,204 @@ +import asyncio +import os +import time +from datetime import datetime +from typing import Dict, List, Optional, Tuple + +from litellm._logging import verbose_logger +from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.integrations.datadog_cost_management import ( + DatadogFOCUSCostEntry, +) +from litellm.types.utils import StandardLoggingPayload + + +class DatadogCostManagementLogger(CustomBatchLogger): + def __init__(self, **kwargs): + self.dd_api_key = os.getenv("DD_API_KEY") + self.dd_app_key = os.getenv("DD_APP_KEY") + self.dd_site = os.getenv("DD_SITE", "datadoghq.com") + + if not self.dd_api_key or not self.dd_app_key: + verbose_logger.warning( + "Datadog Cost Management: DD_API_KEY and DD_APP_KEY are required. Integration will not work." + ) + + self.upload_url = f"https://api.{self.dd_site}/api/v2/cost/custom_costs" + + self.async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback + ) + + # Initialize lock and start periodic flush task + self.flush_lock = asyncio.Lock() + asyncio.create_task(self.periodic_flush()) + + # Check if flush_lock is already in kwargs to avoid double passing (unlikely but safe) + if "flush_lock" not in kwargs: + kwargs["flush_lock"] = self.flush_lock + + super().__init__(**kwargs) + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + try: + standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get( + "standard_logging_object", None + ) + + if standard_logging_object is None: + return + + # Only log if there is a cost associated + if standard_logging_object.get("response_cost", 0) > 0: + self.log_queue.append(standard_logging_object) + + if len(self.log_queue) >= self.batch_size: + await self.async_send_batch() + + except Exception as e: + verbose_logger.exception( + f"Datadog Cost Management: Error in async_log_success_event: {str(e)}" + ) + + async def async_send_batch(self): + if not self.log_queue: + return + + try: + # Aggregate costs from the batch + aggregated_entries = self._aggregate_costs(self.log_queue) + + if not aggregated_entries: + return + + # Send to Datadog + await self._upload_to_datadog(aggregated_entries) + + # Clear queue only on success (or if we decide to drop on failure) + # CustomBatchLogger clears queue in flush_queue, so we just process here + + except Exception as e: + verbose_logger.exception( + f"Datadog Cost Management: Error in async_send_batch: {str(e)}" + ) + + def _aggregate_costs( + self, logs: List[StandardLoggingPayload] + ) -> List[DatadogFOCUSCostEntry]: + """ + Aggregates costs by Provider, Model, and Date. + Returns a list of DatadogFOCUSCostEntry. + """ + aggregator: Dict[Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry] = {} + + for log in logs: + try: + # Extract keys for aggregation + provider = log.get("custom_llm_provider") or "unknown" + model = log.get("model") or "unknown" + cost = log.get("response_cost", 0) + + if cost == 0: + continue + + # Get date strings (FOCUS format requires specific keys, but for aggregation we group by Day) + # UTC date + # We interpret "ChargePeriod" as the day of the request. + ts = log.get("startTime") or time.time() + dt = datetime.fromtimestamp(ts) + date_str = dt.strftime("%Y-%m-%d") + + # ChargePeriodStart and End + # If we want daily granularity, end date is usually same day or next day? + # Datadog Custom Costs usually expects periods. + # "ChargePeriodStart": "2023-01-01", "ChargePeriodEnd": "2023-12-31" in example. + # If we send daily, we can say Start=Date, End=Date. + + # Grouping Key: Provider + Model + Date + Tags? + # For simplicity, let's aggregate by Provider + Model + Date first. + # If we handle tags, we need to include them in the key. + + tags = self._extract_tags(log) + tags_key = tuple(sorted(tags.items())) if tags else () + + key = (provider, model, date_str, tags_key) + + if key not in aggregator: + aggregator[key] = { + "ProviderName": provider, + "ChargeDescription": f"LLM Usage for {model}", + "ChargePeriodStart": date_str, + "ChargePeriodEnd": date_str, + "BilledCost": 0.0, + "BillingCurrency": "USD", + "Tags": tags if tags else None, + } + + aggregator[key]["BilledCost"] += cost + + except Exception as e: + verbose_logger.warning( + f"Error processing log for cost aggregation: {e}" + ) + continue + + return list(aggregator.values()) + + def _extract_tags(self, log: StandardLoggingPayload) -> Dict[str, str]: + from litellm.integrations.datadog.datadog_handler import ( + get_datadog_env, + get_datadog_hostname, + get_datadog_pod_name, + get_datadog_service, + ) + + tags = { + "env": get_datadog_env(), + "service": get_datadog_service(), + "host": get_datadog_hostname(), + "pod_name": get_datadog_pod_name(), + } + + # Add metadata as tags + metadata = log.get("metadata", {}) + if metadata: + # Add user info + if "user_api_key_alias" in metadata: + tags["user"] = str(metadata["user_api_key_alias"]) + if "user_api_key_team_alias" in metadata: + tags["team"] = str(metadata["user_api_key_team_alias"]) + # model_group is not in StandardLoggingMetadata TypedDict, so we need to access it via dict.get() + model_group = metadata.get("model_group") # type: ignore[misc] + if model_group: + tags["model_group"] = str(model_group) + + return tags + + async def _upload_to_datadog(self, payload: List[Dict]): + if not self.dd_api_key or not self.dd_app_key: + return + + headers = { + "Content-Type": "application/json", + "DD-API-KEY": self.dd_api_key, + "DD-APPLICATION-KEY": self.dd_app_key, + } + + # The API endpoint expects a list of objects directly in the body (file content behavior) + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + + data_json = safe_dumps(payload) + + response = await self.async_client.put( + self.upload_url, content=data_json, headers=headers + ) + + response.raise_for_status() + + verbose_logger.debug( + f"Datadog Cost Management: Uploaded {len(payload)} cost entries. Status: {response.status_code}" + ) diff --git a/litellm/integrations/datadog/datadog_handler.py b/litellm/integrations/datadog/datadog_handler.py index 26fab77759e..e2f30f2f614 100644 --- a/litellm/integrations/datadog/datadog_handler.py +++ b/litellm/integrations/datadog/datadog_handler.py @@ -20,6 +20,14 @@ def get_datadog_hostname() -> str: return os.getenv("HOSTNAME", "") +def get_datadog_base_url_from_env() -> Optional[str]: + """ + Get base URL override from common DD_BASE_URL env var. + This is useful for testing or custom endpoints. + """ + return os.getenv("DD_BASE_URL") + + def get_datadog_env() -> str: return os.getenv("DD_ENV", "unknown") diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 6ffdbc0a005..4f6a5b339a7 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -21,6 +21,7 @@ from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.integrations.datadog.datadog_handler import ( get_datadog_service, get_datadog_tags, + get_datadog_base_url_from_env, ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -43,24 +44,22 @@ class DataDogLLMObsLogger(CustomBatchLogger): def __init__(self, **kwargs): try: verbose_logger.debug("DataDogLLMObs: Initializing logger") - if os.getenv("DD_API_KEY", None) is None: - raise Exception("DD_API_KEY is not set, set 'DD_API_KEY=<>'") - if os.getenv("DD_SITE", None) is None: - raise Exception( - "DD_SITE is not set, set 'DD_SITE=<>', example sit = `us5.datadoghq.com`" - ) + # Configure DataDog endpoint (Agent or Direct API) + # Use LITELLM_DD_AGENT_HOST to avoid conflicts with ddtrace's DD_AGENT_HOST + dd_agent_host = os.getenv("LITELLM_DD_AGENT_HOST") self.async_client = get_async_httpx_client( llm_provider=httpxSpecialProvider.LoggingCallback ) self.DD_API_KEY = os.getenv("DD_API_KEY") - self.DD_SITE = os.getenv("DD_SITE") - self.intake_url = ( - f"https://api.{self.DD_SITE}/api/intake/llm-obs/v1/trace/spans" - ) - # testing base url - dd_base_url = os.getenv("DD_BASE_URL") + if dd_agent_host: + self._configure_dd_agent(dd_agent_host=dd_agent_host) + else: + self._configure_dd_direct_api() + + # Optional override for testing + dd_base_url = get_datadog_base_url_from_env() if dd_base_url: self.intake_url = f"{dd_base_url}/api/intake/llm-obs/v1/trace/spans" @@ -78,6 +77,38 @@ class DataDogLLMObsLogger(CustomBatchLogger): verbose_logger.exception(f"DataDogLLMObs: Error initializing - {str(e)}") raise e + def _configure_dd_agent(self, dd_agent_host: str): + """ + Configure the Datadog logger to send traces to the Agent. + """ + # When using the Agent, LLM Observability Intake does NOT require the API Key + # Reference: https://docs.datadoghq.com/llm_observability/setup/sdk/#agent-setup + + # Use specific port for LLM Obs (Trace Agent) to avoid conflict with Logs Agent (10518) + agent_port = os.getenv("LITELLM_DD_LLM_OBS_PORT", "8126") + self.DD_SITE = "localhost" # Not used for URL construction in agent mode + self.intake_url = ( + f"http://{dd_agent_host}:{agent_port}/api/intake/llm-obs/v1/trace/spans" + ) + verbose_logger.debug(f"DataDogLLMObs: Using DD Agent at {self.intake_url}") + + def _configure_dd_direct_api(self): + """ + Configure the Datadog logger to send traces directly to the Datadog API. + """ + if not self.DD_API_KEY: + raise Exception("DD_API_KEY is not set, set 'DD_API_KEY=<>'") + + self.DD_SITE = os.getenv("DD_SITE") + if not self.DD_SITE: + raise Exception( + "DD_SITE is not set, set 'DD_SITE=<>', example site = `us5.datadoghq.com`" + ) + + self.intake_url = ( + f"https://api.{self.DD_SITE}/api/intake/llm-obs/v1/trace/spans" + ) + def _get_datadog_llm_obs_params(self) -> Dict: """ Get the datadog_llm_observability_params from litellm.datadog_llm_observability_params @@ -164,13 +195,14 @@ class DataDogLLMObsLogger(CustomBatchLogger): json_payload = safe_dumps(payload) + headers = {"Content-Type": "application/json"} + if self.DD_API_KEY: + headers["DD-API-KEY"] = self.DD_API_KEY + response = await self.async_client.post( url=self.intake_url, content=json_payload, - headers={ - "DD-API-KEY": self.DD_API_KEY, - "Content-Type": "application/json", - }, + headers=headers, ) if response.status_code != 202: diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 46ada3c3930..7bf97665fd2 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -23,6 +23,7 @@ from litellm.constants import MAX_LANGFUSE_INITIALIZED_CLIENTS from litellm.litellm_core_utils.core_helpers import ( safe_deep_copy, reconstruct_model_name, + filter_exceptions_from_params, ) from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info from litellm.integrations.langfuse.langfuse_mock_client import ( @@ -75,9 +76,8 @@ def _extract_cache_read_input_tokens(usage_obj) -> int: # Check prompt_tokens_details.cached_tokens (used by Gemini and other providers) if hasattr(usage_obj, "prompt_tokens_details"): prompt_tokens_details = getattr(usage_obj, "prompt_tokens_details", None) - if ( - prompt_tokens_details is not None - and hasattr(prompt_tokens_details, "cached_tokens") + if prompt_tokens_details is not None and hasattr( + prompt_tokens_details, "cached_tokens" ): cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None) if ( @@ -540,7 +540,6 @@ class LangFuseLogger: verbose_logger.debug("Langfuse Layer Logging - logging to langfuse v2") try: - metadata = metadata or {} standard_logging_object: Optional[StandardLoggingPayload] = cast( Optional[StandardLoggingPayload], kwargs.get("standard_logging_object", None), @@ -706,9 +705,10 @@ class LangFuseLogger: clean_metadata["litellm_response_cost"] = cost if standard_logging_object is not None: - clean_metadata["hidden_params"] = standard_logging_object[ - "hidden_params" - ] + hidden_params = standard_logging_object.get("hidden_params", {}) + clean_metadata["hidden_params"] = filter_exceptions_from_params( + hidden_params + ) if ( litellm.langfuse_default_tags is not None diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 9cb0a00d9fc..00695cbfb5b 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -351,9 +351,9 @@ def filter_exceptions_from_params(data: Any, max_depth: int = 20) -> Any: # Skip callable objects (functions, methods, lambdas) but not classes (type objects) if callable(data) and not isinstance(data, type): return None - # Skip known non-serializable object types (Logging, etc.) + # Skip known non-serializable object types (Logging, Router, etc.) obj_type_name = type(data).__name__ - if obj_type_name in ["Logging", "LiteLLMLoggingObj"]: + if obj_type_name in ["Logging", "LiteLLMLoggingObj", "Router"]: return None if isinstance(data, dict): diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e5412a650b7..66b14946c32 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3304,6 +3304,7 @@ def _get_masked_values( "token", "key", "secret", + "vertex_credentials", ] return { k: ( diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index bbe28e3ec2c..25ad0a570cb 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -21,11 +21,13 @@ from litellm.types.utils import ( ChatCompletionMessageToolCall, ChatCompletionRedactedThinkingBlock, Choices, + CompletionTokensDetailsWrapper, Delta, EmbeddingResponse, Function, HiddenParams, ImageResponse, + PromptTokensDetailsWrapper, ) from litellm.types.utils import Logprobs as TextCompletionLogprobs from litellm.types.utils import ( @@ -304,6 +306,22 @@ class LiteLLMResponseObjectHandler: "text_tokens": 0, } + # Map Responses API naming to Chat Completions API naming for cost calculator + if usage.get("prompt_tokens") is None: + usage["prompt_tokens"] = usage.get("input_tokens", 0) + if usage.get("completion_tokens") is None: + usage["completion_tokens"] = usage.get("output_tokens", 0) + + # Convert dicts to wrapper objects so getattr() works in cost calculation + if isinstance(usage.get("input_tokens_details"), dict): + usage["prompt_tokens_details"] = PromptTokensDetailsWrapper( + **usage["input_tokens_details"] + ) + if isinstance(usage.get("output_tokens_details"), dict): + usage["completion_tokens_details"] = CompletionTokensDetailsWrapper( + **usage["output_tokens_details"] + ) + if model_response_object is None: model_response_object = ImageResponse(**response_object) return model_response_object diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 30263543fc6..1d1e38c09da 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4408,7 +4408,7 @@ def _bedrock_tools_pt(tools: List) -> List[BedrockToolBlock]: ] """ """ - Bedrock toolConfig looks like: + Bedrock toolConfig looks like: "tools": [ { "toolSpec": { @@ -4436,6 +4436,7 @@ def _bedrock_tools_pt(tools: List) -> List[BedrockToolBlock]: tool_block_list: List[BedrockToolBlock] = [] for tool in tools: + # Handle regular function tools parameters = tool.get("function", {}).get( "parameters", {"type": "object", "properties": {}} ) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 9d50cc4d92d..71d74121a30 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -110,6 +110,10 @@ class AnthropicMessagesHandler(BaseTranslation): inputs["tools"] = tools_to_check if structured_messages: inputs["structured_messages"] = structured_messages + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, request_data=data, @@ -309,6 +313,14 @@ class AnthropicMessagesHandler(BaseTranslation): inputs["images"] = images_to_check if tool_calls_to_check: inputs["tool_calls"] = tool_calls_to_check + # Include model information from the response if available + response_model = None + if isinstance(response, dict): + response_model = response.get("model") + elif hasattr(response, "model"): + response_model = getattr(response, "model", None) + if response_model: + inputs["model"] = response_model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -552,7 +564,7 @@ class AnthropicMessagesHandler(BaseTranslation): response_content = response.get("content", []) else: response_content = getattr(response, "content", None) or [] - + if not response_content: return False for content_block in response_content: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 1706f045f14..5ba0754b744 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -168,6 +168,36 @@ class LiteLLMAnthropicMessagesAdapter: return provider_specific_fields.get("signature") return None + def _add_cache_control_if_applicable( + self, + source: Any, + target: Any, + model: Optional[str], + ) -> None: + """ + Extract cache_control from source and add to target if it should be preserved. + + This method accepts Any type to support both regular dicts and TypedDict objects. + TypedDict objects (like ChatCompletionTextObject, ChatCompletionImageObject, etc.) + are dicts at runtime but have specific types at type-check time. Using Any allows + this method to work with both while maintaining runtime correctness. + + Args: + source: Dict or TypedDict containing potential cache_control field + target: Dict or TypedDict to add cache_control to + model: Model name to check if cache_control should be preserved + """ + # TypedDict objects are dicts at runtime, so .get() works + cache_control = source.get("cache_control") if isinstance(source, dict) else getattr(source, "cache_control", None) + if cache_control and model and self.is_anthropic_claude_model(model): + # TypedDict objects support dict operations at runtime + # Use type ignore consistent with codebase pattern (see anthropic/chat/transformation.py:432) + if isinstance(target, dict): + target["cache_control"] = cache_control # type: ignore[typeddict-item] + else: + # Fallback for non-dict objects (shouldn't happen in practice) + cast(Dict[str, Any], target)["cache_control"] = cache_control + def translatable_anthropic_params(self) -> List: """ Which anthropic params, we need to translate to the openai format. @@ -205,12 +235,8 @@ class LiteLLMAnthropicMessagesAdapter: text_obj = ChatCompletionTextObject( type="text", text=content.get("text", "") ) - # Preserve cache_control if present (for prompt caching) - # Only for Anthropic models that support prompt caching - cache_control = content.get("cache_control") - if cache_control and model and self.is_anthropic_claude_model(model): - text_obj["cache_control"] = cache_control # type: ignore - new_user_content_list.append(text_obj) + self._add_cache_control_if_applicable(content, text_obj, model) + new_user_content_list.append(text_obj) # type: ignore elif content.get("type") == "image": # Convert Anthropic image format to OpenAI format source = content.get("source", {}) @@ -225,7 +251,24 @@ class LiteLLMAnthropicMessagesAdapter: image_obj = ChatCompletionImageObject( type="image_url", image_url=image_url_obj ) - new_user_content_list.append(image_obj) + self._add_cache_control_if_applicable(content, image_obj, model) + new_user_content_list.append(image_obj) # type: ignore + elif content.get("type") == "document": + # Convert Anthropic document format (PDF, etc.) to OpenAI format + source = content.get("source", {}) + openai_image_url = ( + self._translate_anthropic_image_to_openai(cast(dict, source)) + ) + + if openai_image_url: + image_url_obj = ChatCompletionImageUrlObject( + url=openai_image_url + ) + doc_obj = ChatCompletionImageObject( + type="image_url", image_url=image_url_obj + ) + self._add_cache_control_if_applicable(content, doc_obj, model) + new_user_content_list.append(doc_obj) # type: ignore elif content.get("type") == "tool_result": if "content" not in content: tool_result = ChatCompletionToolMessage( @@ -233,14 +276,16 @@ class LiteLLMAnthropicMessagesAdapter: tool_call_id=content.get("tool_use_id", ""), content="", ) - tool_message_list.append(tool_result) + self._add_cache_control_if_applicable(content, tool_result, model) + tool_message_list.append(tool_result) # type: ignore[arg-type] elif isinstance(content.get("content"), str): tool_result = ChatCompletionToolMessage( role="tool", tool_call_id=content.get("tool_use_id", ""), content=str(content.get("content", "")), ) - tool_message_list.append(tool_result) + self._add_cache_control_if_applicable(content, tool_result, model) + tool_message_list.append(tool_result) # type: ignore[arg-type] elif isinstance(content.get("content"), list): # Combine all content items into a single tool message # to avoid creating multiple tool_result blocks with the same ID @@ -256,7 +301,8 @@ class LiteLLMAnthropicMessagesAdapter: tool_call_id=content.get("tool_use_id", ""), content=c, ) - tool_message_list.append(tool_result) + self._add_cache_control_if_applicable(content, tool_result, model) + tool_message_list.append(tool_result) # type: ignore[arg-type] elif isinstance(c, dict): if c.get("type") == "text": tool_result = ChatCompletionToolMessage( @@ -266,7 +312,8 @@ class LiteLLMAnthropicMessagesAdapter: ), content=c.get("text", ""), ) - tool_message_list.append(tool_result) + self._add_cache_control_if_applicable(content, tool_result, model) + tool_message_list.append(tool_result) # type: ignore[arg-type] elif c.get("type") == "image": source = c.get("source", {}) openai_image_url = ( @@ -282,7 +329,8 @@ class LiteLLMAnthropicMessagesAdapter: ), content=openai_image_url, ) - tool_message_list.append(tool_result) + self._add_cache_control_if_applicable(content, tool_result, model) + tool_message_list.append(tool_result) # type: ignore[arg-type] else: # For multiple content items, combine into a single tool message # with list content to preserve all items while having one tool_use_id @@ -331,7 +379,8 @@ class LiteLLMAnthropicMessagesAdapter: tool_call_id=content.get("tool_use_id", ""), content=combined_content_parts, # type: ignore ) - tool_message_list.append(tool_result) + self._add_cache_control_if_applicable(content, tool_result, model) + tool_message_list.append(tool_result) # type: ignore[arg-type] if len(tool_message_list) > 0: new_messages.extend(tool_message_list) @@ -344,6 +393,8 @@ class LiteLLMAnthropicMessagesAdapter: ## ASSISTANT MESSAGE ## assistant_message_str: Optional[str] = None + assistant_content_list: List[Dict[str, Any]] = [] # For content blocks with cache_control + has_cache_control_in_text = False tool_calls: List[ChatCompletionAssistantToolCall] = [] thinking_blocks: List[ Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock] @@ -357,10 +408,14 @@ class LiteLLMAnthropicMessagesAdapter: assistant_message_str = str(content) elif isinstance(content, dict): if content.get("type") == "text": - if assistant_message_str is None: - assistant_message_str = content.get("text", "") - else: - assistant_message_str += content.get("text", "") + text_block: Dict[str, Any] = { + "type": "text", + "text": content.get("text", ""), + } + self._add_cache_control_if_applicable(content, text_block, model) + if "cache_control" in text_block: + has_cache_control_in_text = True + assistant_content_list.append(text_block) elif content.get("type") == "tool_use": function_chunk: ChatCompletionToolCallFunctionChunk = { "name": content.get("name", ""), @@ -384,13 +439,13 @@ class LiteLLMAnthropicMessagesAdapter: provider_specific_fields ) - tool_calls.append( - ChatCompletionAssistantToolCall( - id=content.get("id", ""), - type="function", - function=function_chunk, - ) + tool_call = ChatCompletionAssistantToolCall( + id=content.get("id", ""), + type="function", + function=function_chunk, ) + self._add_cache_control_if_applicable(content, tool_call, model) + tool_calls.append(tool_call) elif content.get("type") == "thinking": thinking_block = ChatCompletionThinkingBlock( type="thinking", @@ -411,18 +466,30 @@ class LiteLLMAnthropicMessagesAdapter: if ( assistant_message_str is not None + or len(assistant_content_list) > 0 or len(tool_calls) > 0 or len(thinking_blocks) > 0 ): + # Use list format if any text block has cache_control, otherwise use string + if has_cache_control_in_text and len(assistant_content_list) > 0: + assistant_content: Any = assistant_content_list + elif len(assistant_content_list) > 0 and not has_cache_control_in_text: + # Concatenate text blocks into string when no cache_control + assistant_content = "".join( + block.get("text", "") for block in assistant_content_list + ) + else: + assistant_content = assistant_message_str + assistant_message = ChatCompletionAssistantMessage( role="assistant", - content=assistant_message_str, + content=assistant_content, thinking_blocks=( thinking_blocks if len(thinking_blocks) > 0 else None ), ) if len(tool_calls) > 0: - assistant_message["tool_calls"] = tool_calls + assistant_message["tool_calls"] = tool_calls # type: ignore if len(thinking_blocks) > 0: assistant_message["thinking_blocks"] = thinking_blocks # type: ignore new_messages.append(assistant_message) @@ -532,10 +599,10 @@ class LiteLLMAnthropicMessagesAdapter: ) def translate_anthropic_tools_to_openai( - self, tools: List[AllAnthropicToolsValues] + self, tools: List[AllAnthropicToolsValues], model: Optional[str] = None ) -> List[ChatCompletionToolParam]: new_tools: List[ChatCompletionToolParam] = [] - mapped_tool_params = ["name", "input_schema", "description"] + mapped_tool_params = ["name", "input_schema", "description", "cache_control"] for tool in tools: function_chunk = ChatCompletionToolParamFunctionChunk( name=tool["name"], @@ -548,11 +615,11 @@ class LiteLLMAnthropicMessagesAdapter: for k, v in tool.items(): if k not in mapped_tool_params: # pass additional computer kwargs function_chunk.setdefault("parameters", {}).update({k: v}) - new_tools.append( - ChatCompletionToolParam(type="function", function=function_chunk) - ) + tool_param = ChatCompletionToolParam(type="function", function=function_chunk) + self._add_cache_control_if_applicable(tool, tool_param, model) + new_tools.append(tool_param) # type: ignore[arg-type] - return new_tools + return new_tools # type: ignore[return-value] def translate_anthropic_output_format_to_openai( self, output_format: Any @@ -590,6 +657,41 @@ class LiteLLMAnthropicMessagesAdapter: }, } + def _add_system_message_to_messages( + self, + new_messages: List[AllMessageValues], + anthropic_message_request: AnthropicMessagesRequest, + ) -> None: + """Add system message to messages list if present in request.""" + if "system" not in anthropic_message_request: + return + system_content = anthropic_message_request["system"] + if not system_content: + return + # Handle system as string or array of content blocks + if isinstance(system_content, str): + new_messages.insert( + 0, + ChatCompletionSystemMessage(role="system", content=system_content), + ) + elif isinstance(system_content, list): + # Convert Anthropic system content blocks to OpenAI format + openai_system_content: List[Dict[str, Any]] = [] + model_name = anthropic_message_request.get("model", "") + for block in system_content: + if isinstance(block, dict) and block.get("type") == "text": + text_block: Dict[str, Any] = { + "type": "text", + "text": block.get("text", ""), + } + self._add_cache_control_if_applicable(block, text_block, model_name) + openai_system_content.append(text_block) + if openai_system_content: + new_messages.insert( + 0, + ChatCompletionSystemMessage(role="system", content=openai_system_content), # type: ignore + ) + def translate_anthropic_to_openai( self, anthropic_message_request: AnthropicMessagesRequest ) -> ChatCompletionRequest: @@ -618,13 +720,7 @@ class LiteLLMAnthropicMessagesAdapter: model=anthropic_message_request.get("model"), ) ## ADD SYSTEM MESSAGE TO MESSAGES - if "system" in anthropic_message_request: - system_content = anthropic_message_request["system"] - if system_content: - new_messages.insert( - 0, - ChatCompletionSystemMessage(role="system", content=system_content), - ) + self._add_system_message_to_messages(new_messages, anthropic_message_request) new_kwargs: ChatCompletionRequest = { "model": anthropic_message_request["model"], @@ -655,7 +751,8 @@ class LiteLLMAnthropicMessagesAdapter: tools = anthropic_message_request["tools"] if tools: new_kwargs["tools"] = self.translate_anthropic_tools_to_openai( - tools=cast(List[AllAnthropicToolsValues], tools) + tools=cast(List[AllAnthropicToolsValues], tools), + model=new_kwargs.get("model"), ) ## CONVERT THINKING @@ -843,7 +940,7 @@ class LiteLLMAnthropicMessagesAdapter: role="assistant", model=response.model or "unknown-model", stop_sequence=None, - usage=anthropic_usage, + usage=anthropic_usage, # type: ignore content=anthropic_content, # type: ignore stop_reason=anthropic_finish_reason, ) @@ -992,7 +1089,7 @@ class LiteLLMAnthropicMessagesAdapter: else: usage_delta = UsageDelta(input_tokens=0, output_tokens=0) return MessageBlockDelta( - type="message_delta", delta=delta, usage=usage_delta + type="message_delta", delta=delta, usage=usage_delta # type: ignore ) ( type_of_content, diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index cb26a22edfa..ec665142073 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -298,6 +298,39 @@ class AmazonConverseConfig(BaseConfig): # Check if the model is specifically Nova Lite 2 return "nova-2-lite" in model_without_region + def _map_web_search_options( + self, + web_search_options: dict, + model: str + ) -> Optional[BedrockToolBlock]: + """ + Map web_search_options to Nova grounding systemTool. + + Nova grounding (web search) is only supported on Amazon Nova models. + Returns None for non-Nova models. + + Args: + web_search_options: The web_search_options dict from the request + model: The model identifier string + + Returns: + BedrockToolBlock with systemTool for Nova models, None otherwise + + Reference: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html + """ + # Only Nova models support nova_grounding + # Model strings can be like: "amazon.nova-pro-v1:0", "us.amazon.nova-pro-v1:0", etc. + if "nova" not in model.lower(): + verbose_logger.debug( + f"web_search_options passed but model {model} is not a Nova model. " + "Nova grounding is only supported on Amazon Nova models." + ) + return None + + # Nova doesn't support search_context_size or user_location params + # (unlike Anthropic), so we just enable grounding with no options + return BedrockToolBlock(systemTool={"name": "nova_grounding"}) + def _transform_reasoning_effort_to_reasoning_config( self, reasoning_effort: str ) -> dict: @@ -438,6 +471,10 @@ class AmazonConverseConfig(BaseConfig): ): supported_params.append("tools") + # Nova models support web_search_options (mapped to nova_grounding systemTool) + if base_model.startswith("amazon.nova"): + supported_params.append("web_search_options") + if litellm.utils.supports_tool_choice( model=model, custom_llm_provider=self.custom_llm_provider ) or litellm.utils.supports_tool_choice( @@ -730,6 +767,13 @@ class AmazonConverseConfig(BaseConfig): if bedrock_tier in ("default", "flex", "priority"): optional_params["serviceTier"] = {"type": bedrock_tier} + if param == "web_search_options" and value and isinstance(value, dict): + grounding_tool = self._map_web_search_options(value, model) + if grounding_tool is not None: + optional_params = self._add_tools_to_optional_params( + optional_params=optional_params, tools=[grounding_tool] + ) + # Only update thinking tokens for non-GPT-OSS models and non-Nova-Lite-2 models # Nova Lite 2 handles token budgeting differently through reasoningConfig if "gpt-oss" not in model and not self._is_nova_lite_2_model(model): @@ -1388,20 +1432,23 @@ class AmazonConverseConfig(BaseConfig): str, List[ChatCompletionToolCallChunk], Optional[List[BedrockConverseReasoningContentBlock]], + Optional[List[CitationsContentBlock]], ]: """ - Translate the message content to a string and a list of tool calls and reasoning content blocks + Translate the message content to a string and a list of tool calls, reasoning content blocks, and citations. Returns: content_str: str tools: List[ChatCompletionToolCallChunk] reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] + citationsContentBlocks: Optional[List[CitationsContentBlock]] - Citations from Nova grounding """ content_str = "" tools: List[ChatCompletionToolCallChunk] = [] reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( None ) + citationsContentBlocks: Optional[List[CitationsContentBlock]] = None for idx, content in enumerate(content_blocks): """ - Content is either a tool response or text @@ -1446,10 +1493,15 @@ class AmazonConverseConfig(BaseConfig): if reasoningContentBlocks is None: reasoningContentBlocks = [] reasoningContentBlocks.append(content["reasoningContent"]) + # Handle Nova grounding citations content + if "citationsContent" in content: + if citationsContentBlocks is None: + citationsContentBlocks = [] + citationsContentBlocks.append(content["citationsContent"]) - return content_str, tools, reasoningContentBlocks + return content_str, tools, reasoningContentBlocks, citationsContentBlocks - def _transform_response( + def _transform_response( # noqa: PLR0915 self, model: str, response: httpx.Response, @@ -1525,18 +1577,27 @@ class AmazonConverseConfig(BaseConfig): reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( None ) + citationsContentBlocks: Optional[List[CitationsContentBlock]] = None if message is not None: ( content_str, tools, reasoningContentBlocks, + citationsContentBlocks, ) = self._translate_message_content(message["content"]) + # Initialize provider_specific_fields if we have any special content blocks + provider_specific_fields: dict = {} + if reasoningContentBlocks is not None: + provider_specific_fields["reasoningContentBlocks"] = reasoningContentBlocks + if citationsContentBlocks is not None: + provider_specific_fields["citationsContent"] = citationsContentBlocks + + if provider_specific_fields: + chat_completion_message["provider_specific_fields"] = provider_specific_fields + if reasoningContentBlocks is not None: - chat_completion_message["provider_specific_fields"] = { - "reasoningContentBlocks": reasoningContentBlocks, - } chat_completion_message["reasoning_content"] = ( self._transform_reasoning_content(reasoningContentBlocks) ) diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 17474fa022b..1c58a11eebe 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -1476,6 +1476,11 @@ class AWSEventStreamDecoder: reasoning_content = ( "" # set to non-empty string to ensure consistency with Anthropic ) + elif "citationsContent" in delta_obj: + # Handle Nova grounding citations in streaming responses + provider_specific_fields = { + "citationsContent": delta_obj["citationsContent"], + } return ( text, tool_use, diff --git a/litellm/llms/cohere/rerank/guardrail_translation/handler.py b/litellm/llms/cohere/rerank/guardrail_translation/handler.py index 6893a5991c3..b8133c59f7d 100644 --- a/litellm/llms/cohere/rerank/guardrail_translation/handler.py +++ b/litellm/llms/cohere/rerank/guardrail_translation/handler.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -49,8 +50,13 @@ class CohereRerankHandler(BaseTranslation): # Process query only query = data.get("query") if query is not None and isinstance(query, str): + inputs = GenericGuardrailAPIInputs(texts=[query]) + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [query]}, + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index d0ed3f165cc..fb00aa28f45 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -87,6 +87,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation): tools = data.get("tools") if tools: inputs["tools"] = tools + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -297,6 +301,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): inputs["images"] = images_to_check if tool_calls_to_check: inputs["tool_calls"] = tool_calls_to_check # type: ignore + # Include model information from the response if available + if hasattr(response, "model") and response.model: + inputs["model"] = response.model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -417,6 +424,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation): inputs = GenericGuardrailAPIInputs(texts=texts_to_check) if images_to_check: inputs["images"] = images_to_check + # Include model information from the first response if available + if ( + responses_so_far + and hasattr(responses_so_far[0], "model") + and responses_so_far[0].model + ): + inputs["model"] = responses_so_far[0].model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, request_data=request_data, diff --git a/litellm/llms/openai/completion/guardrail_translation/handler.py b/litellm/llms/openai/completion/guardrail_translation/handler.py index 73d08cfead4..1f8c6159da0 100644 --- a/litellm/llms/openai/completion/guardrail_translation/handler.py +++ b/litellm/llms/openai/completion/guardrail_translation/handler.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -53,8 +54,13 @@ class OpenAITextCompletionHandler(BaseTranslation): if isinstance(prompt, str): # Single string prompt + inputs = GenericGuardrailAPIInputs(texts=[prompt]) + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [prompt]}, + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, @@ -80,8 +86,13 @@ class OpenAITextCompletionHandler(BaseTranslation): text_indices.append(idx) if texts_to_check: + inputs = GenericGuardrailAPIInputs(texts=texts_to_check) + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": texts_to_check}, + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, @@ -154,8 +165,12 @@ class OpenAITextCompletionHandler(BaseTranslation): if user_metadata: request_data["litellm_metadata"] = user_metadata + inputs = GenericGuardrailAPIInputs(texts=texts_to_check) + # Include model information from the response if available + if hasattr(response, "model") and response.model: + inputs["model"] = response.model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": texts_to_check}, + inputs=inputs, request_data=request_data, input_type="response", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/openai/image_generation/cost_calculator.py b/litellm/llms/openai/image_generation/cost_calculator.py index 35caaf6e9b1..988d5626134 100644 --- a/litellm/llms/openai/image_generation/cost_calculator.py +++ b/litellm/llms/openai/image_generation/cost_calculator.py @@ -8,8 +8,7 @@ from typing import Optional from litellm import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token -from litellm.responses.utils import ResponseAPILoggingUtils -from litellm.types.utils import ImageResponse +from litellm.types.utils import ImageResponse, Usage def cost_calculator( @@ -39,11 +38,18 @@ def cost_calculator( ) return 0.0 - # Transform ImageUsage to Usage using the existing helper - # ImageUsage has the same format as ResponseAPIUsage - chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( - usage - ) + # If usage is already a Usage object with completion_tokens_details set, + # use it directly (it was already transformed in convert_to_image_response) + if isinstance(usage, Usage) and usage.completion_tokens_details is not None: + chat_usage = usage + else: + # Transform ImageUsage to Usage using the existing helper + # ImageUsage has the same format as ResponseAPIUsage + from litellm.responses.utils import ResponseAPILoggingUtils + + chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( + usage + ) # Use generic_cost_per_token for cost calculation prompt_cost, completion_cost = generic_cost_per_token( diff --git a/litellm/llms/openai/image_generation/guardrail_translation/handler.py b/litellm/llms/openai/image_generation/guardrail_translation/handler.py index 842a64b1878..e6340ba4705 100644 --- a/litellm/llms/openai/image_generation/guardrail_translation/handler.py +++ b/litellm/llms/openai/image_generation/guardrail_translation/handler.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -52,8 +53,13 @@ class OpenAIImageGenerationHandler(BaseTranslation): # Apply guardrail to the prompt if isinstance(prompt, str): + inputs = GenericGuardrailAPIInputs(texts=[prompt]) + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [prompt]}, + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 9b8f15c7623..d943662f9e4 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -105,6 +105,10 @@ class OpenAIResponsesHandler(BaseTranslation): inputs["tools"] = tools_to_check if structured_messages: inputs["structured_messages"] = structured_messages # type: ignore + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -150,6 +154,10 @@ class OpenAIResponsesHandler(BaseTranslation): inputs["tools"] = tools_to_check if structured_messages: inputs["structured_messages"] = structured_messages # type: ignore + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, request_data=data, @@ -344,6 +352,14 @@ class OpenAIResponsesHandler(BaseTranslation): inputs["images"] = images_to_check if tool_calls_to_check: inputs["tool_calls"] = tool_calls_to_check + # Include model information from the response if available + response_model = None + if isinstance(response, dict): + response_model = response.get("model") + elif hasattr(response, "model"): + response_model = getattr(response, "model", None) + if response_model: + inputs["model"] = response_model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -388,12 +404,15 @@ class OpenAIResponsesHandler(BaseTranslation): tool_calls = model_response_stream.choices[0].delta.tool_calls if tool_calls: + inputs = GenericGuardrailAPIInputs() + inputs["tool_calls"] = cast( + List[ChatCompletionToolCallChunk], tool_calls + ) + # Include model information if available + if hasattr(model_response_stream, "model") and model_response_stream.model: + inputs["model"] = model_response_stream.model _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={ - "tool_calls": cast( - List[ChatCompletionToolCallChunk], tool_calls - ) - }, + inputs=inputs, request_data={}, input_type="response", logging_obj=litellm_logging_obj, @@ -417,7 +436,11 @@ class OpenAIResponsesHandler(BaseTranslation): guardrail_inputs["tool_calls"] = cast( List[ChatCompletionToolCallChunk], tool_calls ) - if tool_calls: + # Include model information from the response if available + response_model = final_chunk.get("response", {}).get("model") + if response_model: + guardrail_inputs["model"] = response_model + if tool_calls or text: _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=guardrail_inputs, request_data={}, @@ -429,8 +452,14 @@ class OpenAIResponsesHandler(BaseTranslation): # tool_calls = model_response_stream.choices[0].tool_calls # convert openai response to model response string_so_far = self.get_streaming_string_so_far(responses_so_far) + inputs = GenericGuardrailAPIInputs(texts=[string_so_far]) + # Try to get model from the final chunk if available + if isinstance(final_chunk, dict): + response_model = final_chunk.get("response", {}).get("model") if isinstance(final_chunk.get("response"), dict) else None + if response_model: + inputs["model"] = response_model _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [string_so_far]}, + inputs=inputs, request_data={}, input_type="response", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/openai/speech/guardrail_translation/handler.py b/litellm/llms/openai/speech/guardrail_translation/handler.py index 4c2f71477be..e6796fbac2a 100644 --- a/litellm/llms/openai/speech/guardrail_translation/handler.py +++ b/litellm/llms/openai/speech/guardrail_translation/handler.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -50,8 +51,13 @@ class OpenAITextToSpeechHandler(BaseTranslation): return data if isinstance(input_text, str): + inputs = GenericGuardrailAPIInputs(texts=[input_text]) + # Include model information if available (voice model) + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [input_text]}, + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/openai/transcriptions/guardrail_translation/handler.py b/litellm/llms/openai/transcriptions/guardrail_translation/handler.py index ac416f42c81..3d76a21c389 100644 --- a/litellm/llms/openai/transcriptions/guardrail_translation/handler.py +++ b/litellm/llms/openai/transcriptions/guardrail_translation/handler.py @@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -88,8 +89,12 @@ class OpenAIAudioTranscriptionHandler(BaseTranslation): if user_metadata: request_data["litellm_metadata"] = user_metadata + inputs = GenericGuardrailAPIInputs(texts=[original_text]) + # Include model information from the response if available + if hasattr(response, "model") and response.model: + inputs["model"] = response.model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [original_text]}, + inputs=inputs, request_data=request_data, input_type="response", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index c0979e37e66..40433d53413 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, List, Optional from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.proxy._types import PassThroughGuardrailSettings +from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -118,8 +119,13 @@ class PassThroughEndpointHandler(BaseTranslation): return data # Apply guardrail (pass-through doesn't modify the text, just checks it) + inputs = GenericGuardrailAPIInputs(texts=[text_to_check]) + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [text_to_check]}, + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, @@ -178,8 +184,13 @@ class PassThroughEndpointHandler(BaseTranslation): request_data["litellm_metadata"] = user_metadata # Apply guardrail (pass-through doesn't modify the text, just checks it) + inputs = GenericGuardrailAPIInputs(texts=[text_to_check]) + # Include model information from the response if available + response_model = response.get("model") if isinstance(response, dict) else None + if response_model: + inputs["model"] = response_model _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs={"texts": [text_to_check]}, + inputs=inputs, request_data=request_data, input_type="response", logging_obj=litellm_logging_obj, diff --git a/litellm/llms/vercel_ai_gateway/embedding/__init__.py b/litellm/llms/vercel_ai_gateway/embedding/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/vercel_ai_gateway/embedding/transformation.py b/litellm/llms/vercel_ai_gateway/embedding/transformation.py new file mode 100644 index 00000000000..7238b05f10d --- /dev/null +++ b/litellm/llms/vercel_ai_gateway/embedding/transformation.py @@ -0,0 +1,176 @@ +""" +Vercel AI Gateway Embedding API Configuration. + +This module provides the configuration for Vercel AI Gateway's Embedding API. +Vercel AI Gateway is OpenAI-compatible and supports embeddings via the /v1/embeddings endpoint. + +Docs: https://vercel.com/docs/ai-gateway/openai-compat/embeddings +""" + +from typing import TYPE_CHECKING, Any, Optional + +import httpx + +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllEmbeddingInputValues +from litellm.types.utils import EmbeddingResponse +from litellm.utils import convert_to_model_response_object + +from ..common_utils import VercelAIGatewayException + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig): + """ + Configuration for Vercel AI Gateway's Embedding API. + + Reference: https://vercel.com/docs/ai-gateway/openai-compat/embeddings + """ + + def validate_environment( + self, + headers: dict, + model: str, + messages: list, + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """ + Validate environment and set up headers for Vercel AI Gateway API. + + Vercel AI Gateway requires: + - Authorization header with Bearer token (API key or OIDC token) + """ + vercel_headers = { + "Content-Type": "application/json", + } + + # Add Authorization header if api_key is provided + if api_key: + vercel_headers["Authorization"] = f"Bearer {api_key}" + + # Merge with existing headers (user's extra_headers take priority) + merged_headers = {**vercel_headers, **headers} + + return merged_headers + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + Get the complete URL for Vercel AI Gateway Embedding API endpoint. + """ + if api_base: + api_base = api_base.rstrip("/") + else: + api_base = ( + get_secret_str("VERCEL_AI_GATEWAY_API_BASE") + or "https://ai-gateway.vercel.sh/v1" + ) + + return f"{api_base}/embeddings" + + def transform_embedding_request( + self, + model: str, + input: AllEmbeddingInputValues, + optional_params: dict, + headers: dict, + ) -> dict: + """ + Transform embedding request to Vercel AI Gateway format (OpenAI-compatible). + """ + # Ensure input is a list + if isinstance(input, str): + input = [input] + + # Strip 'vercel_ai_gateway/' prefix if present + if model.startswith("vercel_ai_gateway/"): + model = model.replace("vercel_ai_gateway/", "", 1) + + return { + "model": model, + "input": input, + **optional_params, + } + + def transform_embedding_response( + self, + model: str, + raw_response: httpx.Response, + model_response: EmbeddingResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + request_data: dict, + optional_params: dict, + litellm_params: dict, + ) -> EmbeddingResponse: + """ + Transform embedding response from Vercel AI Gateway format (OpenAI-compatible). + """ + logging_obj.post_call(original_response=raw_response.text) + + # Vercel AI Gateway returns standard OpenAI-compatible embedding response + response_json = raw_response.json() + + return convert_to_model_response_object( + response_object=response_json, + model_response_object=model_response, + response_type="embedding", + ) + + def get_supported_openai_params(self, model: str) -> list: + """ + Get list of supported OpenAI parameters for Vercel AI Gateway embeddings. + + Vercel AI Gateway supports the standard OpenAI embeddings parameters + and auto-maps 'dimensions' to each provider's expected field. + """ + return [ + "timeout", + "dimensions", + "encoding_format", + "user", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to Vercel AI Gateway format. + """ + for param, value in non_default_params.items(): + if param in self.get_supported_openai_params(model): + optional_params[param] = value + return optional_params + + def get_error_class( + self, error_message: str, status_code: int, headers: Any + ) -> Any: + """ + Get the error class for Vercel AI Gateway errors. + """ + return VercelAIGatewayException( + message=error_message, + status_code=status_code, + headers=headers, + ) diff --git a/litellm/main.py b/litellm/main.py index a4bcfdec81b..5b8c569a390 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4867,6 +4867,36 @@ def embedding( # noqa: PLR0915 headers = openrouter_headers + response = base_llm_http_handler.embedding( + model=model, + input=input, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + api_key=api_key, + logging_obj=logging, + timeout=timeout, + model_response=EmbeddingResponse(), + optional_params=optional_params, + client=client, + aembedding=aembedding, + litellm_params=litellm_params_dict, + headers=headers, + ) + elif custom_llm_provider == "vercel_ai_gateway": + api_base = ( + api_base + or litellm.api_base + or get_secret_str("VERCEL_AI_GATEWAY_API_BASE") + or "https://ai-gateway.vercel.sh/v1" + ) + + api_key = ( + api_key + or litellm.api_key + or get_secret_str("VERCEL_AI_GATEWAY_API_KEY") + or get_secret_str("VERCEL_OIDC_TOKEN") + ) + response = base_llm_http_handler.embedding( model=model, input=input, diff --git a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py index 8d6d236b884..4d53ae7059d 100644 --- a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py +++ b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py @@ -46,8 +46,13 @@ class MCPGuardrailTranslationHandler(BaseTranslation): ) return data + inputs = GenericGuardrailAPIInputs(texts=[content]) + # Include model information if available + model = data.get("model") + if model: + inputs["model"] = model guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs=GenericGuardrailAPIInputs(texts=[content]), + inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7fceec005e6..92fd54e8775 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -11,7 +11,7 @@ import datetime import hashlib import json import re -from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast +from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast, Callable from urllib.parse import urlparse from fastapi import HTTPException @@ -1842,6 +1842,7 @@ class MCPServerManager: oauth2_headers: Optional[Dict[str, str]], raw_headers: Optional[Dict[str, str]], proxy_logging_obj: Optional[ProxyLogging], + host_progress_callback: Optional[Callable] = None, ) -> CallToolResult: """ Call a regular MCP tool using the MCP client. @@ -1926,7 +1927,7 @@ class MCPServerManager: ) async def _call_tool_via_client(client, params): - return await client.call_tool(params) + return await client.call_tool(params, host_progress_callback=host_progress_callback) tasks.append( asyncio.create_task(_call_tool_via_client(client, call_tool_params)) @@ -1963,6 +1964,8 @@ class MCPServerManager: proxy_logging_obj: Optional[ProxyLogging] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, + host_progress_callback: Optional[Callable] = None, + ) -> CallToolResult: """ Call a tool with the given name and arguments @@ -2038,6 +2041,7 @@ class MCPServerManager: oauth2_headers=oauth2_headers, raw_headers=raw_headers, proxy_logging_obj=proxy_logging_obj, + host_progress_callback=host_progress_callback, ) # For OpenAPI tools, await outside the client context diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index bd7870a0fa0..6d54c3871e5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -8,8 +8,7 @@ import contextlib from datetime import datetime import traceback import uuid -from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union, cast - +from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union, cast, Callable from fastapi import FastAPI, HTTPException from pydantic import AnyUrl, ConfigDict from starlette.types import Receive, Scope, Send @@ -128,8 +127,8 @@ if MCP_AVAILABLE: session_manager = StreamableHTTPSessionManager( app=server, event_store=None, - json_response=True, # Use JSON responses instead of SSE by default - stateless=True, + json_response=False, # enables SSE streaming + stateless=False, # enables session state ) # Create SSE session manager @@ -282,6 +281,30 @@ if MCP_AVAILABLE: verbose_logger.debug( f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}" ) + host_progress_callback = None + try: + host_ctx = server.request_context + if host_ctx and hasattr(host_ctx, 'meta') and host_ctx.meta: + host_token = getattr(host_ctx.meta, 'progressToken', None) + if host_token and hasattr(host_ctx, 'session') and host_ctx.session: + host_session = host_ctx.session + + async def forward_progress(progress: float, total: float | None): + """Forward progress notifications from external MCP to Host""" + try: + await host_session.send_progress_notification( + progress_token=host_token, + progress=progress, + total=total + ) + verbose_logger.debug(f"Forwarded progress {progress}/{total} to Host") + except Exception as e: + verbose_logger.error(f"Failed to forward progress to Host: {e}") + + host_progress_callback = forward_progress + verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...") + except Exception as e: + verbose_logger.warning(f"Could not capture host progress context: {e}") try: # Create a body date for logging body_data = {"name": name, "arguments": arguments} @@ -311,6 +334,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + host_progress_callback=host_progress_callback, **data, # for logging ) except BlockedPiiEntityError as e: @@ -1345,6 +1369,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, + host_progress_callback: Optional[Callable] = None, **kwargs: Any, ) -> CallToolResult: """ @@ -1442,6 +1467,7 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, + host_progress_callback=host_progress_callback, ) # Fall back to local tool registry with original name (legacy support) @@ -1689,6 +1715,7 @@ if MCP_AVAILABLE: oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, litellm_logging_obj: Optional[Any] = None, + host_progress_callback: Optional[Callable] = None, ) -> CallToolResult: """Handle tool execution for managed server tools""" # Import here to avoid circular import @@ -1704,6 +1731,7 @@ if MCP_AVAILABLE: oauth2_headers=oauth2_headers, raw_headers=raw_headers, proxy_logging_obj=proxy_logging_obj, + host_progress_callback=host_progress_callback, ) verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result) return call_tool_result diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 60f968d9be1..c854d81ec71 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -851,6 +851,13 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): aliases: Optional[dict] = {} object_permission: Optional[LiteLLM_ObjectPermissionBase] = None + @field_validator("max_budget", mode="before") + @classmethod + def check_max_budget(cls, v): + if v == "": + return None + return v + class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase): index_name: str @@ -3421,6 +3428,8 @@ class LitellmMetadataFromRequestHeaders(TypedDict, total=False): """ spend_logs_metadata: Optional[dict] + agent_id: Optional[str] + trace_id: Optional[str] class JWTKeyItem(TypedDict, total=False): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 5e0a211906e..61d9044f925 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -21,6 +21,8 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.caching.dual_cache import LimitedSizeOrderedDict from litellm.constants import ( + CLI_JWT_EXPIRATION_HOURS, + CLI_JWT_TOKEN_NAME, DEFAULT_IN_MEMORY_TTL, DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, DEFAULT_MAX_RECURSE_DEPTH, @@ -1602,7 +1604,10 @@ class ExperimentalUIJWTToken: user_info: LiteLLM_UserTable, team_id: Optional[str] = None ) -> str: """ - Generate a JWT token for CLI authentication with 24-hour expiration. + Generate a JWT token for CLI authentication with configurable expiration. + + The expiration time can be controlled via the LITELLM_CLI_JWT_EXPIRATION_HOURS + environment variable (defaults to 24 hours). Args: user_info: User information from the database @@ -1613,7 +1618,6 @@ class ExperimentalUIJWTToken: """ from datetime import timedelta - from litellm.constants import CLI_JWT_TOKEN_NAME from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) @@ -1621,8 +1625,8 @@ class ExperimentalUIJWTToken: if user_info.user_role is None: raise Exception("User role is required for CLI JWT login") - # Calculate expiration time (24 hours from now - matching old CLI key behavior) - expiration_time = get_utc_datetime() + timedelta(hours=24) + # Calculate expiration time (configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS env var) + expiration_time = get_utc_datetime() + timedelta(hours=CLI_JWT_EXPIRATION_HOURS) # Format the expiration time as ISO 8601 string expires = expiration_time.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "+00:00" diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 5be44f479b8..939cfefadcc 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -197,9 +197,9 @@ async def authenticate_user( # noqa: PLR0915 - Login with UI_USERNAME and UI_PASSWORD - Login with Invite Link `user_email` and `password` combination """ - if secrets.compare_digest(username, ui_username) and secrets.compare_digest( - password, ui_password - ): + if secrets.compare_digest( + username.encode("utf-8"), ui_username.encode("utf-8") + ) and secrets.compare_digest(password.encode("utf-8"), ui_password.encode("utf-8")): # Non SSO -> If user is using UI_USERNAME and UI_PASSWORD they are Proxy admin user_role = LitellmUserRoles.PROXY_ADMIN user_id = LITELLM_PROXY_ADMIN_NAME @@ -313,9 +313,9 @@ async def authenticate_user( # noqa: PLR0915 # check if password == _user_row.password hash_password = hash_token(token=password) - if secrets.compare_digest(password, _password) or secrets.compare_digest( - hash_password, _password - ): + if secrets.compare_digest( + password.encode("utf-8"), _password.encode("utf-8") + ) or secrets.compare_digest(hash_password.encode("utf-8"), _password.encode("utf-8")): if os.getenv("DATABASE_URL") is not None: # Expire any previous UI session tokens for this user await expire_previous_ui_session_tokens( diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index bdc8d56d1c3..2345b0263e3 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -11,6 +11,8 @@ import requests from rich.console import Console from rich.table import Table +from litellm.constants import CLI_JWT_EXPIRATION_HOURS + # Token storage utilities def get_token_file_path() -> str: @@ -592,8 +594,8 @@ def whoami(): age_hours = (time.time() - timestamp) / 3600 click.echo(f"Token age: {age_hours:.1f} hours") - if age_hours > 24: - click.echo("⚠️ Warning: Token is more than 24 hours old and may have expired.") + if age_hours > CLI_JWT_EXPIRATION_HOURS: + click.echo(f"⚠️ Warning: Token is more than {CLI_JWT_EXPIRATION_HOURS} hours old and may have expired.") # Export functions for use by other CLI commands diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index c0cff84bacd..faeca9b2aed 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -274,11 +274,20 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 WebSearchInterceptionLogger, ) - websearch_interception_obj = WebSearchInterceptionLogger.initialize_from_proxy_config( - litellm_settings=litellm_settings, - callback_specific_params=callback_specific_params, + websearch_interception_obj = ( + WebSearchInterceptionLogger.initialize_from_proxy_config( + litellm_settings=litellm_settings, + callback_specific_params=callback_specific_params, + ) ) imported_list.append(websearch_interception_obj) + elif isinstance(callback, str) and callback == "datadog_cost_management": + from litellm.integrations.datadog.datadog_cost_management import ( + DatadogCostManagementLogger, + ) + + datadog_cost_management_obj = DatadogCostManagementLogger() + imported_list.append(datadog_cost_management_obj) elif isinstance(callback, CustomLogger): imported_list.append(callback) else: @@ -353,17 +362,17 @@ def get_remaining_tokens_and_requests_from_request_data(data: Dict) -> Dict[str, remaining_requests_variable_name = f"litellm-key-remaining-requests-{model_group}" remaining_requests = _metadata.get(remaining_requests_variable_name, None) if remaining_requests: - headers[f"x-litellm-key-remaining-requests-{h11_model_group_name}"] = ( - remaining_requests - ) + headers[ + f"x-litellm-key-remaining-requests-{h11_model_group_name}" + ] = remaining_requests # Remaining Tokens remaining_tokens_variable_name = f"litellm-key-remaining-tokens-{model_group}" remaining_tokens = _metadata.get(remaining_tokens_variable_name, None) if remaining_tokens: - headers[f"x-litellm-key-remaining-tokens-{h11_model_group_name}"] = ( - remaining_tokens - ) + headers[ + f"x-litellm-key-remaining-tokens-{h11_model_group_name}" + ] = remaining_tokens return headers @@ -438,9 +447,9 @@ def add_guardrail_response_to_standard_logging_object( ): if litellm_logging_obj is None: return - standard_logging_object: Optional[StandardLoggingPayload] = ( - litellm_logging_obj.model_call_details.get("standard_logging_object") - ) + standard_logging_object: Optional[ + StandardLoggingPayload + ] = litellm_logging_obj.model_call_details.get("standard_logging_object") if standard_logging_object is None: return guardrail_information = standard_logging_object.get("guardrail_information", []) @@ -469,7 +478,9 @@ def get_metadata_variable_name_from_kwargs( return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata" -def process_callback(_callback: str, callback_type: str, environment_variables: dict) -> dict: +def process_callback( + _callback: str, callback_type: str, environment_variables: dict +) -> dict: """Process a single callback and return its data with environment variables""" env_vars = CustomLogger.get_callback_env_vars(_callback) @@ -481,11 +492,9 @@ def process_callback(_callback: str, callback_type: str, environment_variables: else: env_vars_dict[_var] = env_variable - return { - "name": _callback, - "variables": env_vars_dict, - "type": callback_type - } + return {"name": _callback, "variables": env_vars_dict, "type": callback_type} + + def normalize_callback_names(callbacks: Iterable[Any]) -> List[Any]: if callbacks is None: return [] diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index d5a62661460..b37074e25e7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -185,6 +185,7 @@ class GenericGuardrailAPI(CustomGuardrail): tools = inputs.get("tools") structured_messages = inputs.get("structured_messages") tool_calls = inputs.get("tool_calls") + model = inputs.get("model") # Use provided request_data or create an empty dict if request_data is None: @@ -215,6 +216,7 @@ class GenericGuardrailAPI(CustomGuardrail): tool_calls=tool_calls, additional_provider_specific_params=additional_params, input_type=input_type, + model=model, ) # Prepare headers diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 02dcb25c82f..9be78264e85 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -558,6 +558,16 @@ class LiteLLMProxyRequestSetup: ######################################################################################### # Finally update the requests metadata with the `metadata_from_headers` ######################################################################################### + agent_id_from_header = headers.get("x-litellm-agent-id") + trace_id_from_header = headers.get("x-litellm-trace-id") + if agent_id_from_header: + metadata_from_headers["agent_id"] = agent_id_from_header + verbose_proxy_logger.debug(f"Extracted agent_id from header: {agent_id_from_header}") + + if trace_id_from_header: + metadata_from_headers["trace_id"] = trace_id_from_header + verbose_proxy_logger.debug(f"Extracted trace_id from header: {trace_id_from_header}") + if isinstance(data[_metadata_variable_name], dict): data[_metadata_variable_name].update(metadata_from_headers) return data diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 9a986326c03..38a867d031b 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -814,7 +814,9 @@ def _update_internal_user_params( ) -> dict: non_default_values = {} for k, v in data_json.items(): - if ( + if k == "max_budget": + non_default_values[k] = v + elif ( v is not None and v not in ( diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 5adf54c1627..1b827411e5a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -99,24 +99,24 @@ def determine_role_from_groups( ) -> Optional[LitellmUserRoles]: """ Determine the highest privilege role for a user based on their groups. - + Role hierarchy (highest to lowest): - proxy_admin - proxy_admin_viewer - internal_user - internal_user_viewer - + Args: user_groups: List of group names from the SSO token role_mappings: RoleMappings configuration object - + Returns: The highest privilege role found, or default_role if no matches, or None """ if not role_mappings.roles: # No role mappings configured, return default_role return role_mappings.default_role - + # Role hierarchy (highest to lowest) role_hierarchy = [ LitellmUserRoles.PROXY_ADMIN, @@ -124,20 +124,22 @@ def determine_role_from_groups( LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - + # Convert user_groups to a set for efficient lookup user_groups_set = set(user_groups) if isinstance(user_groups, list) else set() - + # Find the highest privilege role the user belongs to for role in role_hierarchy: if role in role_mappings.roles: role_groups = role_mappings.roles[role] - if isinstance(role_groups, list) and user_groups_set.intersection(set(role_groups)): + if isinstance(role_groups, list) and user_groups_set.intersection( + set(role_groups) + ): verbose_proxy_logger.debug( f"User groups {user_groups} matched role '{role.value}' via groups: {role_groups}" ) return role - + # No matching groups found, return default_role verbose_proxy_logger.debug( f"User groups {user_groups} did not match any role mappings, using default_role: {role_mappings.default_role}" @@ -326,9 +328,7 @@ def generic_response_convertor( "GENERIC_USER_PROVIDER_ATTRIBUTE", "provider" ) - generic_user_role_attribute_name = os.getenv( - "GENERIC_USER_ROLE_ATTRIBUTE", "role" - ) + generic_user_role_attribute_name = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role") verbose_proxy_logger.debug( f" generic_user_id_attribute_name: {generic_user_id_attribute_name}\n generic_user_email_attribute_name: {generic_user_email_attribute_name}" @@ -345,12 +345,15 @@ def generic_response_convertor( # Determine user role based on role_mappings if available # Only apply role_mappings for GENERIC SSO provider user_role: Optional[LitellmUserRoles] = None - - if role_mappings is not None and role_mappings.provider.lower() in ["generic", "okta"]: + + if role_mappings is not None and role_mappings.provider.lower() in [ + "generic", + "okta", + ]: # Use role_mappings to determine role from groups group_claim = role_mappings.group_claim user_groups_raw: Any = get_nested_value(response, group_claim) - + # Handle different formats: could be a list, string (comma-separated), or single value user_groups: List[str] = [] if isinstance(user_groups_raw, list): @@ -361,7 +364,7 @@ def generic_response_convertor( elif user_groups_raw is not None: # Single value user_groups = [str(user_groups_raw)] - + if user_groups: user_role = determine_role_from_groups(user_groups, role_mappings) verbose_proxy_logger.debug( @@ -373,10 +376,12 @@ def generic_response_convertor( verbose_proxy_logger.debug( f"No groups found in '{group_claim}', using default_role: {role_mappings.default_role}" ) - + # Fallback to existing logic if role_mappings not used if user_role is None: - user_role_from_sso = get_nested_value(response, generic_user_role_attribute_name) + user_role_from_sso = get_nested_value( + response, generic_user_role_attribute_name + ) if user_role_from_sso is not None: role = get_litellm_user_role(user_role_from_sso) if role is not None: @@ -399,7 +404,9 @@ def generic_response_convertor( ) -def _setup_generic_sso_env_vars(generic_client_id: str, redirect_url: str) -> Tuple[str, List[str], str, str, str, bool]: +def _setup_generic_sso_env_vars( + generic_client_id: str, redirect_url: str +) -> Tuple[str, List[str], str, str, str, bool]: """Setup and validate Generic SSO environment variables.""" generic_client_secret = os.getenv("GENERIC_CLIENT_SECRET", None) generic_scope = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ") @@ -492,7 +499,43 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: verbose_proxy_logger.debug( f"Could not load role_mappings from database: {e}. Continuing with existing role logic." ) + + generic_role_mappings = os.getenv("GENERIC_ROLE_MAPPINGS_ROLES", None) + generic_role_mappings_group_claim = os.getenv( + "GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None + ) + generic_role_mappoings_default_role = os.getenv( + "GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None + ) + if generic_role_mappings is not None: + verbose_proxy_logger.debug( + "Found role_mappings for generic provider in environment variables" + ) + import ast + try: + generic_user_role_mappings_data: Dict[ + LitellmUserRoles, List[str] + ] = ast.literal_eval(generic_role_mappings) + if isinstance(generic_user_role_mappings_data, dict): + from litellm.types.proxy.management_endpoints.ui_sso import ( + RoleMappings, + ) + + role_mappings_data = { + "provider": "generic", + "group_claim": generic_role_mappings_group_claim, + "default_role": generic_role_mappoings_default_role, + "roles": generic_user_role_mappings_data, + } + + role_mappings = RoleMappings(**role_mappings_data) + verbose_proxy_logger.debug( + f"Loaded role_mappings from environments for provider '{role_mappings.provider}'." + ) + return role_mappings + except TypeError as e: + verbose_proxy_logger.warning(f"Error decoding role mappings from environment variables: {e}. Continuing with existing role logic.") return role_mappings @@ -529,7 +572,7 @@ async def get_generic_sso_response( # Get role_mappings from SSO settings if available role_mappings = await _setup_role_mappings() - + def response_convertor(response, client): nonlocal received_response # return for user debugging received_response = response @@ -1156,7 +1199,7 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None): max_budget=litellm.max_ui_session_budget, ) - # Generate CLI JWT on-demand (24hr expiration) + # Generate CLI JWT on-demand (expiration configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS) # Pass selected team_id to ensure JWT has correct team jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( user_info=user_info, team_id=team_id @@ -1217,20 +1260,24 @@ async def insert_sso_user( role_mappings_configured = False try: from litellm.proxy.utils import get_prisma_client_or_throw - + prisma_client = get_prisma_client_or_throw( "Prisma client is None, connect a database to your proxy" ) - + # Get SSO config from dedicated table sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique( where={"id": "sso_config"} ) - + if sso_db_record and sso_db_record.sso_settings: sso_settings_dict = dict(sso_db_record.sso_settings) role_mappings_data = sso_settings_dict.get("role_mappings") role_mappings_configured = role_mappings_data is not None + generic_user_role_mappings = os.getenv("GENERIC_USER_ROLE_MAPPINGS", None) + if generic_user_role_mappings is not None: + role_mappings_configured = True + except Exception as e: # If we can't check role_mappings, continue with existing logic verbose_proxy_logger.debug( @@ -1240,7 +1287,10 @@ async def insert_sso_user( # Apply default_internal_user_params if litellm.default_internal_user_params: # If role_mappings is configured and user_role is already set from SSO, preserve it - if role_mappings_configured and user_defined_values.get("user_role") is not None: + if ( + role_mappings_configured + and user_defined_values.get("user_role") is not None + ): # Preserve the SSO-extracted role, but apply other defaults preserved_role = user_defined_values.get("user_role") user_defined_values.update(litellm.default_internal_user_params) # type: ignore diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bcd51c36632..183c25ed463 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -825,6 +825,7 @@ app = FastAPI( title=_title, description=_description, version=version, + root_path=server_root_path, lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues] ) @@ -5419,6 +5420,24 @@ async def chat_completion( # noqa: PLR0915 global general_settings, user_debug, proxy_logging_obj, llm_model_list global user_temperature, user_request_timeout, user_max_tokens, user_api_base data = await _read_request_body(request=request) + if user_api_key_dict is not None: + if data.get("metadata") is None: + data["metadata"] = {} + if ( + hasattr(user_api_key_dict, "user_id") + and user_api_key_dict.user_id is not None + ): + data["metadata"]["user_api_key_user_id"] = user_api_key_dict.user_id + if ( + hasattr(user_api_key_dict, "team_id") + and user_api_key_dict.team_id is not None + ): + data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id + if ( + hasattr(user_api_key_dict, "org_id") + and user_api_key_dict.org_id is not None + ): + data["metadata"]["user_api_key_org_id"] = user_api_key_dict.org_id base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) try: result = await base_llm_response_processor.base_process_llm_request( @@ -5570,6 +5589,24 @@ async def completion( # noqa: PLR0915 data = {} try: data = await _read_request_body(request=request) + if user_api_key_dict is not None: + if data.get("metadata") is None: + data["metadata"] = {} + if ( + hasattr(user_api_key_dict, "user_id") + and user_api_key_dict.user_id is not None + ): + data["metadata"]["user_api_key_user_id"] = user_api_key_dict.user_id + if ( + hasattr(user_api_key_dict, "team_id") + and user_api_key_dict.team_id is not None + ): + data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id + if ( + hasattr(user_api_key_dict, "org_id") + and user_api_key_dict.org_id is not None + ): + data["metadata"]["user_api_key_org_id"] = user_api_key_dict.org_id base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) return await base_llm_response_processor.base_process_llm_request( request=request, @@ -5789,6 +5826,25 @@ async def embeddings( # noqa: PLR0915 ) data["input"] = input_list + if user_api_key_dict is not None: + if data.get("metadata") is None: + data["metadata"] = {} + if ( + hasattr(user_api_key_dict, "user_id") + and user_api_key_dict.user_id is not None + ): + data["metadata"]["user_api_key_user_id"] = user_api_key_dict.user_id + if ( + hasattr(user_api_key_dict, "team_id") + and user_api_key_dict.team_id is not None + ): + data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id + if ( + hasattr(user_api_key_dict, "org_id") + and user_api_key_dict.org_id is not None + ): + data["metadata"]["user_api_key_org_id"] = user_api_key_dict.org_id + # Use unified request processor (same as chat/completions and responses) base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 79b4fd6873d..79f182817f6 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -26,6 +26,184 @@ from litellm.proxy.common_utils.http_parsing_utils import ( router = APIRouter() +def _build_file_metadata_entry( + response: Any, + file_data: Optional[Tuple[str, bytes, str]] = None, + file_url: Optional[str] = None, +) -> Dict[str, Any]: + """ + Build a file metadata entry for storing in vector_store_metadata. + + Args: + response: The response from litellm.aingest containing file_id + file_data: Optional tuple of (filename, content, content_type) + file_url: Optional URL if file was ingested from URL + + Returns: + Dictionary with file metadata (file_id, filename, file_url, ingested_at, etc.) + """ + from datetime import datetime, timezone + + # Extract file_id from response + file_id = None + if hasattr(response, "get"): + file_id = response.get("file_id") + elif hasattr(response, "file_id"): + file_id = response.file_id + + # Extract file information from file_data tuple + filename = None + file_size = None + content_type = None + + if file_data: + filename = file_data[0] + file_size = len(file_data[1]) if len(file_data) > 1 else None + content_type = file_data[2] if len(file_data) > 2 else None + + # Build file metadata entry + file_entry = { + "file_id": file_id, + "filename": filename, + "file_url": file_url, + "ingested_at": datetime.now(timezone.utc).isoformat(), + } + + # Add optional fields if available + if file_size is not None: + file_entry["file_size"] = file_size + if content_type is not None: + file_entry["content_type"] = content_type + + return file_entry + + +async def _save_vector_store_to_db_from_rag_ingest( + response: Any, + ingest_options: Dict[str, Any], + prisma_client, + user_api_key_dict: UserAPIKeyAuth, + file_data: Optional[Tuple[str, bytes, str]] = None, + file_url: Optional[str] = None, +) -> None: + """ + Helper function to save a newly created vector store from RAG ingest to the database. + + This function: + - Extracts vector store ID and config from the ingest response + - Checks if the vector store already exists in the database + - Creates a new database entry if it doesn't exist + - Adds the vector store to the registry + + Args: + response: The response from litellm.aingest() + ingest_options: The ingest options containing vector store config + prisma_client: The Prisma database client + user_api_key_dict: User API key authentication info + """ + from litellm.proxy.vector_store_endpoints.management_endpoints import ( + create_vector_store_in_db, + ) + + # Handle both dict and object responses + if hasattr(response, "get"): + vector_store_id = response.get("vector_store_id") + elif hasattr(response, "vector_store_id"): + vector_store_id = response.vector_store_id + else: + verbose_proxy_logger.warning( + f"Unable to extract vector_store_id from response type: {type(response)}" + ) + return + + if vector_store_id is None or not isinstance(vector_store_id, str): + verbose_proxy_logger.warning( + "Vector store ID is None or not a string, skipping database save" + ) + return + + vector_store_config = ingest_options.get("vector_store", {}) + custom_llm_provider = vector_store_config.get("custom_llm_provider") + + # Extract litellm_vector_store_params for custom name and description + litellm_vector_store_params = ingest_options.get("litellm_vector_store_params", {}) + custom_vector_store_name = litellm_vector_store_params.get("vector_store_name") + custom_vector_store_description = litellm_vector_store_params.get("vector_store_description") + + # Build file metadata entry using helper + file_entry = _build_file_metadata_entry( + response=response, + file_data=file_data, + file_url=file_url, + ) + + try: + # Check if vector store already exists in database + existing_vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.find_unique( + where={"vector_store_id": vector_store_id} + ) + ) + + # Only create if it doesn't exist + if existing_vector_store is None: + verbose_proxy_logger.info( + f"Saving newly created vector store {vector_store_id} to database" + ) + + # Initialize metadata with first file + initial_metadata = { + "ingested_files": [file_entry] + } + + # Use custom name if provided, otherwise default + vector_store_name = custom_vector_store_name or f"RAG Vector Store - {vector_store_id[:8]}" + vector_store_description = custom_vector_store_description or "Created via RAG ingest endpoint" + + await create_vector_store_in_db( + vector_store_id=vector_store_id, + custom_llm_provider=custom_llm_provider or "openai", + prisma_client=prisma_client, + vector_store_name=vector_store_name, + vector_store_description=vector_store_description, + vector_store_metadata=initial_metadata, + ) + + verbose_proxy_logger.info( + f"Vector store {vector_store_id} saved to database successfully" + ) + else: + verbose_proxy_logger.info( + f"Vector store {vector_store_id} already exists, appending file to metadata" + ) + + # Update existing vector store with new file + existing_metadata = existing_vector_store.vector_store_metadata or {} + if isinstance(existing_metadata, str): + import json + existing_metadata = json.loads(existing_metadata) + + ingested_files = existing_metadata.get("ingested_files", []) + ingested_files.append(file_entry) + existing_metadata["ingested_files"] = ingested_files + + # Update the vector store + from litellm.proxy.utils import safe_dumps + await prisma_client.db.litellm_managedvectorstorestable.update( + where={"vector_store_id": vector_store_id}, + data={"vector_store_metadata": safe_dumps(existing_metadata)} + ) + + verbose_proxy_logger.info( + f"Added file {file_entry.get('filename') or file_entry.get('file_url', 'Unknown')} to vector store {vector_store_id} metadata" + ) + except Exception as db_error: + # Log the error but don't fail the request since ingestion succeeded + verbose_proxy_logger.exception( + f"Failed to save vector store {vector_store_id} to database: {db_error}" + ) + + async def parse_rag_ingest_request( request: Request, ) -> Tuple[Dict[str, Any], Optional[Tuple[str, bytes, str]], Optional[str], Optional[str]]: @@ -158,6 +336,7 @@ async def rag_ingest( add_litellm_data_to_request, general_settings, llm_router, + prisma_client, proxy_config, version, ) @@ -189,6 +368,25 @@ async def rag_ingest( **request_data, ) + # Save vector store to database if it was newly created and prisma_client is available + verbose_proxy_logger.debug( + f"RAG Ingest - Checking database save conditions: prisma_client={prisma_client is not None}, response={response is not None}, response_type={type(response)}" + ) + + if prisma_client is not None and response is not None: + await _save_vector_store_to_db_from_rag_ingest( + response=response, + ingest_options=ingest_options, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + file_data=file_data, + file_url=file_url, + ) + else: + verbose_proxy_logger.warning( + f"Skipping database save: prisma_client={prisma_client is not None}, response={response is not None}" + ) + return response except HTTPException: diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 5d348d35cb2..db4e4beec21 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -396,7 +396,7 @@ def get_logging_payload( # noqa: PLR0915 ) # Extract agent_id for A2A requests (set directly on model_call_details) - agent_id: Optional[str] = kwargs.get("agent_id") + agent_id: Optional[str] = kwargs.get("agent_id") or metadata.get("agent_id") custom_llm_provider = kwargs.get("custom_llm_provider") raw_model = cast(str, kwargs.get("model") or "") model_name = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {}) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5dcb0339739..13f42a2f71d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -982,7 +982,9 @@ class ProxyLogging: try: # Check if load balancing should be used - if guardrail_name and self._should_use_guardrail_load_balancing(guardrail_name): + if guardrail_name and self._should_use_guardrail_load_balancing( + guardrail_name + ): response = await self._execute_guardrail_with_load_balancing( guardrail_name=guardrail_name, hook_type="pre_call", @@ -1017,7 +1019,11 @@ class ProxyLogging: latency_seconds = guardrail_end_time - guardrail_start_time # Get guardrail name for metrics (fallback if not set) - metrics_guardrail_name = guardrail_name or getattr(callback, "guardrail_name", callback.__class__.__name__) or "unknown" + metrics_guardrail_name = ( + guardrail_name + or getattr(callback, "guardrail_name", callback.__class__.__name__) + or "unknown" + ) # Find PrometheusLogger in callbacks and record metrics for prom_callback in litellm.callbacks: @@ -1793,9 +1799,11 @@ class ProxyLogging: ################################################################# for callback in other_callbacks: - await callback.async_post_call_success_hook( + callback_response = await callback.async_post_call_success_hook( user_api_key_dict=user_api_key_dict, data=data, response=response ) + if callback_response is not None: + response = callback_response except Exception as e: raise e return response @@ -1852,16 +1860,14 @@ class ProxyLogging: complete_response = str_so_far + response_str else: complete_response = response_str - potential_error_response = ( + callback_response = ( await _callback.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, response=complete_response, ) ) - if isinstance( - potential_error_response, str - ) and potential_error_response.startswith("data: "): - return potential_error_response + if callback_response is not None: + response = callback_response except Exception as e: raise e return response @@ -2218,17 +2224,17 @@ class PrismaClient: ) -> Optional[dict]: """ Execute a query with automatic fallback for PostgreSQL cached plan errors. - + This handles the "cached plan must not change result type" error that occurs during rolling deployments when schema changes are applied while old pods still have cached query plans expecting the old schema. - + Args: sql_query: SQL query string to execute - + Returns: Query result or None - + Raises: Original exception if not a cached plan error """ @@ -2241,7 +2247,7 @@ class PrismaClient: # Add a unique comment to make the query different sql_query_retry = sql_query.replace( "SELECT", - f"SELECT /* cache_invalidated_{int(time.time() * 1000)} */" + f"SELECT /* cache_invalidated_{int(time.time() * 1000)} */", ) verbose_proxy_logger.warning( "PostgreSQL cached plan error detected for token lookup, " @@ -2583,7 +2589,9 @@ class PrismaClient: WHERE v.token = '{token}' """ - response = await self._query_first_with_cached_plan_fallback(sql_query) + response = await self._query_first_with_cached_plan_fallback( + sql_query + ) if response is not None: if response["team_models"] is None: @@ -4227,7 +4235,7 @@ def get_server_root_path() -> str: - If SERVER_ROOT_PATH is set, return it. - Otherwise, default to "/". """ - return os.getenv("SERVER_ROOT_PATH", "/") + return os.getenv("SERVER_ROOT_PATH", "") def get_prisma_client_or_throw(message: str): diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index bc61a60fe5a..f3787e62f4d 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -133,6 +133,112 @@ async def _resolve_embedding_config_from_db( return None +######################################################## +# Helper Functions +######################################################## +async def create_vector_store_in_db( + vector_store_id: str, + custom_llm_provider: str, + prisma_client, + vector_store_name: Optional[str] = None, + vector_store_description: Optional[str] = None, + vector_store_metadata: Optional[Dict] = None, + litellm_params: Optional[Dict] = None, + litellm_credential_name: Optional[str] = None, +) -> LiteLLM_ManagedVectorStore: + """ + Helper function to create a vector store in the database. + + This function handles: + - Checking if vector store already exists + - Creating the vector store in the database + - Adding it to the vector store registry + + Returns: + LiteLLM_ManagedVectorStore: The created vector store object + + Raises: + HTTPException: If vector store already exists or database error occurs + """ + from litellm.types.router import GenericLiteLLMParams + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + # Check if vector store already exists + existing_vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.find_unique( + where={"vector_store_id": vector_store_id} + ) + ) + if existing_vector_store is not None: + raise HTTPException( + status_code=400, + detail=f"Vector store with ID {vector_store_id} already exists", + ) + + # Prepare data for database + data_to_create: Dict[str, Any] = { + "vector_store_id": vector_store_id, + "custom_llm_provider": custom_llm_provider, + } + + if vector_store_name is not None: + data_to_create["vector_store_name"] = vector_store_name + if vector_store_description is not None: + data_to_create["vector_store_description"] = vector_store_description + if vector_store_metadata is not None: + data_to_create["vector_store_metadata"] = safe_dumps(vector_store_metadata) + if litellm_credential_name is not None: + data_to_create["litellm_credential_name"] = litellm_credential_name + + # Handle litellm_params - always provide at least an empty dict + if litellm_params: + # Auto-resolve embedding config if embedding model is provided but config is not + embedding_model = litellm_params.get("litellm_embedding_model") + if embedding_model and not litellm_params.get("litellm_embedding_config"): + resolved_config = await _resolve_embedding_config_from_db( + embedding_model=embedding_model, + prisma_client=prisma_client + ) + if resolved_config: + litellm_params["litellm_embedding_config"] = resolved_config + verbose_proxy_logger.info( + f"Auto-resolved embedding config for model {embedding_model}" + ) + + litellm_params_dict = GenericLiteLLMParams( + **litellm_params + ).model_dump(exclude_none=True) + data_to_create["litellm_params"] = safe_dumps(litellm_params_dict) + else: + # Provide empty dict if no litellm_params provided + data_to_create["litellm_params"] = safe_dumps({}) + + # Create in database + _new_vector_store = ( + await prisma_client.db.litellm_managedvectorstorestable.create( + data=data_to_create + ) + ) + + new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore( + **_new_vector_store.model_dump() + ) + + # Add vector store to registry + if litellm.vector_store_registry is not None: + litellm.vector_store_registry.add_vector_store_to_registry( + vector_store=new_vector_store + ) + + verbose_proxy_logger.info( + f"Vector store {vector_store_id} created in database successfully" + ) + + return new_vector_store + + ######################################################## # Management Endpoints ######################################################## @@ -156,71 +262,34 @@ async def new_vector_store( - vector_store_metadata: Optional[Dict] - Additional metadata for the vector store """ from litellm.proxy.proxy_server import prisma_client - from litellm.types.router import GenericLiteLLMParams - - if prisma_client is None: - raise HTTPException(status_code=500, detail="Database not connected") try: - # Check if vector store already exists - existing_vector_store = ( - await prisma_client.db.litellm_managedvectorstorestable.find_unique( - where={"vector_store_id": vector_store.get("vector_store_id")} - ) - ) - if existing_vector_store is not None: + vector_store_id = vector_store.get("vector_store_id") + custom_llm_provider = vector_store.get("custom_llm_provider") + + if not vector_store_id or not custom_llm_provider: raise HTTPException( status_code=400, - detail=f"Vector store with ID {vector_store.get('vector_store_id')} already exists", - ) - - if vector_store.get("vector_store_metadata") is not None: - vector_store["vector_store_metadata"] = safe_dumps( - vector_store.get("vector_store_metadata") - ) - - # Safely handle JSON serialization of litellm_params - litellm_params_json: Optional[str] = None - _input_litellm_params: dict = vector_store.get("litellm_params", {}) or {} - if _input_litellm_params is not None: - # Auto-resolve embedding config if embedding model is provided but config is not - embedding_model = _input_litellm_params.get("litellm_embedding_model") - if embedding_model and not _input_litellm_params.get("litellm_embedding_config"): - resolved_config = await _resolve_embedding_config_from_db( - embedding_model=embedding_model, - prisma_client=prisma_client - ) - if resolved_config: - _input_litellm_params["litellm_embedding_config"] = resolved_config - verbose_proxy_logger.info( - f"Auto-resolved embedding config for model {embedding_model}" - ) - - litellm_params_dict = GenericLiteLLMParams( - **_input_litellm_params - ).model_dump(exclude_none=True) - litellm_params_json = safe_dumps(litellm_params_dict) - del vector_store["litellm_params"] - - _new_vector_store = ( - await prisma_client.db.litellm_managedvectorstorestable.create( - data={ - **vector_store, - "litellm_params": litellm_params_json, - } + detail="vector_store_id and custom_llm_provider are required" ) + + # Extract and validate metadata + metadata = vector_store.get("vector_store_metadata") + validated_metadata: Optional[Dict] = None + if metadata is not None and isinstance(metadata, dict): + validated_metadata = metadata + + new_vector_store = await create_vector_store_in_db( + vector_store_id=vector_store_id, + custom_llm_provider=custom_llm_provider, + prisma_client=prisma_client, + vector_store_name=vector_store.get("vector_store_name"), + vector_store_description=vector_store.get("vector_store_description"), + vector_store_metadata=validated_metadata, + litellm_params=vector_store.get("litellm_params"), + litellm_credential_name=vector_store.get("litellm_credential_name"), ) - new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore( - **_new_vector_store.model_dump() - ) - - # Add vector store to registry - if litellm.vector_store_registry is not None: - litellm.vector_store_registry.add_vector_store_to_registry( - vector_store=new_vector_store - ) - return { "status": "success", "message": f"Vector store {vector_store.get('vector_store_id')} created successfully", diff --git a/litellm/rag/main.py b/litellm/rag/main.py index 0ccbb435e53..620097f83dd 100644 --- a/litellm/rag/main.py +++ b/litellm/rag/main.py @@ -198,6 +198,10 @@ async def _execute_query_pipeline( """ Execute the RAG query pipeline. """ + # Extract router from kwargs - use it for completion if available + # to properly resolve virtual model names + router: Optional["Router"] = kwargs.pop("router", None) + # 1. Extract query from last user message query_text = RAGQuery.extract_query_from_messages(messages) if not query_text: @@ -233,12 +237,21 @@ async def _execute_query_pipeline( context_message = RAGQuery.build_context_message(context_chunks) modified_messages = messages[:-1] + [context_message] + [messages[-1]] - response = await litellm.acompletion( - model=model, - messages=modified_messages, - stream=stream, - **kwargs, - ) + # Use router if available to properly resolve virtual model names + if router is not None: + response = await router.acompletion( + model=model, + messages=modified_messages, + stream=stream, + **kwargs, + ) + else: + response = await litellm.acompletion( + model=model, + messages=modified_messages, + stream=stream, + **kwargs, + ) # 5. Attach search results to response if not stream and isinstance(response, ModelResponse): diff --git a/litellm/router.py b/litellm/router.py index 54650b120a7..9fe37efa3ad 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1707,8 +1707,11 @@ class Router: litellm_params = deployment.get("litellm_params", {}) dep_num_retries = litellm_params.get("num_retries") - if dep_num_retries is not None and isinstance(dep_num_retries, int): - exception.num_retries = dep_num_retries # type: ignore + if dep_num_retries is not None: + try: + exception.num_retries = int(dep_num_retries) # type: ignore # Handle both int and str + except (ValueError, TypeError): + pass # Skip if value can't be converted to int def _update_kwargs_with_default_litellm_params( self, kwargs: dict, metadata_variable_name: Optional[str] = "metadata" diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index a093fe2d2fd..38405f058c3 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -17,6 +17,22 @@ from litellm.secret_managers.secret_manager_handler import get_secret_from_manag oidc_cache = DualCache() +def _get_oidc_http_handler(timeout: Optional[httpx.Timeout] = None) -> HTTPHandler: + """ + Factory function to create HTTPHandler for OIDC requests. + This function can be mocked in tests. + + Args: + timeout: Optional timeout for HTTP requests. Defaults to 600.0 seconds with 5.0 connect timeout. + + Returns: + HTTPHandler instance configured for OIDC requests. + """ + if timeout is None: + timeout = httpx.Timeout(timeout=600.0, connect=5.0) + return HTTPHandler(timeout=timeout) + + ######### Secret Manager ############################ # checks if user has passed in a secret manager client # if passed in then checks the secret there @@ -103,7 +119,7 @@ def get_secret( # noqa: PLR0915 if oidc_token is not None: return oidc_token - oidc_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) + oidc_client = _get_oidc_http_handler() # https://cloud.google.com/compute/docs/instances/verifying-instance-identity#request_signature response = oidc_client.get( "http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity", @@ -141,7 +157,7 @@ def get_secret( # noqa: PLR0915 if oidc_token is not None: return oidc_token - oidc_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) + oidc_client = _get_oidc_http_handler() response = oidc_client.get( actions_id_token_request_url, params={"audience": oidc_aud}, diff --git a/litellm/types/integrations/datadog_cost_management.py b/litellm/types/integrations/datadog_cost_management.py new file mode 100644 index 00000000000..fe04f43ea03 --- /dev/null +++ b/litellm/types/integrations/datadog_cost_management.py @@ -0,0 +1,27 @@ +from typing import Dict, Optional, TypedDict + + +from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams + + +class DatadogCostManagementInitParams(StandardCustomLoggerInitParams): + """ + Init params for Datadog Cost Management + """ + + datadog_cost_management_params: Optional[Dict] = None + + +class DatadogFOCUSCostEntry(TypedDict): + """ + Represents a single cost line item in the FOCUS format. + Ref: https://focus.finops.org/#specification + """ + + ProviderName: str + ChargeDescription: str + ChargePeriodStart: str + ChargePeriodEnd: str + BilledCost: float + BillingCurrency: str + Tags: Optional[Dict[str, str]] diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index ef2f1ba4d5e..a85aaafe23d 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -93,6 +93,67 @@ class GuardrailConverseContentBlock(TypedDict, total=False): text: GuardrailConverseTextBlock +class CitationWebLocationBlock(TypedDict, total=False): + """ + Web location block for Nova grounding citations. + Contains the URL and domain from web search results. + + Reference: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html + """ + + url: str + domain: str + + +class CitationLocationBlock(TypedDict, total=False): + """ + Location block containing the web location for a citation. + """ + + web: CitationWebLocationBlock + + +class CitationReferenceBlock(TypedDict, total=False): + """ + Citation reference block containing a single citation with its location. + + Each citation contains: + - location.web.url: The URL of the source + - location.web.domain: The domain of the source + """ + + location: CitationLocationBlock + + +class CitationsContentBlock(TypedDict, total=False): + """ + Citations content block returned by Nova grounding (web search) tool. + + When Nova grounding is enabled via systemTool, the model may return + citationsContent blocks containing web search citation references. + + Reference: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html + + Example response structure: + { + "citationsContent": { + "citations": [ + { + "location": { + "web": { + "url": "https://example.com/article", + "domain": "example.com" + } + } + } + ] + } + } + """ + + citations: List[CitationReferenceBlock] + + class ContentBlock(TypedDict, total=False): text: str image: ImageBlock @@ -103,6 +164,7 @@ class ContentBlock(TypedDict, total=False): cachePoint: CachePointBlock reasoningContent: BedrockConverseReasoningContentBlock guardContent: GuardrailConverseContentBlock + citationsContent: CitationsContentBlock class MessageBlock(TypedDict): @@ -159,8 +221,24 @@ class ToolSpecBlock(TypedDict, total=False): description: str +class SystemToolBlock(TypedDict, total=False): + """ + System tool block for Nova grounding and other built-in tools. + + Example: + { + "systemTool": { + "name": "nova_grounding" + } + } + """ + + name: Required[str] + + class ToolBlock(TypedDict, total=False): toolSpec: Optional[ToolSpecBlock] + systemTool: Optional[SystemToolBlock] cachePoint: Optional[CachePointBlock] @@ -210,11 +288,13 @@ class ContentBlockStartEvent(TypedDict, total=False): class ContentBlockDeltaEvent(TypedDict, total=False): """ Either 'text' or 'toolUse' will be specified for Converse API streaming response. + May also include 'citationsContent' when Nova grounding is enabled. """ text: str toolUse: ToolBlockDeltaEvent reasoningContent: BedrockConverseReasoningContentBlockDelta + citationsContent: CitationsContentBlock class PerformanceConfigBlock(TypedDict): @@ -879,3 +959,8 @@ class BedrockGetBatchResponse(TypedDict, total=False): outputDataConfig: BedrockOutputDataConfig timeoutDurationInHours: Optional[int] clientRequestToken: Optional[str] + +class BedrockToolBlock(TypedDict, total=False): + toolSpec: Optional[ToolSpecBlock] + systemTool: Optional[SystemToolBlock] # For Nova grounding + cachePoint: Optional[CachePointBlock] diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index cbca58e6516..96d78cf8827 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -51,19 +51,20 @@ class GenericGuardrailAPIRequest(BaseModel): """Request model for the Generic Guardrail API""" input_type: Literal["request", "response"] - litellm_call_id: Optional[str] # the call id of the individual LLM call + litellm_call_id: Optional[str] = None # the call id of the individual LLM call litellm_trace_id: Optional[ str - ] # the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation - structured_messages: Optional[List[AllMessageValues]] - images: Optional[List[str]] - tools: Optional[List[ChatCompletionToolParam]] - texts: Optional[List[str]] + ] = None # the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation + structured_messages: Optional[List[AllMessageValues]] = None + images: Optional[List[str]] = None + tools: Optional[List[ChatCompletionToolParam]] = None + texts: Optional[List[str]] = None request_data: GenericGuardrailAPIMetadata - additional_provider_specific_params: Optional[Dict[str, Any]] + additional_provider_specific_params: Optional[Dict[str, Any]] = None tool_calls: Optional[ Union[List[ChatCompletionToolCallChunk], List[ChatCompletionMessageToolCall]] - ] + ] = None + model: Optional[str] = None # the model being used for the LLM call class GenericGuardrailAPIResponse: diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 2ac5443b3bc..2d67e13e92d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3431,3 +3431,4 @@ class GenericGuardrailAPIInputs(TypedDict, total=False): structured_messages: List[ AllMessageValues ] # structured messages sent to the LLM - indicates if text is from system or user + model: Optional[str] # the model being used for the LLM call diff --git a/litellm/utils.py b/litellm/utils.py index 584ab8805a0..bb95be05b5d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8053,6 +8053,12 @@ class ProviderConfigManager: ) return OpenrouterEmbeddingConfig() + elif litellm.LlmProviders.VERCEL_AI_GATEWAY == provider: + from litellm.llms.vercel_ai_gateway.embedding.transformation import ( + VercelAIGatewayEmbeddingConfig, + ) + + return VercelAIGatewayEmbeddingConfig() elif litellm.LlmProviders.GIGACHAT == provider: return litellm.GigaChatEmbeddingConfig() elif litellm.LlmProviders.SAGEMAKER == provider: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index d958ea4503a..20462db12cc 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -10232,6 +10232,48 @@ "mode": "completion", "output_cost_per_token": 5e-07 }, + "deepseek-v3-2-251201": { + "input_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "max_input_tokens": 98304, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, + "glm-4-7-251222": { + "input_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, + "kimi-k2-thinking-251104": { + "input_cost_per_token": 0.0, + "litellm_provider": "volcengine", + "max_input_tokens": 229376, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "doubao-embedding": { "input_cost_per_token": 0.0, "litellm_provider": "volcengine", diff --git a/pyproject.toml b/pyproject.toml index 981ff97c976..1e4836ed19b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "1.81.3" +version = "1.81.4" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT" @@ -145,6 +145,7 @@ mypy = "^1.0" pytest = "^7.4.3" pytest-mock = "^3.12.0" pytest-asyncio = "^0.21.1" +pytest-retry = "^1.6.3" requests-mock = "^1.12.1" responses = "^0.25.7" respx = "^0.22.0" @@ -173,7 +174,7 @@ requires = ["poetry-core", "wheel"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "1.81.3" +version = "1.81.4" version_files = [ "pyproject.toml:^version" ] @@ -183,6 +184,8 @@ plugins = "pydantic.mypy" [tool.pytest.ini_options] asyncio_mode = "auto" +retries = 20 +retry_delay = 5 markers = [ "asyncio: mark test as an asyncio test", "limit_leaks: mark test with memory limit for leak detection (e.g., '40 MB')", diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index d700e56a06f..49288dd3775 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -6,7 +6,7 @@ def get_function_names_from_file(file_path): """ Extracts all function names from a given Python file. """ - with open(file_path, "r") as file: + with open(file_path, "r", encoding="utf-8") as file: tree = ast.parse(file.read()) function_names = [] @@ -45,7 +45,7 @@ def get_all_functions_called_in_tests(base_dir): if file.endswith(".py") and "router" in file.lower(): print("file: ", file) file_path = os.path.join(root, file) - with open(file_path, "r") as f: + with open(file_path, "r", encoding="utf-8") as f: try: tree = ast.parse(f.read()) except SyntaxError: diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 7c0db41d13a..d48dd1bfd98 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3954,3 +3954,288 @@ def test_bedrock_openai_error_handling(): assert exc_info.value.status_code == 422 print("✓ Error handling works correctly") + +# ============================================================================ +# Nova Grounding (web_search_options) Unit Tests (Mocked) +# ============================================================================ + +def test_bedrock_nova_grounding_web_search_options_non_streaming(): + """ + Unit test for Nova grounding using web_search_options parameter (non-streaming). + + This test mocks the HTTP call to verify: + 1. web_search_options is correctly mapped to systemTool for Nova models + 2. The request structure is correct + + Related: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html + """ + from unittest.mock import patch, MagicMock + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + + messages = [ + { + "role": "user", + "content": "What is the current population of Tokyo, Japan?", + } + ] + + with patch.object(client, "post") as mock_post: + try: + completion( + model="us.amazon.nova-pro-v1:0", # No bedrock/ prefix when using api_base + messages=messages, + web_search_options={}, # Enables Nova grounding + max_tokens=500, + client=client, + api_base="https://bedrock-runtime.us-east-1.amazonaws.com", + ) + except Exception: + pass # Expected - we're just checking the request structure + + # Verify the request was made correctly + if mock_post.called: + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + print(f"Request body: {json.dumps(request_body, indent=2)}") + + # Verify toolConfig is present with systemTool + assert "toolConfig" in request_body, "toolConfig should be in request" + tool_config = request_body["toolConfig"] + assert "tools" in tool_config, "tools should be in toolConfig" + + # Find the systemTool for nova_grounding + system_tool_found = False + for tool in tool_config["tools"]: + if "systemTool" in tool: + assert tool["systemTool"]["name"] == "nova_grounding" + system_tool_found = True + break + + assert system_tool_found, "systemTool with nova_grounding should be present" + print(f"✓ web_search_options correctly transformed to systemTool (non-streaming)") + + +def test_bedrock_nova_grounding_with_function_tools(): + """ + Unit test for Nova grounding combined with regular function tools. + + This tests the scenario where users want both web grounding AND + custom function calling capabilities. + """ + from unittest.mock import patch + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + + # Regular function tool + tools = [ + { + "type": "function", + "function": { + "name": "get_stock_price", + "description": "Get the current stock price for a given ticker symbol", + "parameters": { + "type": "object", + "properties": { + "ticker": { + "type": "string", + "description": "The stock ticker symbol, e.g. AAPL, GOOGL", + } + }, + "required": ["ticker"], + }, + }, + } + ] + + messages = [ + { + "role": "user", + "content": "What is the current market cap of Apple Inc?", + } + ] + + with patch.object(client, "post") as mock_post: + try: + completion( + model="us.amazon.nova-pro-v1:0", # No bedrock/ prefix when using api_base + messages=messages, + tools=tools, + web_search_options={}, # Also enable web grounding + max_tokens=500, + client=client, + api_base="https://bedrock-runtime.us-east-1.amazonaws.com", + ) + except Exception: + pass # Expected - we're just checking the request structure + + # Verify the request was made correctly + if mock_post.called: + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + print(f"Request body: {json.dumps(request_body, indent=2)}") + + # Verify toolConfig has both function tool and systemTool + assert "toolConfig" in request_body, "toolConfig should be in request" + tool_config = request_body["toolConfig"] + assert "tools" in tool_config, "tools should be in toolConfig" + + tools_in_request = tool_config["tools"] + + # Should have both the function tool and the systemTool + function_tool_found = False + system_tool_found = False + + for tool in tools_in_request: + if "toolSpec" in tool: + assert tool["toolSpec"]["name"] == "get_stock_price" + function_tool_found = True + if "systemTool" in tool: + assert tool["systemTool"]["name"] == "nova_grounding" + system_tool_found = True + + assert function_tool_found, "Function tool (get_stock_price) should be present" + assert system_tool_found, "systemTool (nova_grounding) should be present" + print(f"✓ Both function tools and web_search_options correctly combined") + + +@pytest.mark.asyncio +async def test_bedrock_nova_grounding_async(): + """ + Async unit test for Nova grounding via web_search_options. + + This test verifies the request transformation for async calls. + """ + from unittest.mock import patch, AsyncMock + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + client = AsyncHTTPHandler() + + messages = [ + { + "role": "user", + "content": "What is the weather forecast for New York City today?", + } + ] + + with patch.object(client, "post", new=AsyncMock()) as mock_post: + try: + await litellm.acompletion( + model="us.amazon.nova-pro-v1:0", # No bedrock/ prefix when using api_base + messages=messages, + web_search_options={}, + max_tokens=500, + client=client, + api_base="https://bedrock-runtime.us-east-1.amazonaws.com", + ) + except Exception: + pass # Expected - we're just checking the request structure + + # Verify the request was made correctly + if mock_post.called: + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + print(f"Request body: {json.dumps(request_body, indent=2)}") + + # Verify toolConfig is present with systemTool + assert "toolConfig" in request_body, "toolConfig should be in request" + tool_config = request_body["toolConfig"] + assert "tools" in tool_config, "tools should be in toolConfig" + + # Find the systemTool for nova_grounding + system_tool_found = False + for tool in tool_config["tools"]: + if "systemTool" in tool: + assert tool["systemTool"]["name"] == "nova_grounding" + system_tool_found = True + break + + assert system_tool_found, "systemTool with nova_grounding should be present" + print(f"✓ Async web_search_options correctly transformed to systemTool") + + +def test_bedrock_nova_web_search_options_ignored_for_non_nova(): + """ + Test that web_search_options is ignored for non-Nova Bedrock models. + + Nova grounding is only supported on Nova models. For other models, + the parameter should be silently ignored. + """ + from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig + + config = AmazonConverseConfig() + + # Should return None for non-Nova models + result = config._map_web_search_options({}, "anthropic.claude-3-sonnet-v1") + assert result is None + + result = config._map_web_search_options({}, "amazon.titan-text-express-v1") + assert result is None + + # Should return systemTool for Nova models + result = config._map_web_search_options({}, "amazon.nova-pro-v1:0") + assert result is not None + system_tool = result.get("systemTool") + assert system_tool is not None + assert system_tool["name"] == "nova_grounding" + + result2 = config._map_web_search_options({}, "us.amazon.nova-premier-v1:0") + assert result2 is not None + system_tool2 = result2.get("systemTool") + assert system_tool2 is not None + assert system_tool2["name"] == "nova_grounding" + + +def test_bedrock_nova_grounding_request_transformation(): + """ + Unit test to verify that web_search_options transforms to systemTool in the request. + """ + from unittest.mock import patch, MagicMock + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + + messages = [{"role": "user", "content": "What is the population of Tokyo?"}] + + with patch.object(client, "post") as mock_post: + mock_post.return_value = MagicMock( + status_code=200, + json=lambda: { + "output": {"message": {"role": "assistant", "content": [{"text": "Test"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 5} + } + ) + + try: + response = completion( + model="bedrock/us.amazon.nova-pro-v1:0", + messages=messages, + web_search_options={}, + max_tokens=100, + client=client, + ) + except Exception: + pass # Expected - we're just checking the request + + if mock_post.called: + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + print(f"Request body: {json.dumps(request_body, indent=2)}") + + # Verify toolConfig is present with systemTool + assert "toolConfig" in request_body, "toolConfig should be in request" + + tool_config = request_body["toolConfig"] + assert "tools" in tool_config, "tools should be in toolConfig" + + tools_in_request = tool_config["tools"] + + # Find the systemTool + system_tool_found = False + for tool in tools_in_request: + if "systemTool" in tool: + assert tool["systemTool"]["name"] == "nova_grounding" + system_tool_found = True + break + + assert system_tool_found, "systemTool with nova_grounding should be present" + print("✓ web_search_options correctly transformed to systemTool") diff --git a/tests/mcp_tests/test_mcp_client_unit.py b/tests/mcp_tests/test_mcp_client_unit.py index fcee3208be1..9533e56bcc8 100644 --- a/tests/mcp_tests/test_mcp_client_unit.py +++ b/tests/mcp_tests/test_mcp_client_unit.py @@ -5,7 +5,7 @@ import base64 import os import sys import pytest -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch, ANY # Add the project root to the path sys.path.insert(0, os.path.abspath("../../..")) @@ -183,7 +183,7 @@ class TestMCPClientUnitTests: assert result == mock_result mock_session_instance.initialize.assert_called_once() mock_session_instance.call_tool.assert_called_once_with( - name="test_tool", arguments={"arg1": "value1"} + name="test_tool", arguments={"arg1": "value1"},progress_callback=ANY ) diff --git a/tests/proxy_admin_ui_tests/test_key_management.py b/tests/proxy_admin_ui_tests/test_key_management.py index a196080eada..fa39a05a27d 100644 --- a/tests/proxy_admin_ui_tests/test_key_management.py +++ b/tests/proxy_admin_ui_tests/test_key_management.py @@ -473,18 +473,41 @@ async def test_get_users_key_count(prisma_client): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") await litellm.proxy.proxy_server.prisma_client.connect() - # Get initial user list and select the first user - initial_users = await get_users(role=None, page=1, page_size=20) + # Create a test user with no initial keys to ensure deterministic behavior + test_user_id = f"test_user_key_count-{uuid.uuid4()}" + test_user_request = NewUserRequest( + user_id=test_user_id, + user_role=LitellmUserRoles.INTERNAL_USER.value, + auto_create_key=False, # Ensure we start with 0 keys + ) + + await new_user( + test_user_request, + UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="admin", + ), + ) + + # Get initial key count for the test user + initial_users = await get_users( + user_ids=test_user_id, + role=None, + page=1, + page_size=20, + ) print("initial_users", initial_users) - assert len(initial_users["users"]) > 0, "No users found to test with" - + assert len(initial_users["users"]) == 1, "Test user should be found" test_user = initial_users["users"][0] + assert test_user.user_id == test_user_id initial_key_count = test_user.key_count + assert initial_key_count == 0, f"Expected initial key count to be 0, but got {initial_key_count}" - # Create a new key for the selected user + # Create a new key for the test user new_key = await generate_key_fn( data=GenerateKeyRequest( - user_id=test_user.user_id, + user_id=test_user_id, key_alias=f"test_key_{uuid.uuid4()}", models=["fake-model"], ), @@ -496,19 +519,26 @@ async def test_get_users_key_count(prisma_client): ) # Get updated user list and check key count - updated_users = await get_users(role=None, page=1, page_size=20) + updated_users = await get_users( + user_ids=test_user_id, + role=None, + page=1, + page_size=20, + ) print("updated_users", updated_users) - updated_key_count = None - for user in updated_users["users"]: - if user.user_id == test_user.user_id: - updated_key_count = user.key_count - break + assert len(updated_users["users"]) == 1, "Test user should still be found" + updated_user = updated_users["users"][0] + updated_key_count = updated_user.key_count - assert updated_key_count is not None, "Test user not found in updated users list" assert ( updated_key_count == initial_key_count + 1 ), f"Expected key count to increase by 1, but got {updated_key_count} (was {initial_key_count})" + # Clean up test user and keys + await prisma_client.db.litellm_usertable.delete( + where={"user_id": test_user_id} + ) + async def cleanup_existing_teams(prisma_client): all_teams = await prisma_client.db.litellm_teamtable.find_many() diff --git a/tests/proxy_unit_tests/test_server_root_path.py b/tests/proxy_unit_tests/test_server_root_path.py new file mode 100644 index 00000000000..4b39558e15a --- /dev/null +++ b/tests/proxy_unit_tests/test_server_root_path.py @@ -0,0 +1,64 @@ +import os +from unittest import mock +from litellm.proxy import utils + + +# Test the utility function logic +def test_get_server_root_path_unset(): + """ + Test that get_server_root_path returns empty string when SERVER_ROOT_PATH is unset + """ + with mock.patch.dict(os.environ, {}, clear=True): + # We need to make sure SERVER_ROOT_PATH is not in env + if "SERVER_ROOT_PATH" in os.environ: + del os.environ["SERVER_ROOT_PATH"] + + root_path = utils.get_server_root_path() + assert ( + root_path == "" + ), "Should return empty string when unset to allow X-Forwarded-Prefix" + + +def test_get_server_root_path_set(): + """ + Test that get_server_root_path returns the value when SERVER_ROOT_PATH is set + """ + with mock.patch.dict(os.environ, {"SERVER_ROOT_PATH": "/my-path"}, clear=True): + root_path = utils.get_server_root_path() + assert root_path == "/my-path", "Should return the set value" + + +def test_get_server_root_path_empty_string(): + """ + Test that get_server_root_path returns empty string when SERVER_ROOT_PATH is explicitly empty + """ + with mock.patch.dict(os.environ, {"SERVER_ROOT_PATH": ""}, clear=True): + root_path = utils.get_server_root_path() + assert ( + root_path == "" + ), "Should return empty string when explicitly set to empty" + + +# Integration test simulation for FastAPI app initialization +def test_fastapi_app_initialization_mock(): + """ + Simulate how proxy_server.py initializes FastAPI app with the root_path. + We don't import proxy_server because it has global side effects/singletons. + Instead we verify the logic flow. + """ + from fastapi import FastAPI + + # CASE 1: Proxy Mode (Unset) + with mock.patch.dict(os.environ, {}, clear=True): + if "SERVER_ROOT_PATH" in os.environ: + del os.environ["SERVER_ROOT_PATH"] + + server_root_path = utils.get_server_root_path() + app = FastAPI(root_path=server_root_path) + assert app.root_path == "" + + # CASE 2: Direct Mode (Set) + with mock.patch.dict(os.environ, {"SERVER_ROOT_PATH": "/custom-root"}, clear=True): + server_root_path = utils.get_server_root_path() + app = FastAPI(root_path=server_root_path) + assert app.root_path == "/custom-root" diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index 073433cb9e5..673c606a7d1 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -2169,3 +2169,32 @@ def test_resolve_model_name_from_model_id(): result = router.resolve_model_name_from_model_id("gpt-3.5-turbo") assert result == "gpt-3.5-turbo" + + +def test_get_valid_args(): + """Test get_valid_args static method returns valid Router.__init__ arguments""" + # Call the static method + valid_args = Router.get_valid_args() + + # Verify it returns a list + assert isinstance(valid_args, list) + assert len(valid_args) > 0 + + # Verify it contains expected Router.__init__ arguments + expected_args = [ + "model_list", + "routing_strategy", + "cache_responses", + "num_retries", + "timeout", + "fallbacks", + ] + for arg in expected_args: + assert arg in valid_args, f"Expected argument '{arg}' not found in valid_args" + + # Verify "self" is not in the list (since it's removed) + assert "self" not in valid_args + + # Verify it contains keyword-only arguments too + # These are common Router.__init__ parameters + assert "assistants_config" in valid_args or "search_tools" in valid_args diff --git a/tests/test_litellm/containers/test_container_integration.py b/tests/test_litellm/containers/test_container_integration.py index e83ae921c19..d36918c63b9 100644 --- a/tests/test_litellm/containers/test_container_integration.py +++ b/tests/test_litellm/containers/test_container_integration.py @@ -357,15 +357,15 @@ class TestContainerIntegration: def test_error_handling_integration(self): """Test error handling in the integration flow.""" - with patch('litellm.containers.main.base_llm_http_handler') as mock_handler: - # Simulate an API error - mock_handler.container_create_handler.side_effect = litellm.APIError( - status_code=400, - message="API Error occurred", - llm_provider="openai", - model="" - ) - + # Simulate an API error + api_error = litellm.APIError( + status_code=400, + message="API Error occurred", + llm_provider="openai", + model="" + ) + + with patch.object(litellm.main.base_llm_http_handler, 'container_create_handler', side_effect=api_error): with pytest.raises(litellm.APIError): create_container( name="Error Test Container", @@ -385,12 +385,12 @@ class TestContainerIntegration: name="Provider Test Container" ) - with patch('litellm.containers.main.base_llm_http_handler') as mock_handler: - mock_handler.container_create_handler.return_value = mock_response - + with patch.object(litellm.main.base_llm_http_handler, 'container_create_handler', return_value=mock_response) as mock_handler: response = create_container( name="Provider Test Container", custom_llm_provider=provider ) assert response.name == "Provider Test Container" + # Verify the mock was actually called (not making real API calls) + mock_handler.assert_called_once() diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py index 60b8582089b..8db6d98f13d 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py @@ -2,7 +2,9 @@ import os import sys import unittest.mock as mock +import httpx import pytest +import respx from httpx import Response sys.path.insert(0, os.path.abspath("../../..")) @@ -11,6 +13,8 @@ from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ( ResendEmailLogger, ) +# Test file for Resend email integration + @pytest.fixture def mock_env_vars(): @@ -23,21 +27,27 @@ def mock_httpx_client(): with mock.patch( "litellm_enterprise.enterprise_callbacks.send_emails.resend_email.get_async_httpx_client" ) as mock_client: - # Create a mock response - mock_response = mock.AsyncMock(spec=Response) + + mock_response = mock.Mock(spec=Response) mock_response.status_code = 200 mock_response.json.return_value = {"id": "test_email_id"} + mock_response.raise_for_status.return_value = None - # Create a mock client mock_async_client = mock.AsyncMock() mock_async_client.post.return_value = mock_response - mock_client.return_value = mock_async_client + mock_client.return_value = mock_async_client yield mock_async_client @pytest.mark.asyncio +@respx.mock async def test_send_email_success(mock_env_vars, mock_httpx_client): + # Block all HTTP requests at network level to prevent real API calls + respx.post("https://api.resend.com/emails").mock( + return_value=httpx.Response(200, json={"id": "test_email_id"}) + ) + # Initialize the logger logger = ResendEmailLogger() @@ -71,7 +81,13 @@ async def test_send_email_success(mock_env_vars, mock_httpx_client): @pytest.mark.asyncio +@respx.mock async def test_send_email_missing_api_key(mock_httpx_client): + # Block all HTTP requests at network level to prevent real API calls + respx.post("https://api.resend.com/emails").mock( + return_value=httpx.Response(200, json={"id": "test_email_id"}) + ) + # Remove the API key from environment before initializing logger original_key = os.environ.pop("RESEND_API_KEY", None) @@ -86,7 +102,9 @@ async def test_send_email_missing_api_key(mock_httpx_client): html_body = "

Test email body

" # Mock the response to avoid making real HTTP requests - mock_response = mock.AsyncMock(spec=Response) + mock_response = mock.Mock(spec=Response) + mock_response.raise_for_status.return_value = None + mock_response.status_code = 200 mock_response.json.return_value = {"id": "test_email_id"} mock_httpx_client.post.return_value = mock_response @@ -107,7 +125,13 @@ async def test_send_email_missing_api_key(mock_httpx_client): @pytest.mark.asyncio +@respx.mock async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client): + # Block all HTTP requests at network level to prevent real API calls + respx.post("https://api.resend.com/emails").mock( + return_value=httpx.Response(200, json={"id": "test_email_id"}) + ) + # Initialize the logger logger = ResendEmailLogger() @@ -118,7 +142,9 @@ async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client): html_body = "

Test email body

" # Mock the response to avoid making real HTTP requests - mock_response = mock.AsyncMock(spec=Response) + mock_response = mock.Mock(spec=Response) + mock_response.raise_for_status.return_value = None + mock_response.status_code = 200 mock_response.json.return_value = {"id": "test_email_id"} mock_httpx_client.post.return_value = mock_response diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py index 2b0bb31751c..836b717bd6e 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_sendgrid_email.py @@ -2,7 +2,9 @@ import os import sys import unittest.mock as mock +import httpx import pytest +import respx from httpx import Response sys.path.insert(0, os.path.abspath("../../..")) @@ -14,8 +16,26 @@ from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import ( @pytest.fixture def mock_env_vars(): - with mock.patch.dict(os.environ, {"SENDGRID_API_KEY": "test_api_key"}): + # Store original values + original_api_key = os.environ.get("SENDGRID_API_KEY") + original_sender_email = os.environ.get("SENDGRID_SENDER_EMAIL") + + # Set test API key and remove SENDGRID_SENDER_EMAIL to ensure isolation + os.environ["SENDGRID_API_KEY"] = "test_api_key" + if "SENDGRID_SENDER_EMAIL" in os.environ: + del os.environ["SENDGRID_SENDER_EMAIL"] + + try: yield + finally: + # Restore original values + if original_api_key is not None: + os.environ["SENDGRID_API_KEY"] = original_api_key + elif "SENDGRID_API_KEY" in os.environ: + del os.environ["SENDGRID_API_KEY"] + + if original_sender_email is not None: + os.environ["SENDGRID_SENDER_EMAIL"] = original_sender_email @pytest.fixture @@ -23,14 +43,16 @@ def mock_httpx_client(): with mock.patch( "litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email.get_async_httpx_client" ) as mock_client: - mock_response = mock.AsyncMock(spec=Response) + + mock_response = mock.Mock(spec=Response) mock_response.status_code = 202 mock_response.text = "accepted" + mock_response.raise_for_status.return_value = None mock_async_client = mock.AsyncMock() mock_async_client.post.return_value = mock_response - mock_client.return_value = mock_async_client + mock_client.return_value = mock_async_client yield mock_async_client @@ -62,18 +84,12 @@ async def test_send_email_success(mock_env_vars, mock_httpx_client): @pytest.mark.asyncio -async def test_send_email_missing_api_key(mock_httpx_client): - # Remove the API key from environment before initializing logger +async def test_send_email_missing_api_key(): original_key = os.environ.pop("SENDGRID_API_KEY", None) - + try: logger = SendGridEmailLogger() - # Mock the response to avoid making real HTTP requests - mock_response = mock.AsyncMock(spec=Response) - mock_response.status_code = 401 - mock_httpx_client.post.return_value = mock_response - with pytest.raises(ValueError): await logger.send_email( from_email="test@example.com", @@ -81,16 +97,19 @@ async def test_send_email_missing_api_key(mock_httpx_client): subject="Test Subject", html_body="

Test email body

", ) - - mock_httpx_client.post.assert_not_called() finally: - # Restore the original key if it existed if original_key is not None: os.environ["SENDGRID_API_KEY"] = original_key @pytest.mark.asyncio +@respx.mock async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client): + # Block all HTTP requests at network level to prevent real API calls + respx.post("https://api.sendgrid.com/v3/mail/send").mock( + return_value=httpx.Response(202, text="accepted") + ) + logger = SendGridEmailLogger() from_email = "test@example.com" @@ -98,10 +117,10 @@ async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client): subject = "Test Subject" html_body = "

Test email body

" - # Mock the response to avoid making real HTTP requests - mock_response = mock.AsyncMock(spec=Response) + mock_response = mock.Mock(spec=Response) mock_response.status_code = 202 mock_response.text = "accepted" + mock_response.raise_for_status.return_value = None mock_httpx_client.post.return_value = mock_response await logger.send_email( diff --git a/tests/test_litellm/integrations/datadog/test_datadog_cost_management.py b/tests/test_litellm/integrations/datadog/test_datadog_cost_management.py new file mode 100644 index 00000000000..be2084969a5 --- /dev/null +++ b/tests/test_litellm/integrations/datadog/test_datadog_cost_management.py @@ -0,0 +1,169 @@ +import os +import time +from unittest.mock import AsyncMock + +import pytest +from httpx import Response + +from litellm.integrations.datadog.datadog_cost_management import ( + DatadogCostManagementLogger, +) +from litellm.types.utils import StandardLoggingPayload + + +@pytest.fixture +def clean_env(): + # Save original env + original_api_key = os.environ.get("DD_API_KEY") + original_app_key = os.environ.get("DD_APP_KEY") + original_site = os.environ.get("DD_SITE") + + # Set test env + os.environ["DD_API_KEY"] = "test_api_key" + os.environ["DD_APP_KEY"] = "test_app_key" + os.environ["DD_SITE"] = "test.datadoghq.com" + + yield + + # Restore original env + if original_api_key: + os.environ["DD_API_KEY"] = original_api_key + else: + del os.environ["DD_API_KEY"] + + if original_app_key: + os.environ["DD_APP_KEY"] = original_app_key + else: + del os.environ["DD_APP_KEY"] + + if original_site: + os.environ["DD_SITE"] = original_site + else: + del os.environ["DD_SITE"] + + +@pytest.mark.asyncio +async def test_init(clean_env): + """ + Test initialization sets up clients and url correctly + """ + logger = DatadogCostManagementLogger() + assert logger.dd_api_key == "test_api_key" + assert logger.dd_app_key == "test_app_key" + assert ( + logger.upload_url == "https://api.test.datadoghq.com/api/v2/cost/custom_costs" + ) + + +@pytest.mark.asyncio +async def test_aggregate_costs(clean_env): + """ + Test that costs are correctly aggregated by provider, model, and date + """ + logger = DatadogCostManagementLogger() + + # Mock some log payloads + now = time.time() + day_str = time.strftime("%Y-%m-%d", time.localtime(now)) + + logs = [ + StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4", + response_cost=0.01, + startTime=now, + metadata={"user_api_key_team_alias": "team-a"}, + ), + StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4", + response_cost=0.02, + startTime=now, + metadata={"user_api_key_team_alias": "team-a"}, + ), + StandardLoggingPayload( + custom_llm_provider="anthropic", + model="claude-3", + response_cost=0.05, + startTime=now, + ), + ] + + aggregated = logger._aggregate_costs(logs) + + assert len(aggregated) == 2 + + # Check OpenAI entry + openai_entry = next(e for e in aggregated if e["ProviderName"] == "openai") + assert openai_entry["BilledCost"] == 0.03 + assert openai_entry["ChargeDescription"] == "LLM Usage for gpt-4" + assert openai_entry["ChargePeriodStart"] == day_str + assert openai_entry["Tags"]["team"] == "team-a" + assert "env" in openai_entry["Tags"] + assert "service" in openai_entry["Tags"] + + # Check Anthropic entry + anthropic_entry = next(e for e in aggregated if e["ProviderName"] == "anthropic") + assert anthropic_entry["BilledCost"] == 0.05 + + +@pytest.mark.asyncio +async def test_async_log_success_event(clean_env): + """ + Test that logs are added to queue + """ + logger = DatadogCostManagementLogger(batch_size=10) + + await logger.async_log_success_event( + kwargs={"standard_logging_object": {"response_cost": 0.01}}, + response_obj={}, + start_time=time.time(), + end_time=time.time(), + ) + + assert len(logger.log_queue) == 1 + assert logger.log_queue[0]["response_cost"] == 0.01 + + # Test zero cost ignored + await logger.async_log_success_event( + kwargs={"standard_logging_object": {"response_cost": 0.0}}, + response_obj={}, + start_time=time.time(), + end_time=time.time(), + ) + + assert len(logger.log_queue) == 1 + + +@pytest.mark.asyncio +async def test_async_send_batch(clean_env): + """ + Test that batch is aggregated and uploaded + """ + logger = DatadogCostManagementLogger() + logger.async_client = AsyncMock() + logger.async_client.put.return_value = Response(202, json={"status": "ok"}) + + # Add logs directly to queue + logger.log_queue = [ + StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4", + response_cost=0.01, + startTime=time.time(), + ) + ] + + await logger.async_send_batch() + + # Verify API called + assert logger.async_client.put.called + call_args = logger.async_client.put.call_args + assert call_args[0][0] == "https://api.test.datadoghq.com/api/v2/cost/custom_costs" + + import json + + # Use call_args.kwargs['content'] + content = json.loads(call_args[1]["content"]) + assert content[0]["ProviderName"] == "openai" + assert content[0]["BilledCost"] == 0.01 diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py b/tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py new file mode 100644 index 00000000000..2bb51e1e1b7 --- /dev/null +++ b/tests/test_litellm/integrations/datadog/test_datadog_llm_obs_agent.py @@ -0,0 +1,62 @@ +import os +from unittest.mock import patch +from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger + + +def test_datadog_llm_obs_agent_configuration(): + """ + Test that DataDog LLM Obs logger correctly configures agent endpoint. + """ + test_env = { + "LITELLM_DD_AGENT_HOST": "localhost", + "LITELLM_DD_LLM_OBS_PORT": "10518", + "DD_API_KEY": "test-api-key", # Optional, but checking if it's preserved + } + + # Ensure DD_SITE is NOT set to verify we don't need it in agent mode + + with patch.dict(os.environ, test_env, clear=True): + with patch("asyncio.create_task"): # Prevent periodic flush task from running + dd_logger = DataDogLLMObsLogger() + + expected_url = "http://localhost:10518/api/intake/llm-obs/v1/trace/spans" + assert dd_logger.intake_url == expected_url + assert dd_logger.DD_API_KEY == "test-api-key" + + +def test_datadog_llm_obs_agent_no_api_key_ok(): + """ + Test that agent mode works WITHOUT DD_API_KEY (agent handles auth). + """ + test_env = { + "LITELLM_DD_AGENT_HOST": "localhost", + # No DD_API_KEY + } + + with patch.dict(os.environ, test_env, clear=True): + with patch("asyncio.create_task"): + # Should NOT raise exception anymore + dd_logger = DataDogLLMObsLogger() + + assert dd_logger.DD_API_KEY is None + # Default port is 8126 if not set + expected_url = "http://localhost:8126/api/intake/llm-obs/v1/trace/spans" + assert dd_logger.intake_url == expected_url + + +def test_datadog_llm_obs_direct_api_configuration(): + """ + Test that direct API configuration still works as expected. + """ + test_env = { + "DD_API_KEY": "direct-api-key", + "DD_SITE": "us5.datadoghq.com", + } + + with patch.dict(os.environ, test_env, clear=True): + with patch("asyncio.create_task"): + dd_logger = DataDogLLMObsLogger() + + expected_url = "https://api.us5.datadoghq.com/api/intake/llm-obs/v1/trace/spans" + assert dd_logger.intake_url == expected_url + assert dd_logger.DD_API_KEY == "direct-api-key" diff --git a/tests/test_litellm/integrations/test_custom_guardrail_recursion.py b/tests/test_litellm/integrations/test_custom_guardrail_recursion.py new file mode 100644 index 00000000000..f05b5848bdf --- /dev/null +++ b/tests/test_litellm/integrations/test_custom_guardrail_recursion.py @@ -0,0 +1,73 @@ +import pytest +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.types.guardrails import GuardrailEventHooks +import json + + +class TestCustomGuardrailRecursion: + """ + Specific tests for the circular reference / RecursionError fix in logging. + """ + + def test_log_guardrail_information_handles_circular_references(self): + """ + Test that add_standard_logging method sanitizes input data containing circular references + instead of crashing. + + This reproduces the Langfuse crash scenario: + Request -> Metadata -> GuardrailResponse -> DebugContext -> Request + """ + guardrail = CustomGuardrail( + guardrail_name="recursion_test_guardrail", + event_hook=GuardrailEventHooks.pre_call, + ) + + # 1. Setup Circular Data + request_data = {"user_id": "test_recursive_user"} + metadata = {"session_id": "123"} + request_data["metadata"] = metadata + + # Create the danger: Guardrail Response holding a reference back to request_data + dirty_response = { + "flagged": False, + "debug_context": request_data, # <--- ACCESS TO ROOT (Circular Ref) + } + + # 2. Invoke the logging method + # If the fix is working, this will NOT raise RecursionError + try: + guardrail.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=dirty_response, + request_data=request_data, + guardrail_status="success", + start_time=1.0, + end_time=2.0, + duration=1.0, + masked_entity_count={}, + event_type=GuardrailEventHooks.pre_call, + ) + except RecursionError: + pytest.fail( + "RecursionError raised! The cyclic reference sanitization failed." + ) + + # 3. Verify the data stored is safe + stored_info = request_data["metadata"][ + "standard_logging_guardrail_information" + ][0] + stored_response = stored_info["guardrail_response"] + + # Check that we can dump it to JSON without crashing (Ultimate proof) + try: + json.dumps(stored_response) + except Exception as e: + pytest.fail(f"Stored data is not JSON serializable: {e}") + + # Check content - keys should be preserved but recursion broken + assert "debug_context" in stored_response + debug_context = stored_response["debug_context"] + + # In a sanitized copy, the nested metadata should be a copy, not the original live dict + assert debug_context["user_id"] == "test_recursive_user" + # The 'metadata' inside 'debug_context' would be where recursion stops or is filtered + assert "metadata" in debug_context diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index a22fe13798f..e87233a52a3 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1138,6 +1138,73 @@ def test_bedrock_create_bedrock_block_different_document_formats(): assert block["document"]["name"].endswith(f"_{format_type}") assert block["document"]["format"] == format_type +def test_bedrock_nova_web_search_options_mapping(): + """ + Test that web_search_options is correctly mapped to Nova grounding. + + This follows the LiteLLM pattern for web search where: + - Vertex AI maps web_search_options to {"googleSearch": {}} + - Anthropic maps web_search_options to {"type": "web_search_20250305", ...} + - Nova should map web_search_options to {"systemTool": {"name": "nova_grounding"}} + """ + from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig + + config = AmazonConverseConfig() + + # Test basic mapping for Nova model + result = config._map_web_search_options({}, "amazon.nova-pro-v1:0") + + assert result is not None + system_tool = result.get("systemTool") + assert system_tool is not None + assert system_tool["name"] == "nova_grounding" + + # Test with search_context_size (should be ignored for Nova) + result2 = config._map_web_search_options( + {"search_context_size": "high"}, + "us.amazon.nova-premier-v1:0" + ) + + assert result2 is not None + system_tool2 = result2.get("systemTool") + assert system_tool2 is not None + assert system_tool2["name"] == "nova_grounding" + # Nova doesn't support search_context_size, so it's just ignored + +def test_bedrock_tools_pt_does_not_handle_system_tool(): + """ + Verify that _bedrock_tools_pt does NOT handle system_tool format. + + System tools (nova_grounding) should be added via web_search_options, + not via the tools parameter directly. + """ + + from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt + + # Regular function tools should still work + tools = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"} + }, + "required": ["location"] + } + } + } + ] + + result = _bedrock_tools_pt(tools=tools) + + assert len(result) == 1 + tool_spec = result[0].get("toolSpec") + assert tool_spec is not None + assert tool_spec["name"] == "get_weather" def test_convert_to_anthropic_tool_result_image_with_cache_control(): """ @@ -1305,12 +1372,12 @@ def test_convert_to_anthropic_tool_result_image_url_as_http(): assert result["content"][0]["cache_control"]["type"] == "ephemeral" def test_anthropic_messages_pt_server_tool_use_passthrough(): """ - Test that anthropic_messages_pt passes through server_tool_use and + Test that anthropic_messages_pt passes through server_tool_use and tool_search_tool_result blocks in assistant message content. - + These are Anthropic-native content types used for tool search functionality that need to be preserved when reconstructing multi-turn conversations. - + Fixes: https://github.com/BerriAI/litellm/issues/XXXXX """ from litellm.litellm_core_utils.prompt_templates.factory import anthropic_messages_pt @@ -1359,15 +1426,15 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): # Verify we have 3 messages (user, assistant, user) assert len(result) == 3 - + # Verify the assistant message content assistant_msg = result[1] assert assistant_msg["role"] == "assistant" assert isinstance(assistant_msg["content"], list) - + # Find the different content block types content_types = [block.get("type") for block in assistant_msg["content"]] - + # Verify server_tool_use block is preserved assert "server_tool_use" in content_types server_tool_use_block = next( @@ -1376,7 +1443,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): assert server_tool_use_block["id"] == "srvtoolu_01ABC123" assert server_tool_use_block["name"] == "tool_search_tool_regex" assert server_tool_use_block["input"] == {"query": ".*time.*"} - + # Verify tool_search_tool_result block is preserved assert "tool_search_tool_result" in content_types tool_result_block = next( @@ -1385,7 +1452,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): assert tool_result_block["tool_use_id"] == "srvtoolu_01ABC123" assert tool_result_block["content"]["type"] == "tool_search_tool_search_result" assert tool_result_block["content"]["tool_references"][0]["tool_name"] == "get_time" - + # Verify text block is also preserved assert "text" in content_types text_block = next( diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index e035e193fe1..1f3f558a498 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -787,11 +787,13 @@ def test_get_masked_values(): "presidio_ad_hoc_recognizers": None, "aws_bedrock_runtime_endpoint": None, "presidio_anonymizer_api_base": None, + "vertex_credentials": "{sensitive_api_key}", } masked_values = _get_masked_values( sensitive_object, unmasked_length=4, number_of_asterisks=4 ) assert masked_values["presidio_anonymizer_api_base"] is None + assert masked_values["vertex_credentials"] == "{s****y}" @pytest.mark.asyncio 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 6aadbc058d1..c26d057fbf1 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 @@ -1108,3 +1108,309 @@ def test_streaming_chunk_with_both_text_and_tool_calls_issue_18238(): assert block_type == "tool_use" assert content_block_start["name"] == "Bash" assert content_block_start["id"] == "toolu_bdrk_013xRVejhv3ybmLEGCoZib2b" + + +# ============================================================================ +# Cache Control Transformation Tests +# ============================================================================ + +# Model constant for cache control tests +CACHE_CONTROL_BEDROCK_CONVERSE_MODEL = "bedrock/converse/global.anthropic.claude-opus-4-5-20251101-v1:0" +CACHE_CONTROL_NON_ANTHROPIC_MODEL = "gpt-4" + + +def test_should_add_cache_control_for_anthropic_model(): + """Should add cache_control to target for Anthropic Claude models.""" + adapter = LiteLLMAnthropicMessagesAdapter() + cache_control = {"type": "ephemeral"} + + for model in [ + CACHE_CONTROL_BEDROCK_CONVERSE_MODEL, + "anthropic/claude-sonnet-4-5", + "claude-opus-4-5-20251101", + "vertex_ai/claude-3-sonnet@20240229", + ]: + target = {} + adapter._add_cache_control_if_applicable({"cache_control": cache_control}, target, model) + assert "cache_control" in target + assert target["cache_control"] == cache_control + + +def test_should_not_add_cache_control_for_non_anthropic_model(): + """Should not add cache_control for non-Anthropic models.""" + adapter = LiteLLMAnthropicMessagesAdapter() + cache_control = {"type": "ephemeral"} + + for model in [CACHE_CONTROL_NON_ANTHROPIC_MODEL, "openai/gpt-4-turbo", "gemini-pro"]: + target = {} + adapter._add_cache_control_if_applicable({"cache_control": cache_control}, target, model) + assert "cache_control" not in target + + +def test_should_not_add_cache_control_when_none(): + """Should not add cache_control when source has None or empty cache_control.""" + adapter = LiteLLMAnthropicMessagesAdapter() + + for source in [{"cache_control": None}, {"cache_control": {}}, {"cache_control": ""}, {}]: + target = {} + adapter._add_cache_control_if_applicable(source, target, CACHE_CONTROL_BEDROCK_CONVERSE_MODEL) + assert "cache_control" not in target + + +def test_should_not_add_cache_control_when_model_none(): + """Should not add cache_control when model is None or empty.""" + adapter = LiteLLMAnthropicMessagesAdapter() + cache_control = {"type": "ephemeral"} + + for model in [None, ""]: + target = {} + adapter._add_cache_control_if_applicable({"cache_control": cache_control}, target, model) + assert "cache_control" not in target + + +def test_cache_control_preserved_in_text_content_for_claude(): + """Cache control should be preserved in text content for Claude models.""" + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "text", + "text": "This is cached content", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + assert len(result) == 1 + assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_not_preserved_for_non_claude_model(): + """Cache control should NOT be preserved for non-Claude models.""" + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "text", + "text": "This is cached content", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_NON_ANTHROPIC_MODEL + ) + + assert len(result) == 1 + assert "cache_control" not in result[0]["content"][0] + + +def test_cache_control_preserved_in_image_content_for_claude(): + """Cache control should be preserved in image content for Claude models.""" + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==", + }, + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + assert len(result) == 1 + assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_preserved_in_document_content_for_claude(): + """Cache control should be preserved in document content for Claude models.""" + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "document", + "source": { + "type": "base64", + "media_type": "application/pdf", + "data": "JVBERi0xLjQKJeLjz9MK", + }, + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + assert len(result) == 1 + assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_preserved_in_tool_result_for_claude(): + """Cache control should be preserved in tool_result for Claude models.""" + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "tool_result", + "tool_use_id": "toolu_01234", + "content": "Tool result content", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + tool_message = next(msg for msg in result if msg.get("role") == "tool") + assert tool_message["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_not_preserved_in_tool_result_for_non_claude(): + """Cache control should NOT be preserved in tool_result for non-Claude models.""" + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "tool_result", + "tool_use_id": "toolu_01234", + "content": "Tool result content", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_NON_ANTHROPIC_MODEL + ) + + tool_message = next(msg for msg in result if msg.get("role") == "tool") + assert "cache_control" not in tool_message + + +def test_cache_control_preserved_in_assistant_text_for_claude(): + """Cache control should be preserved in assistant text blocks for Claude models.""" + anthropic_messages = [ + AnthopicMessagesAssistantMessageParam( + role="assistant", + content=[ + { + "type": "text", + "text": "Assistant response", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + assert len(result) == 1 + assert result[0]["role"] == "assistant" + # When cache_control is present, content should be a list + assert isinstance(result[0]["content"], list) + assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_preserved_in_tool_use_for_claude(): + """Cache control should be preserved in tool_use blocks for Claude models.""" + anthropic_messages = [ + AnthopicMessagesAssistantMessageParam( + role="assistant", + content=[ + { + "type": "tool_use", + "id": "toolu_01234", + "name": "get_weather", + "input": {"location": "Boston"}, + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + assert len(result) == 1 + assert "tool_calls" in result[0] + assert result[0]["tool_calls"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_preserved_in_tools_for_claude(): + """Cache control should be preserved in tools for Claude models.""" + tools = [ + { + "name": "get_weather", + "description": "Get weather for a location", + "input_schema": {"type": "object", "properties": {"location": {"type": "string"}}}, + "cache_control": {"type": "ephemeral"}, + } + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_tools_to_openai( + tools=tools, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL + ) + + assert len(result) == 1 + assert result[0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_not_preserved_in_tools_for_non_claude(): + """Cache control should NOT be preserved in tools for non-Claude models.""" + tools = [ + { + "name": "get_weather", + "description": "Get weather for a location", + "input_schema": {"type": "object", "properties": {"location": {"type": "string"}}}, + "cache_control": {"type": "ephemeral"}, + } + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_tools_to_openai( + tools=tools, model=CACHE_CONTROL_NON_ANTHROPIC_MODEL + ) + + assert len(result) == 1 + assert "cache_control" not in result[0] diff --git a/tests/test_litellm/llms/vercel_ai_gateway/embedding/__init__.py b/tests/test_litellm/llms/vercel_ai_gateway/embedding/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/vercel_ai_gateway/embedding/test_vercel_ai_gateway_embedding.py b/tests/test_litellm/llms/vercel_ai_gateway/embedding/test_vercel_ai_gateway_embedding.py new file mode 100644 index 00000000000..af1e1df92fd --- /dev/null +++ b/tests/test_litellm/llms/vercel_ai_gateway/embedding/test_vercel_ai_gateway_embedding.py @@ -0,0 +1,218 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.vercel_ai_gateway.embedding.transformation import ( + VercelAIGatewayEmbeddingConfig, +) +from litellm.llms.vercel_ai_gateway.common_utils import VercelAIGatewayException +from litellm.types.utils import EmbeddingResponse + + +def test_vercel_ai_gateway_embedding_get_complete_url(): + """Test URL generation for embeddings endpoint""" + config = VercelAIGatewayEmbeddingConfig() + + # Test with default API base + url = config.get_complete_url( + api_base=None, + api_key=None, + model="openai/text-embedding-3-small", + optional_params={}, + litellm_params={}, + ) + assert url == "https://ai-gateway.vercel.sh/v1/embeddings" + + # Test with custom API base + url = config.get_complete_url( + api_base="https://custom.vercel.sh/v1", + api_key=None, + model="openai/text-embedding-3-small", + optional_params={}, + litellm_params={}, + ) + assert url == "https://custom.vercel.sh/v1/embeddings" + + # Test with trailing slash + url = config.get_complete_url( + api_base="https://custom.vercel.sh/v1/", + api_key=None, + model="openai/text-embedding-3-small", + optional_params={}, + litellm_params={}, + ) + assert url == "https://custom.vercel.sh/v1/embeddings" + + +def test_vercel_ai_gateway_embedding_transform_request(): + """Test request transformation for embeddings""" + config = VercelAIGatewayEmbeddingConfig() + + # Test with string input + request = config.transform_embedding_request( + model="openai/text-embedding-3-small", + input="Hello world", + optional_params={}, + headers={}, + ) + assert request["model"] == "openai/text-embedding-3-small" + assert request["input"] == ["Hello world"] + + # Test with list input + request = config.transform_embedding_request( + model="openai/text-embedding-3-small", + input=["Hello", "World"], + optional_params={}, + headers={}, + ) + assert request["model"] == "openai/text-embedding-3-small" + assert request["input"] == ["Hello", "World"] + + # Test stripping vercel_ai_gateway/ prefix + request = config.transform_embedding_request( + model="vercel_ai_gateway/openai/text-embedding-3-small", + input="Hello", + optional_params={}, + headers={}, + ) + assert request["model"] == "openai/text-embedding-3-small" + + +def test_vercel_ai_gateway_embedding_transform_request_with_dimensions(): + """Test request transformation with dimensions parameter""" + config = VercelAIGatewayEmbeddingConfig() + + request = config.transform_embedding_request( + model="openai/text-embedding-3-small", + input="Hello world", + optional_params={"dimensions": 768}, + headers={}, + ) + assert request["model"] == "openai/text-embedding-3-small" + assert request["input"] == ["Hello world"] + assert request["dimensions"] == 768 + + +def test_vercel_ai_gateway_embedding_validate_environment(): + """Test header validation and setup""" + config = VercelAIGatewayEmbeddingConfig() + + headers = config.validate_environment( + headers={}, + model="openai/text-embedding-3-small", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test_key", + ) + assert headers["Content-Type"] == "application/json" + assert headers["Authorization"] == "Bearer test_key" + + # Test with existing headers (should merge) + headers = config.validate_environment( + headers={"X-Custom": "value"}, + model="openai/text-embedding-3-small", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test_key", + ) + assert headers["X-Custom"] == "value" + assert headers["Authorization"] == "Bearer test_key" + + +def test_vercel_ai_gateway_embedding_get_supported_params(): + """Test supported OpenAI parameters""" + config = VercelAIGatewayEmbeddingConfig() + supported = config.get_supported_openai_params("openai/text-embedding-3-small") + + assert "dimensions" in supported + assert "encoding_format" in supported + assert "timeout" in supported + assert "user" in supported + + +def test_vercel_ai_gateway_embedding_map_openai_params(): + """Test OpenAI parameter mapping""" + config = VercelAIGatewayEmbeddingConfig() + + optional_params = config.map_openai_params( + non_default_params={"dimensions": 768, "encoding_format": "float"}, + optional_params={}, + model="openai/text-embedding-3-small", + drop_params=False, + ) + assert optional_params["dimensions"] == 768 + assert optional_params["encoding_format"] == "float" + + +def test_vercel_ai_gateway_embedding_error_class(): + """Test error class creation""" + config = VercelAIGatewayEmbeddingConfig() + + error = config.get_error_class( + error_message="Test error", + status_code=400, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, VercelAIGatewayException) + assert error.message == "Test error" + assert error.status_code == 400 + + +def test_vercel_ai_gateway_embedding_transform_response(): + """Test response transformation""" + config = VercelAIGatewayEmbeddingConfig() + + mock_response = MagicMock(spec=httpx.Response) + mock_response.text = '{"object":"list","data":[{"object":"embedding","index":0,"embedding":[0.1,0.2,0.3]}],"model":"openai/text-embedding-3-small","usage":{"prompt_tokens":2,"total_tokens":2}}' + mock_response.json.return_value = { + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}], + "model": "openai/text-embedding-3-small", + "usage": {"prompt_tokens": 2, "total_tokens": 2}, + } + + mock_logging = MagicMock() + + response = config.transform_embedding_response( + model="openai/text-embedding-3-small", + raw_response=mock_response, + model_response=EmbeddingResponse(), + logging_obj=mock_logging, + api_key="test_key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + assert response is not None + mock_logging.post_call.assert_called_once() + + +def test_vercel_ai_gateway_embedding_env_vars(): + """Test environment variable handling""" + config = VercelAIGatewayEmbeddingConfig() + + with patch.dict( + os.environ, + { + "VERCEL_AI_GATEWAY_API_BASE": "https://env.vercel.sh/v1", + }, + ): + url = config.get_complete_url( + api_base=None, + api_key=None, + model="openai/text-embedding-3-small", + optional_params={}, + litellm_params={}, + ) + assert url == "https://env.vercel.sh/v1/embeddings" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index ecdc75ede52..f241b2aa0d5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -1885,7 +1885,7 @@ class TestMCPServerManager: # Create mock client that tracks call_tool usage mock_client = AsyncMock() - async def mock_call_tool(params): + async def mock_call_tool(params, host_progress_callback=None): # Return a mock CallToolResult result = MagicMock(spec=CallToolResult) result.content = [{"type": "text", "text": "Tool executed successfully"}] diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 807559207e6..3df0dc881e8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -15,10 +15,10 @@ import pytest import litellm from litellm.proxy._types import ( CallInfo, + Litellm_EntityType, LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, LiteLLM_UserTable, - Litellm_EntityType, LitellmUserRoles, ProxyErrorTypes, ProxyException, @@ -131,6 +131,60 @@ def test_get_key_object_from_ui_hash_key_invalid(): assert key_object is None +def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values): + """Test generating CLI JWT token with default 24-hour expiration""" + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) + + # Decrypt and verify token contents + decrypted_token = decrypt_value_helper( + token, key="ui_hash_key", exception_type="debug" + ) + assert decrypted_token is not None + token_data = json.loads(decrypted_token) + + assert token_data["user_id"] == "test_user" + assert token_data["user_role"] == LitellmUserRoles.PROXY_ADMIN.value + assert token_data["models"] == ["gpt-3.5-turbo"] + assert token_data["max_budget"] == litellm.max_ui_session_budget + + # Verify expiration time is set to 24 hours (default) + assert "expires" in token_data + expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00")) + assert expires > get_utc_datetime() + assert expires <= get_utc_datetime() + timedelta(hours=24, minutes=1) + assert expires >= get_utc_datetime() + timedelta(hours=23, minutes=59) + + +def test_get_cli_jwt_auth_token_custom_expiration( + valid_sso_user_defined_values, monkeypatch +): + """Test generating CLI JWT token with custom expiration via environment variable""" + # Set custom expiration to 48 hours + monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48") + + # Reload the constants module to pick up the new env var + import importlib + + from litellm import constants + importlib.reload(constants) + + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values) + + # Decrypt and verify token contents + decrypted_token = decrypt_value_helper( + token, key="ui_hash_key", exception_type="debug" + ) + assert decrypted_token is not None + token_data = json.loads(decrypted_token) + + # Verify expiration time is set to 48 hours + assert "expires" in token_data + expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00")) + assert expires > get_utc_datetime() + timedelta(hours=47, minutes=59) + assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1) + + + @pytest.mark.asyncio async def test_default_internal_user_params_with_get_user_object(monkeypatch): """Test that default_internal_user_params is used when creating a new user via get_user_object""" diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index e7b27908c14..a0e29e06100 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -626,3 +626,133 @@ async def test_expire_previous_ui_session_tokens_exception_handling(): # Should not raise exception despite database error await expire_previous_ui_session_tokens(user_id, mock_prisma_client) + + +@pytest.mark.asyncio +async def test_authenticate_user_admin_login_with_non_ascii_characters(): + """Test admin login with non-ASCII characters in password (issue #19559)""" + master_key = "sk-1234" + ui_username = "admin£test" + ui_password = "sk-1234£pass" + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + + with patch.dict( + os.environ, + { + "UI_USERNAME": ui_username, + "UI_PASSWORD": ui_password, + "DATABASE_URL": "postgresql://test:test@localhost/test", + }, + ): + with patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", + new_callable=AsyncMock, + ) as mock_generate_key: + mock_generate_key.return_value = { + "token": "test-token-123", + "user_id": LITELLM_PROXY_ADMIN_NAME, + } + + with patch( + "litellm.proxy.auth.login_utils.user_update", + new_callable=AsyncMock, + return_value=None, + ) as mock_user_update: + with patch( + "litellm.proxy.auth.login_utils.get_secret_bool", + return_value=False, + ): + result = await authenticate_user( + username=ui_username, + password=ui_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + + assert isinstance(result, LoginResult) + assert result.user_id == LITELLM_PROXY_ADMIN_NAME + assert result.key == "test-token-123" + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + + +def test_authenticate_user_non_ascii_direct_comparison(): + """Test that non-ASCII characters can be compared directly (unit test for fix)""" + import secrets + + # This test verifies the fix handles non-ASCII by encoding to bytes + username = "admin£test" + password = "pass£word" + + # This would fail without encoding: + # secrets.compare_digest(username, username) # TypeError! + + # But works with the fix: + result = secrets.compare_digest( + username.encode("utf-8"), username.encode("utf-8") + ) + assert result is True + + # And correctly returns False for different passwords + result = secrets.compare_digest( + password.encode("utf-8"), "different£pass".encode("utf-8") + ) + assert result is False + + +@pytest.mark.asyncio +async def test_authenticate_user_database_login_with_non_ascii_password(): + """Test database user login with non-ASCII characters in password (issue #19559)""" + master_key = "sk-1234" + user_email = "test@example.com" + password_with_special_char = "correct£password" + hashed_password = hash_token(token=password_with_special_char) + + mock_user = MagicMock() + mock_user.user_id = "test-user-123" + mock_user.user_email = user_email + mock_user.password = hashed_password + mock_user.user_role = LitellmUserRoles.INTERNAL_USER + + def mock_find_first(**kwargs): + where = kwargs.get("where", {}) + user_email_filter = where.get("user_email", {}) + if str(user_email_filter.get("equals", "")).lower() == user_email.lower(): + return mock_user + return None + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( + side_effect=mock_find_first + ) + + with patch.dict( + os.environ, + { + "DATABASE_URL": "postgresql://test:test@localhost/test", + "UI_USERNAME": "admin", + "UI_PASSWORD": "admin-password", + }, + ): + with patch( + "litellm.proxy.auth.login_utils.expire_previous_ui_session_tokens", + new_callable=AsyncMock, + return_value=None, + ): + with patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", + new_callable=AsyncMock, + ) as mock_generate_key: + mock_generate_key.return_value = {"token": "token-123"} + + result = await authenticate_user( + username=user_email, + password=password_with_special_char, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + + assert isinstance(result, LoginResult) + assert result.user_id == "test-user-123" + assert result.user_email == user_email diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 61d44e46da5..5c039141928 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -549,6 +549,62 @@ class TestAdditionalParams: ) +class TestModelParameter: + """Test model parameter handling in guardrail requests""" + + @pytest.mark.asyncio + async def test_model_passed_from_inputs( + self, generic_guardrail, mock_request_data_input + ): + """Test that model is passed to the API when provided in inputs""" + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "NONE", + "texts": ["test"], + } + mock_response.raise_for_status = MagicMock() + + with patch.object( + generic_guardrail.async_handler, "post", return_value=mock_response + ) as mock_post: + await generic_guardrail.apply_guardrail( + inputs={"texts": ["test"], "model": "gpt-4"}, + request_data=mock_request_data_input, + input_type="request", + ) + + # Verify API was called with model + call_args = mock_post.call_args + json_payload = call_args.kwargs["json"] + assert json_payload["model"] == "gpt-4" + + @pytest.mark.asyncio + async def test_model_none_when_not_provided( + self, generic_guardrail, mock_request_data_input + ): + """Test that model is None when not provided in inputs""" + mock_response = MagicMock() + mock_response.json.return_value = { + "action": "NONE", + "texts": ["test"], + } + mock_response.raise_for_status = MagicMock() + + with patch.object( + generic_guardrail.async_handler, "post", return_value=mock_response + ) as mock_post: + await generic_guardrail.apply_guardrail( + inputs={"texts": ["test"]}, # No model in inputs + request_data=mock_request_data_input, + input_type="request", + ) + + # Verify API was called with model=None + call_args = mock_post.call_args + json_payload = call_args.kwargs["json"] + assert json_payload["model"] is None + + class TestErrorHandling: """Test error handling scenarios""" diff --git a/tests/test_litellm/proxy/hooks/test_post_call_streaming_hook_integration.py b/tests/test_litellm/proxy/hooks/test_post_call_streaming_hook_integration.py new file mode 100644 index 00000000000..3bc111ef142 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_post_call_streaming_hook_integration.py @@ -0,0 +1,273 @@ +""" +Integration tests for async_post_call_streaming_hook. + +Tests verify that the streaming hook can transform streaming responses sent to clients. +""" + +import os +import sys +import pytest +from typing import Any +from unittest.mock import patch, MagicMock + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import ModelResponseStream, StreamingChoices, Delta + + +class StreamingResponseTransformerLogger(CustomLogger): + """Logger that transforms streaming responses""" + + def __init__(self, transform_content: str = None): + self.called = False + self.transform_content = transform_content + self.received_response = None + + async def async_post_call_streaming_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: str, + ) -> Any: + self.called = True + self.received_response = response + if self.transform_content is not None: + return self.transform_content + return None + + +@pytest.mark.asyncio +async def test_streaming_hook_transforms_response(): + """ + Test that async_post_call_streaming_hook can transform streaming responses. + """ + transformer = StreamingResponseTransformerLogger(transform_content="Modified streaming response") + + with patch("litellm.callbacks", [transformer]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + # Create a mock streaming response + original_response = ModelResponseStream( + id="original-stream", + choices=[ + StreamingChoices( + delta=Delta(content="Original content", role="assistant"), + index=0, + ) + ], + model="test-model", + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + # Call the hook + result = await proxy_logging.async_post_call_streaming_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Verify hook was called + assert transformer.called is True + + # Verify transformed response is returned + assert result == "Modified streaming response" + + +@pytest.mark.asyncio +async def test_streaming_hook_returns_none_keeps_original(): + """ + Test that hook returning None keeps the original response. + """ + + class NoOpLogger(CustomLogger): + def __init__(self): + self.called = False + + async def async_post_call_streaming_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: str, + ): + self.called = True + return None + + logger = NoOpLogger() + + with patch("litellm.callbacks", [logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponseStream( + id="original-stream", + choices=[ + StreamingChoices( + delta=Delta(content="Original content", role="assistant"), + index=0, + ) + ], + model="test-model", + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.async_post_call_streaming_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Should return original response object + assert result.id == "original-stream" + assert logger.called is True + + +@pytest.mark.asyncio +async def test_streaming_hook_works_with_sse_format(): + """ + Test that hook works with SSE-formatted strings (data: prefix). + This was the only supported format before the fix. + """ + transformer = StreamingResponseTransformerLogger( + transform_content="data: {\"error\": \"custom error\"}\n\n" + ) + + with patch("litellm.callbacks", [transformer]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponseStream( + id="original-stream", + choices=[ + StreamingChoices( + delta=Delta(content="Original content", role="assistant"), + index=0, + ) + ], + model="test-model", + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.async_post_call_streaming_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Verify SSE-formatted response is returned + assert result == "data: {\"error\": \"custom error\"}\n\n" + + +@pytest.mark.asyncio +async def test_streaming_hook_chains_multiple_callbacks(): + """ + Test that multiple callbacks can chain modifications. + """ + + class AppendLogger(CustomLogger): + def __init__(self, suffix: str): + self.suffix = suffix + self.called = False + + async def async_post_call_streaming_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: str, + ) -> str: + self.called = True + # Note: response here is the complete_response string, not the chunk + return f"[{self.suffix}]" + + callback1 = AppendLogger("CB1") + callback2 = AppendLogger("CB2") + + with patch("litellm.callbacks", [callback1, callback2]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponseStream( + id="original-stream", + choices=[ + StreamingChoices( + delta=Delta(content="Hello", role="assistant"), + index=0, + ) + ], + model="test-model", + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.async_post_call_streaming_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Both callbacks should have been called + assert callback1.called is True + assert callback2.called is True + + # Last callback's result should be used + assert result == "[CB2]" + + +@pytest.mark.asyncio +async def test_streaming_hook_handles_exceptions(): + """ + Test that hook exceptions are propagated. + """ + + class FailingLogger(CustomLogger): + async def async_post_call_streaming_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: str, + ): + raise RuntimeError("Streaming hook crashed!") + + logger = FailingLogger() + + with patch("litellm.callbacks", [logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponseStream( + id="original-stream", + choices=[ + StreamingChoices( + delta=Delta(content="Hello", role="assistant"), + index=0, + ) + ], + model="test-model", + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + # Exception should be propagated + with pytest.raises(RuntimeError, match="Streaming hook crashed!"): + await proxy_logging.async_post_call_streaming_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) diff --git a/tests/test_litellm/proxy/hooks/test_post_call_success_hook_integration.py b/tests/test_litellm/proxy/hooks/test_post_call_success_hook_integration.py new file mode 100644 index 00000000000..870286f5382 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_post_call_success_hook_integration.py @@ -0,0 +1,260 @@ +""" +Integration tests for async_post_call_success_hook. + +Tests verify that the success hook can transform responses sent to clients. +This mirrors the behavior of CustomGuardrail hooks and streaming iterator hooks. +""" + +import os +import sys +import pytest +from typing import Any +from unittest.mock import patch, MagicMock + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import ModelResponse, Choices, Message, Usage + + +class ResponseTransformerLogger(CustomLogger): + """Logger that transforms successful responses""" + + def __init__(self, transform_content: str = None): + self.called = False + self.transform_content = transform_content + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + self.called = True + if self.transform_content is not None: + # Create a modified response with custom content + return { + "id": "transformed-response", + "choices": [ + { + "message": {"content": self.transform_content, "role": "assistant"}, + "index": 0, + } + ], + "model": "test-model", + "custom_field": "added_by_hook", + } + return response + + +@pytest.mark.asyncio +async def test_success_hook_transforms_response(): + """ + Test that async_post_call_success_hook can transform successful responses. + """ + transformer = ResponseTransformerLogger(transform_content="Modified response") + + with patch("litellm.callbacks", [transformer]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + # Create a mock response + original_response = ModelResponse( + id="original-response", + choices=[ + Choices( + message=Message(content="Original content", role="assistant"), + index=0, + finish_reason="stop", + ) + ], + model="test-model", + usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30), + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + # Call the hook + result = await proxy_logging.post_call_success_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Verify hook was called + assert transformer.called is True + + # Verify transformed response is returned + assert result is not None + assert result["id"] == "transformed-response" + assert result["choices"][0]["message"]["content"] == "Modified response" + assert result["custom_field"] == "added_by_hook" + + +@pytest.mark.asyncio +async def test_success_hook_returns_none_keeps_original(): + """ + Test that hook returning None keeps the original response. + """ + + class NoOpLogger(CustomLogger): + def __init__(self): + self.called = False + + async def async_post_call_success_hook(self, *args, **kwargs): + self.called = True + return None + + logger = NoOpLogger() + + with patch("litellm.callbacks", [logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponse( + id="original-response", + choices=[ + Choices( + message=Message(content="Original content", role="assistant"), + index=0, + finish_reason="stop", + ) + ], + model="test-model", + usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30), + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.post_call_success_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Should return original response + assert result.id == "original-response" + assert logger.called is True + + +@pytest.mark.asyncio +async def test_success_hook_chains_multiple_callbacks(): + """ + Test that multiple callbacks can chain modifications. + """ + + class AddFieldLogger(CustomLogger): + def __init__(self, field_name: str, field_value: Any): + self.field_name = field_name + self.field_value = field_value + self.called = False + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + self.called = True + # Convert response to dict if needed + if hasattr(response, "model_dump"): + resp_dict = response.model_dump() + elif hasattr(response, "dict"): + resp_dict = response.dict() + elif isinstance(response, dict): + resp_dict = response.copy() + else: + resp_dict = {} + + resp_dict[self.field_name] = self.field_value + return resp_dict + + callback1 = AddFieldLogger("field1", "value1") + callback2 = AddFieldLogger("field2", "value2") + + with patch("litellm.callbacks", [callback1, callback2]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponse( + id="original-response", + choices=[ + Choices( + message=Message(content="Original content", role="assistant"), + index=0, + finish_reason="stop", + ) + ], + model="test-model", + usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30), + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.post_call_success_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + # Both callbacks should have been called + assert callback1.called is True + assert callback2.called is True + + # Both fields should be present (chained modifications) + assert result["field1"] == "value1" + assert result["field2"] == "value2" + + +@pytest.mark.asyncio +async def test_success_hook_handles_exceptions(): + """ + Test that hook exceptions are propagated (not silently swallowed). + """ + + class FailingLogger(CustomLogger): + async def async_post_call_success_hook(self, *args, **kwargs): + raise RuntimeError("Hook crashed!") + + logger = FailingLogger() + + with patch("litellm.callbacks", [logger]): + from litellm.proxy.utils import ProxyLogging + from litellm.caching.caching import DualCache + + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + original_response = ModelResponse( + id="original-response", + choices=[ + Choices( + message=Message(content="Original content", role="assistant"), + index=0, + finish_reason="stop", + ) + ], + model="test-model", + usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30), + ) + + data = {"model": "test-model"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + # Exception should be propagated + with pytest.raises(RuntimeError, match="Hook crashed!"): + await proxy_logging.post_call_success_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 397a6af556f..dc436bac087 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1096,3 +1096,57 @@ async def test_get_users_user_id_partial_match(mocker): assert "user_id" in captured_where_conditions assert "in" in captured_where_conditions["user_id"] assert captured_where_conditions["user_id"]["in"] == ["user1", "user2", "user3"] + + +def test_update_internal_user_params_reset_max_budget_with_none(): + """ + Test that _update_internal_user_params allows setting max_budget to None. + This verifies the fix for unsetting/resetting the budget to unlimited. + """ + + # Case 1: max_budget is explicitly None in the input dictionary + data_json = {"max_budget": None, "user_id": "test_user"} + data = UpdateUserRequest(max_budget=None, user_id="test_user") + + # Call the function + non_default_values = _update_internal_user_params(data_json=data_json, data=data) + + # Assertions + assert "max_budget" in non_default_values + assert non_default_values["max_budget"] is None + assert non_default_values["user_id"] == "test_user" + + +def test_update_internal_user_params_ignores_other_nones(): + """ + Test that other fields are still filtered out if None + """ + # Create test data with other None fields + data_json = {"user_alias": None, "user_id": "test_user", "max_budget": 100.0} + data = UpdateUserRequest(user_alias=None, user_id="test_user", max_budget=100.0) + + # Call the function + non_default_values = _update_internal_user_params(data_json=data_json, data=data) + + # Assertions + assert "user_alias" not in non_default_values + assert non_default_values["max_budget"] == 100.0 + + +def test_generate_request_base_validator(): + """ + Test that GenerateRequestBase validator converts empty string to None for max_budget + """ + from litellm.proxy._types import GenerateRequestBase + + # Test with empty string + req = GenerateRequestBase(max_budget="") + assert req.max_budget is None + + # Test with actual float + req = GenerateRequestBase(max_budget=100.0) + assert req.max_budget == 100.0 + + # Test with None + req = GenerateRequestBase(max_budget=None) + assert req.max_budget is None \ No newline at end of file diff --git a/tests/test_litellm/proxy/test_chat_completion_metadata.py b/tests/test_litellm/proxy/test_chat_completion_metadata.py new file mode 100644 index 00000000000..38dcdc13c50 --- /dev/null +++ b/tests/test_litellm/proxy/test_chat_completion_metadata.py @@ -0,0 +1,154 @@ +import pytest +from unittest.mock import MagicMock, AsyncMock, patch +from litellm.proxy.proxy_server import chat_completion, completion, embeddings +from litellm.proxy._types import UserAPIKeyAuth +from fastapi import Request, Response + + +@pytest.mark.asyncio +async def test_chat_completion_metadata_population(): + # Setup + request = MagicMock(spec=Request) + # Mock _read_request_body to return a dict + with patch( + "litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock + ) as mock_read_body: + mock_read_body.return_value = {"model": "gpt-3.5-turbo", "messages": []} + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_id", team_id="test_team_id", org_id="test_org_id" + ) + + fastapi_response = MagicMock(spec=Response) + + # Mock ProxyBaseLLMRequestProcessing + with patch( + "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + mock_instance = MockProcessor.return_value + mock_instance.base_process_llm_request = AsyncMock( + return_value={"choices": []} + ) + + # Execute + await chat_completion( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + ) + + # Verify + # Check if ProxyBaseLLMRequestProcessing was initialized with data containing metadata + call_args = MockProcessor.call_args + assert call_args is not None + data_arg = call_args.kwargs.get("data") + assert data_arg is not None + + assert "metadata" in data_arg + assert data_arg["metadata"]["user_api_key_user_id"] == "test_user_id" + assert data_arg["metadata"]["user_api_key_team_id"] == "test_team_id" + assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id" + + +@pytest.mark.asyncio +async def test_embedding_metadata_population(): + """ + Test that the embedding endpoint correctly populates metadata + from UserAPIKeyAuth. + """ + # Setup + with patch( + "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request" + ): + with patch( + "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.__init__", + return_value=None, + ) as mock_base_process_init: + # Create a mock UserAPIKeyAuth object + mock_user_auth = MagicMock(spec=UserAPIKeyAuth) + mock_user_auth.user_id = "test_user_id_emb" + mock_user_auth.team_id = "test_team_id_emb" + mock_user_auth.org_id = "test_org_id_emb" + + # Create a mock Request object + mock_request = MagicMock(spec=Request) + mock_request.json = AsyncMock( + return_value={"model": "gpt-3.5-turbo", "input": "hello"} + ) + # Mock _read_request_body to return our data + with patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock( + return_value={"model": "gpt-3.5-turbo", "input": "hello"} + ), + ): + # Call the endpoint function directly + await embeddings( + request=mock_request, + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=mock_user_auth, + ) + + # Check if ProxyBaseLLMRequestProcessing was initialized with the correct metadata + mock_base_process_init.assert_called_once() + call_args = mock_base_process_init.call_args + # handle both positional and keyword args for data + if "data" in call_args.kwargs: + data_arg = call_args.kwargs["data"] + else: + data_arg = call_args.args[0] + + assert ( + data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb" + ) + assert ( + data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb" + ) + assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_emb" + + +@pytest.mark.asyncio +async def test_completion_metadata_population(): + # Setup + request = MagicMock(spec=Request) + # Mock _read_request_body to return a dict + with patch( + "litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock + ) as mock_read_body: + mock_read_body.return_value = { + "model": "gpt-3.5-turbo-instruct", + "prompt": "test", + } + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_id_2", team_id="test_team_id_2", org_id="test_org_id_2" + ) + + fastapi_response = MagicMock(spec=Response) + + # Mock ProxyBaseLLMRequestProcessing + with patch( + "litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + mock_instance = MockProcessor.return_value + mock_instance.base_process_llm_request = AsyncMock( + return_value={"choices": []} + ) + + # Execute + await completion( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + ) + + # Verify + call_args = MockProcessor.call_args + assert call_args is not None + data_arg = call_args.kwargs.get("data") + assert data_arg is not None + + assert "metadata" in data_arg + assert data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_2" + assert data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_2" + assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_2" diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index ad4f53dac4b..60bdb7d12cb 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -1227,3 +1227,113 @@ class TestProxySettingEndpoints: assert retrieved_role_mappings["provider"] == "google" assert retrieved_role_mappings["group_claim"] == "groups" assert retrieved_role_mappings["default_role"] == LitellmUserRoles.INTERNAL_USER + + def test_setup_role_mappings_custom_logic_with_env_vars(self, monkeypatch): + """Test the _setup_role_mappings function directly with custom role mapping logic from environment variables""" + import asyncio + import os + from litellm.proxy.management_endpoints.ui_sso import _setup_role_mappings + from litellm.proxy._types import LitellmUserRoles + + # Set up environment variables for custom role mappings using valid Python dict format + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['custom-admin-group'], 'internal_user': ['custom-user-group'], 'proxy_admin_viewer': ['custom-viewer-group']}") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "custom-groups") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", "internal_user_viewer") + + # Debug: Print environment variables + print("GENERIC_ROLE_MAPPINGS_ROLES:", os.getenv("GENERIC_ROLE_MAPPINGS_ROLES")) + print("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM:", os.getenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM")) + print("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE:", os.getenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE")) + + # Run the async function + role_mappings = asyncio.run(_setup_role_mappings()) + + # Debug: Print result + print("role_mappings result:", role_mappings) + + # Verify role_mappings is returned correctly from environment variables + assert role_mappings is not None + assert role_mappings.provider == "generic" + assert role_mappings.group_claim == "custom-groups" + assert role_mappings.default_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY + assert role_mappings.roles[LitellmUserRoles.PROXY_ADMIN] == ["custom-admin-group"] + assert role_mappings.roles[LitellmUserRoles.INTERNAL_USER] == ["custom-user-group"] + assert role_mappings.roles[LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY] == ["custom-viewer-group"] + + def test_setup_role_mappings_custom_logic_with_no_config(self, monkeypatch): + """Test the _setup_role_mappings function returns None when no configuration is available""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy.management_endpoints.ui_sso import _setup_role_mappings + + # Ensure environment variables are not set + monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_ROLES", raising=False) + monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", raising=False) + monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", raising=False) + + # Mock the prisma client to return None (no database record) + mock_prisma = MagicMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + # Run the async function + role_mappings = asyncio.run(_setup_role_mappings()) + + # Should return None when no configuration is available + assert role_mappings is None + + def test_get_sso_settings_with_env_role_mappings(self, mock_proxy_config, mock_auth, monkeypatch): + import json + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", '{"proxy_admin": ["custom-admin-group"], "internal_user": ["custom-user-group"], "proxy_admin_viewer": ["custom-viewer-group"]}') + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "custom-groups") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", "internal_user_viewer") + + mock_prisma = MagicMock() + mock_db_record = MagicMock() + mock_db_record.sso_settings = { + "google_client_id": "test_google_client_id", + "role_mappings": { + "provider": "google", + "group_claim": "db-groups", + "default_role": "proxy_admin", + "roles": { + "proxy_admin": ["db-admin-group"], + }, + }, + } + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_db_record) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr( + proxy_config, "_decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables + ) + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + data = response.json() + + values = data["values"] + assert "role_mappings" in values + assert values["role_mappings"] is not None + + # The database values shoeld override the environment variables + assert values["role_mappings"]["provider"] == "google" + assert values["role_mappings"]["group_claim"] == "db-groups" + assert values["role_mappings"]["default_role"] == LitellmUserRoles.PROXY_ADMIN + assert values["role_mappings"]["roles"][LitellmUserRoles.PROXY_ADMIN] == ["db-admin-group"] + + # Verify that the database was checked but environment variables took priority + mock_prisma.db.litellm_ssoconfig.find_unique.assert_called_once_with( + where={"id": "sso_config"} + ) + + # Verify other SSO settings are still correctly returned + assert values["google_client_id"] == "test_google_client_id" + + # Verify field_schema is still present + assert "field_schema" in data + assert "properties" in data["field_schema"] + assert "role_mappings" in data["field_schema"]["properties"] diff --git a/tests/test_litellm/secret_managers/test_secret_managers_main.py b/tests/test_litellm/secret_managers/test_secret_managers_main.py index 2e5270b5d70..eaef6956cd5 100644 --- a/tests/test_litellm/secret_managers/test_secret_managers_main.py +++ b/tests/test_litellm/secret_managers/test_secret_managers_main.py @@ -47,11 +47,12 @@ def mock_env(): @patch("litellm.secret_managers.main.oidc_cache") -@patch("litellm.secret_managers.main.HTTPHandler") -def test_oidc_google_success(mock_http_handler, mock_oidc_cache): +@patch("litellm.secret_managers.main._get_oidc_http_handler") +@patch("httpx.Client") # Prevent any real HTTP connections +def test_oidc_google_success(mock_httpx_client, mock_get_http_handler, mock_oidc_cache): mock_oidc_cache.get_cache.return_value = None mock_handler = MockHTTPHandler(timeout=600.0) - mock_http_handler.return_value = mock_handler + mock_get_http_handler.return_value = mock_handler secret_name = "oidc/google/[invalid url, do not cite]" result = get_secret(secret_name) @@ -63,29 +64,31 @@ def test_oidc_google_success(mock_http_handler, mock_oidc_cache): @patch("litellm.secret_managers.main.oidc_cache") -def test_oidc_google_cached(mock_oidc_cache): +@patch("litellm.secret_managers.main._get_oidc_http_handler") +def test_oidc_google_cached(mock_get_http_handler, mock_oidc_cache): mock_oidc_cache.get_cache.return_value = "cached_token" secret_name = "oidc/google/[invalid url, do not cite]" - with patch("litellm.secret_managers.main.HTTPHandler") as mock_http: - result = get_secret(secret_name) + result = get_secret(secret_name) - assert result == "cached_token", f"Expected cached token, got {result}" - mock_oidc_cache.get_cache.assert_called_with(key=secret_name) - mock_http.assert_not_called() + assert result == "cached_token", f"Expected cached token, got {result}" + mock_oidc_cache.get_cache.assert_called_with(key=secret_name) + # Verify HTTP handler was never called since we had a cached token + mock_get_http_handler.assert_not_called() @patch("litellm.secret_managers.main.oidc_cache") -def test_oidc_google_failure(mock_oidc_cache): +@patch("litellm.secret_managers.main._get_oidc_http_handler") +def test_oidc_google_failure(mock_get_http_handler, mock_oidc_cache): mock_handler = MockHTTPHandler(timeout=600.0) mock_handler.status_code = 400 + mock_get_http_handler.return_value = mock_handler + mock_oidc_cache.get_cache.return_value = None + + secret_name = "oidc/google/https://example.com/api" - with patch("litellm.secret_managers.main.HTTPHandler", return_value=mock_handler): - mock_oidc_cache.get_cache.return_value = None - secret_name = "oidc/google/https://example.com/api" - - with pytest.raises(ValueError, match="Google OIDC provider failed"): - get_secret(secret_name) + with pytest.raises(ValueError, match="Google OIDC provider failed"): + get_secret(secret_name) def test_oidc_circleci_success(monkeypatch): @@ -106,13 +109,13 @@ def test_oidc_circleci_failure(monkeypatch): @patch("litellm.secret_managers.main.oidc_cache") -@patch("litellm.secret_managers.main.HTTPHandler") -def test_oidc_github_success(mock_http_handler, mock_oidc_cache, mock_env): +@patch("litellm.secret_managers.main._get_oidc_http_handler") +def test_oidc_github_success(mock_get_http_handler, mock_oidc_cache, mock_env): mock_env["ACTIONS_ID_TOKEN_REQUEST_URL"] = "https://github.com/token" mock_env["ACTIONS_ID_TOKEN_REQUEST_TOKEN"] = "github_token" mock_oidc_cache.get_cache.return_value = None mock_handler = MockHTTPHandler(timeout=600.0) - mock_http_handler.return_value = mock_handler + mock_get_http_handler.return_value = mock_handler secret_name = "oidc/github/github-audience" result = get_secret(secret_name) @@ -142,7 +145,7 @@ def test_oidc_azure_file_success(mock_env, tmp_path): mock_env["AZURE_FEDERATED_TOKEN_FILE"] = str(token_file) secret_name = "oidc/azure/azure-audience" - result = get_secret(secret_name) + result = get_secret(secret_name) assert result == "azure_token" @@ -154,16 +157,22 @@ def test_oidc_azure_ad_token_success(mock_get_azure_ad_token_provider): if "AZURE_FEDERATED_TOKEN_FILE" in os.environ: del os.environ["AZURE_FEDERATED_TOKEN_FILE"] + # Mock the token provider function that gets returned and called mock_token_provider = Mock(return_value="azure_ad_token") mock_get_azure_ad_token_provider.return_value = mock_token_provider - secret_name = "oidc/azure/api://azure-audience" - result = get_secret(secret_name) + + # Also mock the Azure Identity SDK to prevent any real Azure calls + with patch("azure.identity.get_bearer_token_provider") as mock_bearer: + mock_bearer.return_value = mock_token_provider + + secret_name = "oidc/azure/api://azure-audience" + result = get_secret(secret_name) - assert result == "azure_ad_token" - mock_get_azure_ad_token_provider.assert_called_once_with( - azure_scope="api://azure-audience" - ) - mock_token_provider.assert_called_once_with() + assert result == "azure_ad_token" + mock_get_azure_ad_token_provider.assert_called_once_with( + azure_scope="api://azure-audience" + ) + mock_token_provider.assert_called_once_with() def test_oidc_file_success(tmp_path): diff --git a/tests/test_litellm/test_gpt_image_cost_calculator.py b/tests/test_litellm/test_gpt_image_cost_calculator.py index 0a2a62b6c97..620c0734980 100644 --- a/tests/test_litellm/test_gpt_image_cost_calculator.py +++ b/tests/test_litellm/test_gpt_image_cost_calculator.py @@ -19,10 +19,13 @@ import pytest import litellm from litellm.types.utils import ( + CompletionTokensDetailsWrapper, ImageResponse, ImageObject, ImageUsage, ImageUsageInputTokensDetails, + PromptTokensDetailsWrapper, + Usage, ) @@ -202,6 +205,71 @@ class TestGPTImageCostRouting: assert cost >= 0 +class TestGPTImage15OutputImageTokens: + """ + Test for GitHub issue #19508: + Image usage calculation does not include image tokens in gpt-image-1.5 + + gpt-image-1.5 returns output_tokens_details with separate image_tokens and text_tokens, + and these must be correctly included in cost calculation. + """ + + def test_gpt_image_15_output_image_tokens_cost(self): + """ + Test that output image tokens are correctly included in cost calculation. + + This tests the fix for issue #19508 where output_tokens_details.image_tokens + were not being included in the cost calculation, causing costs to be + underreported (e.g., $0.046 instead of $0.14). + """ + # Simulate gpt-image-1.5 response with output_tokens_details + # This is what the API returns and what convert_to_image_response transforms + usage = Usage( + prompt_tokens=169, + completion_tokens=4599, + total_tokens=4768, + prompt_tokens_details=PromptTokensDetailsWrapper( + text_tokens=169, + image_tokens=0, + ), + completion_tokens_details=CompletionTokensDetailsWrapper( + text_tokens=439, + image_tokens=4160, + ), + ) + + image_response = ImageResponse( + created=1234567890, + data=[ImageObject(b64_json="test")], + ) + image_response.usage = usage + image_response._hidden_params = {"custom_llm_provider": "openai"} + + cost = litellm.completion_cost( + completion_response=image_response, + model="gpt-image-1.5", + call_type="image_generation", + custom_llm_provider="openai", + ) + + # gpt-image-1.5 pricing: + # - input_cost_per_token: 5e-06 ($5/1M for text input) + # - output_cost_per_token: 1e-05 ($10/1M for text output) + # - output_cost_per_image_token: 3.2e-05 ($32/1M for image output) + # + # Expected cost: + # Input text: 169 * $5/1M = $0.000845 + # Output text: 439 * $10/1M = $0.00439 + # Output image: 4160 * $32/1M = $0.13312 + # Total: $0.138355 + expected_cost = 169 * 5e-06 + 439 * 1e-05 + 4160 * 3.2e-05 + + assert abs(cost - expected_cost) < 1e-6, ( + f"Expected {expected_cost}, got {cost}. " + f"Image tokens may not be included in cost calculation." + ) + + class TestCompletionCostIntegration: """Test the full completion_cost integration for gpt-image-1""" diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index d1185c72a29..80fd9f61298 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -477,6 +477,12 @@ async def test_openai_env_base( respx_mock: respx.MockRouter, env_base, openai_api_response, monkeypatch ): "This tests OpenAI env variables are honored, including legacy OPENAI_API_BASE" + # Clear cache to ensure no cached clients from previous tests interfere + # This prevents cache pollution where a previous test cached a client with + # aiohttp transport, which would bypass respx mocks + if hasattr(litellm, "in_memory_llm_clients_cache"): + litellm.in_memory_llm_clients_cache.flush_cache() + # Ensure aiohttp transport is disabled to use httpx which respx can mock litellm.disable_aiohttp_transport = True diff --git a/tests/test_litellm/test_router_per_deployment_num_retries.py b/tests/test_litellm/test_router_per_deployment_num_retries.py index 4021ca28073..154ba579e4e 100644 --- a/tests/test_litellm/test_router_per_deployment_num_retries.py +++ b/tests/test_litellm/test_router_per_deployment_num_retries.py @@ -32,17 +32,17 @@ class TestPerDeploymentNumRetries: ) deployment = router.model_list[0] - + # Create a mock exception without num_retries class MockException(Exception): pass - + exc = MockException("test error") assert not hasattr(exc, "num_retries") or exc.num_retries is None - + # Call the helper router._set_deployment_num_retries_on_exception(exc, deployment) - + # Verify num_retries was set from deployment assert exc.num_retries == 5 @@ -66,16 +66,16 @@ class TestPerDeploymentNumRetries: ) deployment = router.model_list[0] - + # Create an exception that already has num_retries class MockException(Exception): num_retries = 10 # Already set - + exc = MockException("test error") - + # Call the helper router._set_deployment_num_retries_on_exception(exc, deployment) - + # Verify num_retries was NOT overridden assert exc.num_retries == 10 @@ -99,15 +99,15 @@ class TestPerDeploymentNumRetries: ) deployment = router.model_list[0] - + class MockException(Exception): pass - + exc = MockException("test error") - + # Call the helper router._set_deployment_num_retries_on_exception(exc, deployment) - + # Verify num_retries was not set (deployment has no num_retries) assert not hasattr(exc, "num_retries") or exc.num_retries is None @@ -155,3 +155,36 @@ class TestPerDeploymentNumRetries: kwargs = {} router._update_kwargs_before_fallbacks(model="test-model", kwargs=kwargs) assert kwargs["num_retries"] == 7 # Uses global + + def test_set_deployment_num_retries_with_string_value(self): + """ + Test that _set_deployment_num_retries_on_exception handles string values + from environment variables correctly. + GitHub Issue: #19481 + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/gpt-4", + "api_key": "test-key", + "num_retries": "6", # String value (as from env var) + }, + }, + ], + num_retries=0, # Global setting + ) + + deployment = router.model_list[0] + + class MockException(Exception): + pass + + exc = MockException("test error") + + # Call the helper + router._set_deployment_num_retries_on_exception(exc, deployment) + + # Verify num_retries was converted from string to int + assert exc.num_retries == 6 diff --git a/tests/test_litellm/vector_stores/test_vector_store_registry.py b/tests/test_litellm/vector_stores/test_vector_store_registry.py index 3c4308dd3aa..a3af476bc71 100644 --- a/tests/test_litellm/vector_stores/test_vector_store_registry.py +++ b/tests/test_litellm/vector_stores/test_vector_store_registry.py @@ -119,8 +119,12 @@ def test_add_vector_store_to_registry(): +@respx.mock def test_search_uses_registry_credentials(): """search() should pull credentials from vector_store_registry when available""" + # Block all HTTP requests at the network level to prevent real API calls + respx.route().mock(return_value=httpx.Response(200, json={"object": "list", "data": []})) + vector_store = LiteLLM_ManagedVectorStore( vector_store_id="vs1", custom_llm_provider="bedrock", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts new file mode 100644 index 00000000000..801fbdbb99d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts @@ -0,0 +1,35 @@ +// hooks/useDisableShowPrompts.ts +import { useSyncExternalStore } from "react"; +import { getLocalStorageItem } from "@/utils/localStorageUtils"; +import { LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils"; + +function subscribe(callback: () => void) { + const onStorage = (e: StorageEvent) => { + if (e.key === "disableShowPrompts") { + callback(); + } + }; + + const onCustom = (e: Event) => { + const { key } = (e as CustomEvent).detail; + if (key === "disableShowPrompts") { + callback(); + } + }; + + window.addEventListener("storage", onStorage); + window.addEventListener(LOCAL_STORAGE_EVENT, onCustom); + + return () => { + window.removeEventListener("storage", onStorage); + window.removeEventListener(LOCAL_STORAGE_EVENT, onCustom); + }; +} + +function getSnapshot() { + return getLocalStorageItem("disableShowPrompts") === "true"; +} + +export function useDisableShowPrompts() { + return useSyncExternalStore(subscribe, getSnapshot); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index 97837ff8e0a..97e4c799e72 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -56,8 +56,10 @@ export default function Layout({ children }: { children: React.ReactNode }) { userRole={userRole} premiumUser={premiumUser} proxySettings={undefined} - setProxySettings={() => {}} + setProxySettings={() => { }} accessToken={accessToken} + isDarkMode={false} + toggleDarkMode={() => { }} />
diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index 8ac1f756f96..23c80acf973 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -45,6 +45,7 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { jwtDecode } from "jwt-decode"; import { useSearchParams } from "next/navigation"; import { Suspense, useEffect, useState } from "react"; +import { ConfigProvider, theme } from "antd"; function getCookie(name: string) { // Safer cookie read + decoding; handles '=' inside values @@ -131,6 +132,12 @@ export default function CreateKeyPage() { const [showClaudeCodePrompt, setShowClaudeCodePrompt] = useState(false); const [showClaudeCodeModal, setShowClaudeCodeModal] = useState(false); + // Dark mode state + const [isDarkMode, setIsDarkMode] = useState(false); + const toggleDarkMode = () => { + setIsDarkMode(!isDarkMode); + }; + const invitation_id = searchParams.get("invitation_id"); // Get page from URL, default to 'api-keys' if not present @@ -282,7 +289,7 @@ export default function CreateKeyPage() { const nudgesConfig = await getInProductNudgesCall(accessToken); const isUsingClaudeCode = nudgesConfig?.is_claude_code_enabled || false; setIsClaudeCode(isUsingClaudeCode); - + // Show Claude Code prompt on login if enabled if (isUsingClaudeCode) { setShowClaudeCodePrompt(true); @@ -362,225 +369,231 @@ export default function CreateKeyPage() { return ( }> - - {invitation_id ? ( - - ) : ( -
- + + {invitation_id ? ( + -
-
- -
+ ) : ( +
+ +
+
+ +
- {page == "api-keys" ? ( - - ) : page == "models" ? ( - - ) : page == "llm-playground" ? ( - - ) : page == "users" ? ( - - ) : page == "teams" ? ( - - ) : page == "organizations" ? ( - - ) : page == "admin-panel" ? ( - - ) : page == "api_ref" ? ( - - ) : page == "logging-and-alerts" ? ( - - ) : page == "budgets" ? ( - - ) : page == "guardrails" ? ( - - ) : page == "policies" ? ( - - ) : page == "agents" ? ( - - ) : page == "prompts" ? ( - - ) : page == "transform-request" ? ( - - ) : page == "router-settings" ? ( - - ) : page == "ui-theme" ? ( - - ) : page == "cost-tracking" ? ( - - ) : page == "model-hub-table" ? ( - isAdminRole(userRole) ? ( - + ) : page == "models" ? ( + + ) : page == "llm-playground" ? ( + + ) : page == "users" ? ( + + ) : page == "teams" ? ( + + ) : page == "organizations" ? ( + + ) : page == "admin-panel" ? ( + + ) : page == "api_ref" ? ( + + ) : page == "logging-and-alerts" ? ( + + ) : page == "budgets" ? ( + + ) : page == "guardrails" ? ( + + ) : page == "policies" ? ( + + ) : page == "agents" ? ( + + ) : page == "prompts" ? ( + + ) : page == "transform-request" ? ( + + ) : page == "router-settings" ? ( + + ) : page == "ui-theme" ? ( + + ) : page == "cost-tracking" ? ( + + ) : page == "model-hub-table" ? ( + isAdminRole(userRole) ? ( + + ) : ( + + ) + ) : page == "caching" ? ( + + ) : page == "pass-through-settings" ? ( + + ) : page == "logs" ? ( + + ) : page == "mcp-servers" ? ( + + ) : page == "search-tools" ? ( + + ) : page == "tag-management" ? ( + + ) : page == "claude-code-plugins" ? ( + + ) : page == "vector-stores" ? ( + + ) : page == "new_usage" ? ( + ) : ( - - ) - ) : page == "caching" ? ( - - ) : page == "pass-through-settings" ? ( - - ) : page == "logs" ? ( - - ) : page == "mcp-servers" ? ( - - ) : page == "search-tools" ? ( - - ) : page == "tag-management" ? ( - - ) : page == "claude-code-plugins" ? ( - - ) : page == "vector-stores" ? ( - - ) : page == "new_usage" ? ( - - ) : ( - - )} + + )} +
+ + {/* Survey Components */} + + + + {/* Claude Code Components */} + +
- - {/* Survey Components */} - - - - {/* Claude Code Components */} - - -
- )} -
+ )} + + ); diff --git a/ui/litellm-dashboard/src/components/BulkEditUsers.test.tsx b/ui/litellm-dashboard/src/components/BulkEditUsers.test.tsx new file mode 100644 index 00000000000..4185625e746 --- /dev/null +++ b/ui/litellm-dashboard/src/components/BulkEditUsers.test.tsx @@ -0,0 +1,343 @@ +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi, beforeEach } from "vitest"; +import { renderWithProviders, screen, waitFor } from "../../tests/test-utils"; +import BulkEditUserModal from "./BulkEditUsers"; +import { userBulkUpdateUserCall, teamBulkMemberAddCall } from "./networking"; +import NotificationsManager from "./molecules/notifications_manager"; + +vi.mock("./networking", () => ({ + userBulkUpdateUserCall: vi.fn(), + teamBulkMemberAddCall: vi.fn(), +})); + +vi.mock("./user_edit_view", () => ({ + UserEditView: ({ onSubmit, onCancel }: { onSubmit: (values: any) => void; onCancel: () => void }) => ( +
+ + +
+ ), +})); + +const mockUserBulkUpdateUserCall = vi.mocked(userBulkUpdateUserCall); +const mockTeamBulkMemberAddCall = vi.mocked(teamBulkMemberAddCall); + +const defaultProps = { + open: true, + onCancel: vi.fn(), + selectedUsers: [ + { user_id: "user1", user_email: "user1@example.com", user_role: "user", max_budget: 50 }, + { user_id: "user2", user_email: "user2@example.com", user_role: "admin", max_budget: null }, + ], + possibleUIRoles: { + admin: { ui_label: "Admin", description: "Administrator role" }, + user: { ui_label: "User", description: "Regular user role" }, + }, + accessToken: "test-token", + onSuccess: vi.fn(), + teams: [ + { team_id: "team1", team_alias: "Team 1" }, + { team_id: "team2", team_alias: "Team 2" }, + ], + userRole: "Admin", + userModels: ["gpt-4", "gpt-3.5-turbo"], + allowAllUsers: false, +}; + +describe("BulkEditUserModal", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockUserBulkUpdateUserCall.mockResolvedValue({ + results: [], + total_requested: 2, + successful_updates: 2, + failed_updates: 0, + }); + mockTeamBulkMemberAddCall.mockResolvedValue({ + successful_additions: 2, + failed_additions: 0, + }); + }); + + it("should render without crashing", () => { + renderWithProviders(); + + expect(screen.getByText(`Bulk Edit ${defaultProps.selectedUsers.length} User(s)`)).toBeInTheDocument(); + }); + + it("should display modal title with correct user count", () => { + renderWithProviders(); + + expect(screen.getByText("Bulk Edit 2 User(s)")).toBeInTheDocument(); + }); + + it("should display selected users table when modal is open", () => { + renderWithProviders(); + + expect(screen.getByText("Selected Users (2):")).toBeInTheDocument(); + expect(screen.getByText("user1")).toBeInTheDocument(); + expect(screen.getByText("user2")).toBeInTheDocument(); + expect(screen.getByText("user1@example.com")).toBeInTheDocument(); + expect(screen.getByText("user2@example.com")).toBeInTheDocument(); + }); + + it("should display user roles in table", () => { + renderWithProviders(); + + expect(screen.getByText("User")).toBeInTheDocument(); + expect(screen.getByText("Admin")).toBeInTheDocument(); + }); + + it("should display budget information in table", () => { + renderWithProviders(); + + expect(screen.getByText("$50")).toBeInTheDocument(); + expect(screen.getByText("Unlimited")).toBeInTheDocument(); + }); + + it("should call onCancel when cancel button is clicked", async () => { + const user = userEvent.setup(); + const onCancel = vi.fn(); + renderWithProviders(); + + const cancelButton = screen.getByRole("button", { name: "Cancel" }); + await user.click(cancelButton); + + expect(onCancel).toHaveBeenCalledTimes(1); + }); + + + it("should show update all users checkbox when allowAllUsers is true", () => { + renderWithProviders(); + + expect(screen.getByRole("checkbox", { name: /update all users/i })).toBeInTheDocument(); + }); + + it("should not show update all users checkbox when allowAllUsers is false", () => { + renderWithProviders(); + + expect(screen.queryByRole("checkbox", { name: /update all users/i })).not.toBeInTheDocument(); + }); + + it("should toggle update all users mode", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const checkbox = screen.getByRole("checkbox", { name: /update all users/i }); + expect(checkbox).not.toBeChecked(); + + await user.click(checkbox); + + expect(checkbox).toBeChecked(); + expect(screen.getByText("Bulk Edit All Users")).toBeInTheDocument(); + }); + + it("should show warning message when update all users is enabled", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const checkbox = screen.getByRole("checkbox", { name: /update all users/i }); + await user.click(checkbox); + + expect(screen.getByText(/this will apply changes to all users/i)).toBeInTheDocument(); + }); + + it("should hide selected users table when update all users is enabled", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + expect(screen.getByText("Selected Users (2):")).toBeInTheDocument(); + + const checkbox = screen.getByRole("checkbox", { name: /update all users/i }); + await user.click(checkbox); + + expect(screen.queryByText("Selected Users (2):")).not.toBeInTheDocument(); + }); + + it("should display team management section", () => { + renderWithProviders(); + + expect(screen.getByText("Team Management")).toBeInTheDocument(); + expect(screen.getByRole("checkbox", { name: /add selected users to teams/i })).toBeInTheDocument(); + }); + + it("should show team budget input when add to teams is checked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const addToTeamsCheckbox = screen.getByRole("checkbox", { name: /add selected users to teams/i }); + await user.click(addToTeamsCheckbox); + + expect(screen.getByText("Team Budget (Optional):")).toBeInTheDocument(); + expect(screen.getByPlaceholderText("Max budget per user in team")).toBeInTheDocument(); + }); + + it("should render UserEditView component", () => { + renderWithProviders(); + + expect(screen.getByTestId("user-edit-view")).toBeInTheDocument(); + }); + + it("should show error when access token is missing", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith("Access token not found"); + }); + }); + + it("should call userBulkUpdateUserCall with correct payload for selected users", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(mockUserBulkUpdateUserCall).toHaveBeenCalledWith( + "test-token", + { user_role: "admin", max_budget: 100 }, + ["user1", "user2"], + ); + }); + }); + + it("should call userBulkUpdateUserCall with allUsers flag when update all users is enabled", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const updateAllCheckbox = screen.getByRole("checkbox", { name: /update all users/i }); + await user.click(updateAllCheckbox); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(mockUserBulkUpdateUserCall).toHaveBeenCalledWith( + "test-token", + expect.objectContaining({ user_role: "admin", max_budget: 100 }), + undefined, + true, + ); + }); + }); + + + it("should show success message after successful user update", async () => { + const user = userEvent.setup(); + mockUserBulkUpdateUserCall.mockResolvedValue({ + results: [], + total_requested: 2, + successful_updates: 2, + failed_updates: 0, + }); + + renderWithProviders(); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(NotificationsManager.success).toHaveBeenCalledWith("Updated 2 user(s)"); + }); + }); + + it("should show success message for all users update", async () => { + const user = userEvent.setup(); + mockUserBulkUpdateUserCall.mockResolvedValue({ + results: [], + total_requested: 100, + successful_updates: 100, + failed_updates: 0, + }); + + renderWithProviders(); + + const updateAllCheckbox = screen.getByRole("checkbox", { name: /update all users/i }); + await user.click(updateAllCheckbox); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(NotificationsManager.success).toHaveBeenCalledWith("Updated all users (100 total)"); + }); + }); + + + it("should show error message when bulk update fails", async () => { + const user = userEvent.setup(); + mockUserBulkUpdateUserCall.mockRejectedValueOnce(new Error("Update failed")); + + renderWithProviders(); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith("Failed to perform bulk operations"); + }); + }); + + it("should call onSuccess and onCancel after successful update", async () => { + const user = userEvent.setup(); + const onSuccess = vi.fn(); + const onCancel = vi.fn(); + + renderWithProviders(); + + const submitButton = screen.getByRole("button", { name: "Submit" }); + await user.click(submitButton); + + await waitFor(() => { + expect(onSuccess).toHaveBeenCalledTimes(1); + expect(onCancel).toHaveBeenCalledTimes(1); + }); + }); + + it("should truncate long user IDs in table", () => { + const longUserId = "a".repeat(30); + const propsWithLongId = { + ...defaultProps, + selectedUsers: [{ user_id: longUserId, user_email: "test@example.com", user_role: "user", max_budget: null }], + }; + + renderWithProviders(); + + expect(screen.getByText(new RegExp(`${longUserId.slice(0, 20)}...`))).toBeInTheDocument(); + }); + + it("should display no email text when user email is missing", () => { + const propsWithoutEmail = { + ...defaultProps, + selectedUsers: [{ user_id: "user1", user_email: null, user_role: "user", max_budget: null }], + }; + + renderWithProviders(); + + expect(screen.getByText("No email")).toBeInTheDocument(); + }); + + it("should display role label from possibleUIRoles when available", () => { + renderWithProviders(); + + expect(screen.getByText("Admin")).toBeInTheDocument(); + expect(screen.getByText("User")).toBeInTheDocument(); + }); + + it("should display role key when ui_label is not available", () => { + const propsWithoutUIRoles = { + ...defaultProps, + possibleUIRoles: null, + }; + + renderWithProviders(); + + expect(screen.getByText("user")).toBeInTheDocument(); + expect(screen.getByText("admin")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/bulk_edit_user.tsx b/ui/litellm-dashboard/src/components/BulkEditUsers.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/bulk_edit_user.tsx rename to ui/litellm-dashboard/src/components/BulkEditUsers.tsx index b3847911cc0..2f3e57ff2a8 100644 --- a/ui/litellm-dashboard/src/components/bulk_edit_user.tsx +++ b/ui/litellm-dashboard/src/components/BulkEditUsers.tsx @@ -18,7 +18,7 @@ import NotificationsManager from "./molecules/notifications_manager"; const { Text, Title } = Typography; interface BulkEditUserModalProps { - visible: boolean; + open: boolean; onCancel: () => void; selectedUsers: any[]; possibleUIRoles: Record> | null; @@ -31,7 +31,7 @@ interface BulkEditUserModalProps { } const BulkEditUserModal: React.FC = ({ - visible, + open, onCancel, selectedUsers, possibleUIRoles, @@ -75,7 +75,7 @@ const BulkEditUserModal: React.FC = ({ keys: [], teams: teams || [], }), - [teams, visible], + [teams, open], ); const handleSubmit = async (formValues: any) => { @@ -145,7 +145,7 @@ const BulkEditUserModal: React.FC = ({ if (updateAllUsers) { members = null; } else { - const members = selectedUsers.map((user) => ({ + members = selectedUsers.map((user) => ({ user_id: user.user_id, role: "user" as const, // Default role for bulk add user_email: user.user_email || null, @@ -214,7 +214,7 @@ const BulkEditUserModal: React.FC = ({ return ( { + const Option = ({ children, value }: any) => ( + + ); + const Select = ({ children, value, onChange, placeholder }: any) => ( + + ); + Select.Option = Option; + return { + Select, + Tooltip: ({ children, title }: any) => ( +
+ {children} +
+ ), + Switch: ({ checked, onChange }: any) => ( + onChange(e.target.checked)} + /> + ), + Divider: () =>
, + }; +}); + +vi.mock("@ant-design/icons", () => ({ + InfoCircleOutlined: () => ℹ, +})); + +vi.mock("@tremor/react", () => ({ + TextInput: ({ value, onValueChange, onChange, placeholder, name, className }: any) => { + const handleChange = (e: React.ChangeEvent) => { + if (onChange) { + onChange(e); + } + if (onValueChange) { + onValueChange(e.target.value); + } + }; + return ( + + ); + }, +})); + +describe("KeyLifecycleSettings", () => { + const mockForm = { + getFieldValue: vi.fn(), + setFieldValue: vi.fn(), + setFieldsValue: vi.fn(), + }; + + const defaultProps = { + form: mockForm, + autoRotationEnabled: false, + onAutoRotationChange: vi.fn(), + rotationInterval: "", + onRotationIntervalChange: vi.fn(), + isCreateMode: false, + }; + + beforeEach(() => { + vi.clearAllMocks(); + mockForm.getFieldValue.mockReturnValue(""); + }); + + it("should render without crashing", () => { + renderWithProviders(); + + expect(screen.getByText("Key Expiry Settings")).toBeInTheDocument(); + expect(screen.getByText("Auto-Rotation Settings")).toBeInTheDocument(); + }); + + describe("Key Expiry Settings", () => { + it("should render expiry input field", () => { + renderWithProviders(); + + expect(screen.getByText("Expire Key")).toBeInTheDocument(); + expect(screen.getByTestId("duration-input")).toBeInTheDocument(); + }); + + it("should show correct placeholder in create mode", () => { + renderWithProviders(); + + const input = screen.getByTestId("duration-input"); + expect(input).toHaveAttribute( + "placeholder", + "e.g., 30d or leave empty to never expire" + ); + }); + + it("should show correct placeholder in edit mode", () => { + renderWithProviders(); + + const input = screen.getByTestId("duration-input"); + expect(input).toHaveAttribute("placeholder", "e.g., 30d or -1 to never expire"); + }); + + it("should show correct tooltip in create mode", () => { + renderWithProviders(); + + const tooltips = screen.getAllByTestId("tooltip"); + const expiryTooltip = tooltips.find((tooltip) => + tooltip.getAttribute("title")?.includes("Leave empty to never expire") + ); + expect(expiryTooltip).toBeInTheDocument(); + expect(expiryTooltip).toHaveAttribute( + "title", + "Set when this key should expire. Format: 30s (seconds), 30m (minutes), 30h (hours), 30d (days). Leave empty to never expire." + ); + }); + + it("should show correct tooltip in edit mode", () => { + renderWithProviders(); + + const tooltips = screen.getAllByTestId("tooltip"); + const expiryTooltip = tooltips.find((tooltip) => + tooltip.getAttribute("title")?.includes("Use -1 to never expire") + ); + expect(expiryTooltip).toBeInTheDocument(); + expect(expiryTooltip).toHaveAttribute( + "title", + "Set when this key should expire. Format: 30s (seconds), 30m (minutes), 30h (hours), 30d (days). Use -1 to never expire." + ); + }); + + it("should initialize with form value if present", () => { + mockForm.getFieldValue.mockReturnValue("30d"); + renderWithProviders(); + + const input = screen.getByTestId("duration-input") as HTMLInputElement; + expect(input.value).toBe("30d"); + }); + + it("should update form using setFieldValue when duration changes", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const input = screen.getByTestId("duration-input"); + await user.type(input, "60d"); + + expect(mockForm.setFieldValue).toHaveBeenCalledWith("duration", "60d"); + }); + + it("should update form using setFieldsValue when setFieldValue is not available", async () => { + const user = userEvent.setup(); + const formWithoutSetFieldValue = { + getFieldValue: vi.fn().mockReturnValue(""), + setFieldsValue: vi.fn(), + }; + renderWithProviders( + + ); + + const input = screen.getByTestId("duration-input"); + await user.type(input, "90d"); + + expect(formWithoutSetFieldValue.setFieldsValue).toHaveBeenCalledWith({ duration: "90d" }); + }); + }); + + describe("Auto-Rotation Settings", () => { + it("should render auto-rotation switch", () => { + renderWithProviders(); + + expect(screen.getByText("Enable Auto-Rotation")).toBeInTheDocument(); + expect(screen.getByTestId("switch")).toBeInTheDocument(); + }); + + it("should show switch as unchecked when autoRotationEnabled is false", () => { + renderWithProviders(); + + const switchElement = screen.getByTestId("switch") as HTMLInputElement; + expect(switchElement.checked).toBe(false); + }); + + it("should show switch as checked when autoRotationEnabled is true", () => { + renderWithProviders(); + + const switchElement = screen.getByTestId("switch") as HTMLInputElement; + expect(switchElement.checked).toBe(true); + }); + + it("should call onAutoRotationChange when switch is toggled", async () => { + const user = userEvent.setup(); + const onAutoRotationChange = vi.fn(); + renderWithProviders( + + ); + + const switchElement = screen.getByTestId("switch"); + await user.click(switchElement); + + expect(onAutoRotationChange).toHaveBeenCalledWith(true); + }); + + it("should not show rotation interval section when auto-rotation is disabled", () => { + renderWithProviders(); + + expect(screen.queryByText("Rotation Interval")).not.toBeInTheDocument(); + expect(screen.queryByTestId("select")).not.toBeInTheDocument(); + }); + + it("should show rotation interval section when auto-rotation is enabled", () => { + renderWithProviders( + + ); + + expect(screen.getByText("Rotation Interval")).toBeInTheDocument(); + expect(screen.getByTestId("select")).toBeInTheDocument(); + }); + + it("should show all predefined interval options", () => { + renderWithProviders( + + ); + + expect(screen.getByText("7 days")).toBeInTheDocument(); + expect(screen.getByText("30 days")).toBeInTheDocument(); + expect(screen.getByText("90 days")).toBeInTheDocument(); + expect(screen.getByText("180 days")).toBeInTheDocument(); + expect(screen.getByText("365 days")).toBeInTheDocument(); + expect(screen.getByText("Custom interval")).toBeInTheDocument(); + }); + + it("should display current rotation interval in select", () => { + renderWithProviders( + + ); + + const select = screen.getByTestId("select") as HTMLSelectElement; + expect(select.value).toBe("90d"); + }); + + it("should call onRotationIntervalChange when predefined interval is selected", async () => { + const user = userEvent.setup(); + const onRotationIntervalChange = vi.fn(); + renderWithProviders( + + ); + + const select = screen.getByTestId("select"); + await user.selectOptions(select, "30d"); + + expect(onRotationIntervalChange).toHaveBeenCalledWith("30d"); + }); + + it("should show custom input when custom option is selected", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + const select = screen.getByTestId("select"); + await user.selectOptions(select, "custom"); + + expect(screen.getByTestId("custom-interval-input")).toBeInTheDocument(); + expect(screen.getByText("Supported formats: seconds (s), minutes (m), hours (h), days (d)")).toBeInTheDocument(); + }); + + it("should hide custom input when predefined interval is selected after custom", async () => { + const user = userEvent.setup(); + const onRotationIntervalChange = vi.fn(); + renderWithProviders( + + ); + + const select = screen.getByTestId("select"); + await user.selectOptions(select, "7d"); + + expect(screen.queryByTestId("custom-interval-input")).not.toBeInTheDocument(); + expect(onRotationIntervalChange).toHaveBeenCalledWith("7d"); + }); + + it("should call onRotationIntervalChange when custom interval is entered", async () => { + const user = userEvent.setup(); + const onRotationIntervalChange = vi.fn(); + renderWithProviders( + + ); + + const select = screen.getByTestId("select"); + await user.selectOptions(select, "custom"); + + const customInput = screen.getByTestId("custom-interval-input"); + await user.type(customInput, "14d"); + + expect(onRotationIntervalChange).toHaveBeenCalledWith("14d"); + }); + + it("should show info message when auto-rotation is enabled", () => { + renderWithProviders(); + + expect( + screen.getByText( + "When rotation occurs, you'll receive a notification with the new key. The old key will be deactivated after a brief grace period." + ) + ).toBeInTheDocument(); + }); + + it("should not show info message when auto-rotation is disabled", () => { + renderWithProviders(); + + expect( + screen.queryByText( + "When rotation occurs, you'll receive a notification with the new key. The old key will be deactivated after a brief grace period." + ) + ).not.toBeInTheDocument(); + }); + + it("should initialize with custom interval input visible when custom interval is provided", () => { + renderWithProviders( + + ); + + expect(screen.getByTestId("custom-interval-input")).toBeInTheDocument(); + const customInput = screen.getByTestId("custom-interval-input") as HTMLInputElement; + expect(customInput.value).toBe("14d"); + }); + + it("should show custom option selected when custom interval is provided", () => { + renderWithProviders( + + ); + + const select = screen.getByTestId("select") as HTMLSelectElement; + expect(select.value).toBe("custom"); + }); + + it("should not call onRotationIntervalChange when selecting custom option", async () => { + const user = userEvent.setup(); + const onRotationIntervalChange = vi.fn(); + renderWithProviders( + + ); + + const select = screen.getByTestId("select"); + await user.selectOptions(select, "custom"); + + expect(onRotationIntervalChange).not.toHaveBeenCalled(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx b/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx index 81d22b56347..0f29a47d1dc 100644 --- a/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx +++ b/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx @@ -11,6 +11,7 @@ interface KeyLifecycleSettingsProps { onAutoRotationChange: (enabled: boolean) => void; rotationInterval: string; onRotationIntervalChange: (interval: string) => void; + isCreateMode?: boolean; // If true, shows "leave empty to never expire" instead of "-1 to never expire" } const KeyLifecycleSettings: React.FC = ({ @@ -19,6 +20,7 @@ const KeyLifecycleSettings: React.FC = ({ onAutoRotationChange, rotationInterval, onRotationIntervalChange, + isCreateMode = false, }) => { // Predefined intervals const predefinedIntervals = ["7d", "30d", "90d", "180d", "365d"]; @@ -64,13 +66,19 @@ const KeyLifecycleSettings: React.FC = ({
({ updateGuardrailCall: vi.fn(), })); + +// Mock ContentFilterManager +vi.mock("./content_filter/ContentFilterManager", () => ({ + __esModule: true, + default: ({ onUnsavedChanges, onDataChange, isEditing }: any) => ( +
+ {isEditing && ( + + )} +
+ ), + formatContentFilterDataForAPI: (patterns: any[], blockedWords: any[]) => ({ + patterns, + blocked_words: blockedWords, + }), +})); + describe("Guardrail Info", () => { afterEach(() => { vi.clearAllMocks(); @@ -147,4 +169,101 @@ describe("Guardrail Info", () => { expect(getByText("PII Entity Configuration")).toBeInTheDocument(); }); }); + it("should handle content filter updates correctly", async () => { + // Mock the network responses + vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ + guardrail_id: "123", + guardrail_name: "Content Filter Guardrail", + litellm_params: { + guardrail: "litellm_content_filter", + mode: "pre_call", + default_on: true, + patterns: ["initial_pattern"], + blocked_words: ["initial_word"], + }, + created_at: "2024-01-01T00:00:00Z", + updated_at: "2024-01-01T00:00:00Z", + guardrail_definition_location: "database", + }); + + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + supported_entities: [], + supported_actions: [], + pii_entity_categories: [], + supported_modes: ["pre_call", "post_call"], + }); + + vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({}); + vi.mocked(networking.updateGuardrailCall).mockResolvedValue({ status: "success" }); + + const { getByText, getByRole, getAllByRole, getByLabelText } = render( + { }} accessToken="123" isAdmin={true} />, + ); + + await waitFor(() => { + expect(getByText("Settings")).toBeInTheDocument(); + }); + + // Go to Settings tab + fireEvent.click(getByText("Settings")); + + await waitFor(() => { + expect(getByText("Guardrail Settings")).toBeInTheDocument(); + }); + + // Enter Edit Mode + fireEvent.click(getByText("Edit Settings")); + + // Modify Guardrail Name to force an update + const nameInput = getByLabelText("Guardrail Name"); + fireEvent.change(nameInput, { target: { value: "Updated Name" } }); + + // Save with only name change + const saveButton = getByText("Save Changes"); + fireEvent.click(saveButton); + + await waitFor(() => { + expect(networking.updateGuardrailCall).toHaveBeenCalled(); + }); + + // Verify call did NOT include patterns or blocked_words (because no changes) + // updateGuardrailCall(accessToken, guardrailId, updateData) -> index 2 is updateData + const firstCallArgs: any = vi.mocked(networking.updateGuardrailCall).mock.calls[0][2]; + + // Verify attributes that definitely changed + expect(firstCallArgs.guardrail_name).toBe("Updated Name"); + + // litellm_params might be undefined if empty, which is correct. + // If it exists, ensure patterns/blocked_words are not in it. + if (firstCallArgs.litellm_params) { + expect(firstCallArgs.litellm_params.patterns).toBeUndefined(); + expect(firstCallArgs.litellm_params.blocked_words).toBeUndefined(); + } + + // Clear mocks to reset call count + vi.clearAllMocks(); + + // Enter Edit Mode again to make changes + await waitFor(() => { + expect(getByText("Edit Settings")).toBeInTheDocument(); + }); + fireEvent.click(getByText("Edit Settings")); + + // Now modify the values using the mock button + const simulateChangeButton = getByText("Simulate Change"); + fireEvent.click(simulateChangeButton); + + // Save again + fireEvent.click(getByText("Save Changes")); + + await waitFor(() => { + expect(networking.updateGuardrailCall).toHaveBeenCalled(); + }); + + // Verify call INCLUDES patterns and blocked_words + const secondCallArgs: any = vi.mocked(networking.updateGuardrailCall).mock.calls[0][2]; + expect(secondCallArgs.litellm_params).toBeDefined(); + expect(secondCallArgs.litellm_params.patterns).toEqual(["new_pattern"]); + expect(secondCallArgs.litellm_params.blocked_words).toEqual(["new_word"]); + }); }); diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx index 29c2e7d13ee..0a15446db97 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx @@ -270,8 +270,11 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, updateData.litellm_params.pii_entities_config = newPiiEntitiesConfig; } - // Always include Content Filter patterns to ensure they persist when other fields change - if (guardrailData.litellm_params?.guardrail === "litellm_content_filter") { + // Only add Content Filter patterns if there are changes + if (guardrailData.litellm_params?.guardrail === "litellm_content_filter" && hasUnsavedContentFilterChanges) { + const originalPatterns = guardrailData.litellm_params?.patterns || []; + const originalBlockedWords = guardrailData.litellm_params?.blocked_words || []; + const formattedData = formatContentFilterDataForAPI( contentFilterDataRef.current.patterns || [], contentFilterDataRef.current.blockedWords || [], @@ -349,6 +352,9 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, console.log("allowedParams: ", allowedParams); allowedParams.forEach((paramName) => { + if (paramName === "patterns" || paramName === "blocked_words") { + return; + } // Check for both direct parameter name and nested optional_params object let paramValue = values[paramName]; if (paramValue === undefined || paramValue === null || paramValue === "") { diff --git a/ui/litellm-dashboard/src/components/navbar.test.tsx b/ui/litellm-dashboard/src/components/navbar.test.tsx index 5d3c7254ff9..9fa32cf9cb0 100644 --- a/ui/litellm-dashboard/src/components/navbar.test.tsx +++ b/ui/litellm-dashboard/src/components/navbar.test.tsx @@ -16,6 +16,7 @@ vi.mock("@/utils/proxyUtils", () => ({ let mockUseThemeImpl = () => ({ logoUrl: null as string | null }); let mockUseHealthReadinessImpl = () => ({ data: null as any }); let mockGetLocalStorageItemImpl = () => null as string | null; +let mockUseDisableShowPromptsImpl = () => false; vi.mock("@/contexts/ThemeContext", () => ({ useTheme: () => mockUseThemeImpl(), @@ -25,7 +26,12 @@ vi.mock("@/app/(dashboard)/hooks/healthReadiness/useHealthReadiness", () => ({ useHealthReadiness: () => mockUseHealthReadinessImpl(), })); +vi.mock("@/app/(dashboard)/hooks/useDisableShowPrompts", () => ({ + useDisableShowPrompts: () => mockUseDisableShowPromptsImpl(), +})); + vi.mock("@/utils/localStorageUtils", () => ({ + LOCAL_STORAGE_EVENT: "local-storage-change", getLocalStorageItem: () => mockGetLocalStorageItemImpl(), setLocalStorageItem: vi.fn(), removeLocalStorageItem: vi.fn(), @@ -52,6 +58,8 @@ describe("Navbar", () => { setProxySettings: vi.fn(), accessToken: "test-token", isPublicPage: false, + isDarkMode: false, + toggleDarkMode: vi.fn(), }; it("should render without crashing", () => { @@ -198,4 +206,11 @@ describe("Navbar", () => { expect(cookieUtils.clearTokenCookies).toHaveBeenCalled(); expect(window.location.href).toBe(""); }); + + it("should not render dark mode toggle slider", () => { + renderWithProviders(); + + // DO NOT RENDER THIS UNTIL ALL COMPONENTS ARE CONFIRMED TO SUPPORT DARK MODE STYLES. IT IS AN ISSUE IF THIS TEST FAILS. + expect(screen.queryByTestId("dark-mode-toggle")).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/navbar.tsx b/ui/litellm-dashboard/src/components/navbar.tsx index 6dac073b3a6..c78f355ff19 100644 --- a/ui/litellm-dashboard/src/components/navbar.tsx +++ b/ui/litellm-dashboard/src/components/navbar.tsx @@ -1,4 +1,5 @@ import { useHealthReadiness } from "@/app/(dashboard)/hooks/healthReadiness/useHealthReadiness"; +import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts"; import { getProxyBaseUrl } from "@/components/networking"; import { useTheme } from "@/contexts/ThemeContext"; import { clearTokenCookies } from "@/utils/cookieUtils"; @@ -16,8 +17,10 @@ import { MailOutlined, MenuFoldOutlined, MenuUnfoldOutlined, + MoonOutlined, SafetyOutlined, SlackOutlined, + SunOutlined, UserOutlined, } from "@ant-design/icons"; import type { MenuProps } from "antd"; @@ -36,6 +39,8 @@ interface NavbarProps { isPublicPage: boolean; sidebarCollapsed?: boolean; onToggleSidebar?: () => void; + isDarkMode: boolean; + toggleDarkMode: () => void; } const Navbar: React.FC = ({ @@ -49,11 +54,13 @@ const Navbar: React.FC = ({ isPublicPage = false, sidebarCollapsed = false, onToggleSidebar, + isDarkMode, + toggleDarkMode }) => { const baseUrl = getProxyBaseUrl(); - console.log("baseUrl", baseUrl); const [logoutUrl, setLogoutUrl] = useState(""); const [disableShowNewBadge, setDisableShowNewBadge] = useState(false); + const disableShowPrompts = useDisableShowPrompts(); const { logoUrl } = useTheme(); const { data: healthData } = useHealthReadiness(); const version = healthData?.litellm_version; @@ -152,6 +159,27 @@ const Navbar: React.FC = ({ aria-label="Toggle hide new feature indicators" />
+
e.stopPropagation()} + > + Hide All Prompts + { + if (checked) { + setLocalStorageItem("disableShowPrompts", "true"); + emitLocalStorageChange("disableShowPrompts"); + } else { + removeLocalStorageItem("disableShowPrompts"); + emitLocalStorageChange("disableShowPrompts"); + } + }} + aria-label="Toggle hide all prompts" + /> +
), @@ -227,6 +255,15 @@ const Navbar: React.FC = ({ > Star us on GitHub + {/* Dark mode is currently a work in progress. To test, you can change 'false' to 'true' below. + Do not set this to true by default until all components are confirmed to support dark mode styles. */} + {false && } + unCheckedChildren={} + />} => { + try { + let url = proxyBaseUrl ? `${proxyBaseUrl}/rag/ingest` : `/rag/ingest`; + + const formData = new FormData(); + formData.append("file", file); + + const ingestOptions: any = { + ingest_options: { + vector_store: { + custom_llm_provider: customLlmProvider, + ...(vectorStoreId && { vector_store_id: vectorStoreId }), + }, + }, + }; + + // Add litellm_vector_store_params if name or description provided + if (vectorStoreName || vectorStoreDescription) { + ingestOptions.ingest_options.litellm_vector_store_params = {}; + if (vectorStoreName) { + ingestOptions.ingest_options.litellm_vector_store_params.vector_store_name = vectorStoreName; + } + if (vectorStoreDescription) { + ingestOptions.ingest_options.litellm_vector_store_params.vector_store_description = vectorStoreDescription; + } + } + + formData.append("request", JSON.stringify(ingestOptions)); + + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + }, + body: formData, + }); + + if (!response.ok) { + const error = await response.json(); + throw new Error(error.error?.message || error.detail || "Failed to ingest document"); + } + + return await response.json(); + } catch (error) { + console.error("Error ingesting document:", error); + throw error; + } +}; + export const getEmailEventSettings = async (accessToken: string): Promise => { try { const url = proxyBaseUrl ? `${proxyBaseUrl}/email/event_settings` : `/email/event_settings`; diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 1edbba28afc..8e89b77bc17 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -162,7 +162,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { const [userSearchLoading, setUserSearchLoading] = useState(false); const [mcpAccessGroups, setMcpAccessGroups] = useState([]); const [disabledCallbacks, setDisabledCallbacks] = useState([]); - const [keyType, setKeyType] = useState("default"); + const [keyType, setKeyType] = useState("llm_api"); const [modelAliases, setModelAliases] = useState<{ [key: string]: string }>({}); const [autoRotationEnabled, setAutoRotationEnabled] = useState(false); const [rotationInterval, setRotationInterval] = useState("30d"); @@ -173,7 +173,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { form.resetFields(); setLoggingSettings([]); setDisabledCallbacks([]); - setKeyType("default"); + setKeyType("llm_api"); setModelAliases({}); setAutoRotationEnabled(false); setRotationInterval("30d"); @@ -188,7 +188,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { form.resetFields(); setLoggingSettings([]); setDisabledCallbacks([]); - setKeyType("default"); + setKeyType("llm_api"); setModelAliases({}); setAutoRotationEnabled(false); setRotationInterval("30d"); @@ -321,9 +321,9 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { formValues.rotation_interval = rotationInterval; } - // Handle duration field for key expiry - if (formValues.duration) { - formValues.duration = formValues.duration; + // Handle duration field for key expiry - convert empty string to null + if (!formValues.duration || formValues.duration.trim() === "") { + formValues.duration = null; } // Update the formValues with the final metadata @@ -707,11 +707,11 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { } name="key_type" - initialValue="default" + initialValue="llm_api" className="mt-4" > setVectorStoreName(e.target.value)} + placeholder="e.g., Product Documentation, Customer Support KB" + size="large" + className="rounded-md" + /> + + + + Description{" "} + + + + + } + > + setVectorStoreDescription(e.target.value)} + placeholder="e.g., Contains all product documentation and user guides" + rows={2} + size="large" + className="rounded-md" + /> + + + + Provider{" "} + + + + + } + required + > + + + + +
+ +
+
+ + + {/* Success Message */} + {ingestResults.length > 0 && ( + +

+ Vector Store ID: {ingestResults[0]?.vector_store_id} +

+

+ Documents Ingested: {ingestResults.length} +

+ + } + type="success" + showIcon + closable + /> + )} + + ); +}; + +export default CreateVectorStore; diff --git a/ui/litellm-dashboard/src/components/vector_store_management/DocumentsTable.test.tsx b/ui/litellm-dashboard/src/components/vector_store_management/DocumentsTable.test.tsx new file mode 100644 index 00000000000..761bc2b7df7 --- /dev/null +++ b/ui/litellm-dashboard/src/components/vector_store_management/DocumentsTable.test.tsx @@ -0,0 +1,102 @@ +import { render, screen, fireEvent, act } from "@testing-library/react"; +import { describe, it, expect, vi } from "vitest"; +import DocumentsTable from "./DocumentsTable"; +import { DocumentUpload } from "./types"; + +// Mock antd message +vi.mock("antd", async () => { + const actual = await vi.importActual("antd"); + return { + ...actual, + message: { + success: vi.fn(), + }, + }; +}); + +describe("DocumentsTable", () => { + const mockDocuments: DocumentUpload[] = [ + { + uid: "1", + name: "test1.pdf", + status: "done", + size: 1024000, + type: "application/pdf", + }, + { + uid: "2", + name: "test2.txt", + status: "uploading", + size: 2048000, + type: "text/plain", + }, + { + uid: "3", + name: "test3.docx", + status: "error", + size: 512000, + type: "application/vnd.openxmlformats-officedocument.wordprocessingml.document", + }, + ]; + + it("should render the table successfully", () => { + const onRemove = vi.fn(); + render(); + + expect(screen.getByText("test1.pdf")).toBeInTheDocument(); + expect(screen.getByText("test2.txt")).toBeInTheDocument(); + expect(screen.getByText("test3.docx")).toBeInTheDocument(); + }); + + it("should display correct status badges", () => { + const onRemove = vi.fn(); + render(); + + expect(screen.getByText("Ready")).toBeInTheDocument(); + expect(screen.getByText("Uploading")).toBeInTheDocument(); + expect(screen.getByText("Error")).toBeInTheDocument(); + }); + + it("should display file sizes", () => { + const onRemove = vi.fn(); + render(); + + expect(screen.getByText(/1000.00 KB/)).toBeInTheDocument(); + expect(screen.getByText(/2.00 MB/)).toBeInTheDocument(); + expect(screen.getByText(/500.00 KB/)).toBeInTheDocument(); + }); + + it("should call onRemove when delete button is clicked", () => { + const onRemove = vi.fn(); + render(); + + const deleteButtons = screen.getAllByLabelText(/delete/i); + + act(() => { + fireEvent.click(deleteButtons[0]); + }); + + expect(onRemove).toHaveBeenCalledWith("1"); + }); + + it("should show empty state when no documents", () => { + const onRemove = vi.fn(); + render(); + + expect(screen.getByText(/No documents uploaded yet/)).toBeInTheDocument(); + }); + + it("should have action buttons for each document", () => { + const onRemove = vi.fn(); + render(); + + // Each document should have 3 action buttons (view, copy, delete) + const viewButtons = screen.getAllByLabelText(/eye/i); + const copyButtons = screen.getAllByLabelText(/copy/i); + const deleteButtons = screen.getAllByLabelText(/delete/i); + + expect(viewButtons).toHaveLength(3); + expect(copyButtons).toHaveLength(3); + expect(deleteButtons).toHaveLength(3); + }); +}); diff --git a/ui/litellm-dashboard/src/components/vector_store_management/DocumentsTable.tsx b/ui/litellm-dashboard/src/components/vector_store_management/DocumentsTable.tsx new file mode 100644 index 00000000000..aeb4240d366 --- /dev/null +++ b/ui/litellm-dashboard/src/components/vector_store_management/DocumentsTable.tsx @@ -0,0 +1,98 @@ +import React from "react"; +import { Table, Badge, Tooltip, message } from "antd"; +import { EyeOutlined, CopyOutlined, DeleteOutlined } from "@ant-design/icons"; +import { DocumentUpload } from "./types"; + +interface DocumentsTableProps { + documents: DocumentUpload[]; + onRemove: (uid: string) => void; +} + +const DocumentsTable: React.FC = ({ documents, onRemove }) => { + const handleCopyId = (uid: string) => { + navigator.clipboard.writeText(uid); + message.success("Document ID copied to clipboard"); + }; + + const getStatusBadge = (status: DocumentUpload["status"]) => { + const statusConfig = { + uploading: { color: "blue", text: "Uploading" }, + done: { color: "green", text: "Ready" }, + error: { color: "red", text: "Error" }, + removed: { color: "default", text: "Removed" }, + }; + + const config = statusConfig[status]; + return ; + }; + + const formatFileSize = (bytes?: number) => { + if (!bytes) return "-"; + const kb = bytes / 1024; + if (kb < 1024) return `${kb.toFixed(2)} KB`; + return `${(kb / 1024).toFixed(2)} MB`; + }; + + const columns = [ + { + title: "Name", + dataIndex: "name", + key: "name", + render: (name: string, record: DocumentUpload) => ( +
+ {name} + {record.size && ({formatFileSize(record.size)})} +
+ ), + }, + { + title: "Status", + dataIndex: "status", + key: "status", + width: 150, + render: (status: DocumentUpload["status"]) => getStatusBadge(status), + }, + { + title: "Actions", + key: "actions", + width: 120, + render: (_: any, record: DocumentUpload) => ( +
+ + console.log("View", record)} + /> + + + handleCopyId(record.uid)} + /> + + + onRemove(record.uid)} + /> + +
+ ), + }, + ]; + + return ( + + ); +}; + +export default DocumentsTable; diff --git a/ui/litellm-dashboard/src/components/vector_store_management/TestVectorStoreTab.test.tsx b/ui/litellm-dashboard/src/components/vector_store_management/TestVectorStoreTab.test.tsx new file mode 100644 index 00000000000..ad6ef02e87b --- /dev/null +++ b/ui/litellm-dashboard/src/components/vector_store_management/TestVectorStoreTab.test.tsx @@ -0,0 +1,90 @@ +import { render, screen, fireEvent } from "@testing-library/react"; +import { describe, it, expect, vi } from "vitest"; +import TestVectorStoreTab from "./TestVectorStoreTab"; +import { VectorStore } from "./types"; + +// Mock VectorStoreTester component +vi.mock("./VectorStoreTester", () => ({ + VectorStoreTester: ({ vectorStoreId, accessToken }: { vectorStoreId: string; accessToken: string }) => ( +
+
{vectorStoreId}
+
{accessToken}
+
+ ), +})); + +const mockVectorStores: VectorStore[] = [ + { + vector_store_id: "vs_123", + custom_llm_provider: "openai", + vector_store_name: "Test Store 1", + vector_store_description: "Description 1", + created_at: "2024-01-01T00:00:00Z", + updated_at: "2024-01-01T00:00:00Z", + }, + { + vector_store_id: "vs_456", + custom_llm_provider: "bedrock", + vector_store_name: "Test Store 2", + vector_store_description: "Description 2", + created_at: "2024-01-02T00:00:00Z", + updated_at: "2024-01-02T00:00:00Z", + }, +]; + +describe("TestVectorStoreTab", () => { + it("should render the component successfully", () => { + render(); + + expect(screen.getByText("Select Vector Store")).toBeInTheDocument(); + expect(screen.getByText("Choose a vector store to test search queries against")).toBeInTheDocument(); + }); + + it("should show message when no access token", () => { + render(); + + expect(screen.getByText("Access token is required to test vector stores.")).toBeInTheDocument(); + }); + + it("should show message when no vector stores available", () => { + render(); + + expect(screen.getByText("No vector stores available. Create one first to test it.")).toBeInTheDocument(); + }); + + it("should render VectorStoreTester with first vector store by default", () => { + render(); + + expect(screen.getByTestId("vector-store-tester")).toBeInTheDocument(); + expect(screen.getByTestId("tester-vector-store-id")).toHaveTextContent("vs_123"); + expect(screen.getByTestId("tester-access-token")).toHaveTextContent("test-token"); + }); + + it("should update VectorStoreTester when selecting different vector store", () => { + render(); + + // Find the select component + const selectElement = screen.getByRole("combobox"); + + // Change selection + fireEvent.mouseDown(selectElement); + + // Wait for options to appear and click the second one + const option2 = screen.getByText("Test Store 2"); + fireEvent.click(option2); + + // Verify the tester component updated + expect(screen.getByTestId("tester-vector-store-id")).toHaveTextContent("vs_456"); + }); + + it("should display vector store names in select options", () => { + render(); + + const selectElement = screen.getByRole("combobox"); + fireEvent.mouseDown(selectElement); + + // Use getAllByText since the selected value also shows the name + expect(screen.getAllByText("Test Store 1").length).toBeGreaterThan(0); + expect(screen.getByText("Test Store 2")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/vector_store_management/TestVectorStoreTab.tsx b/ui/litellm-dashboard/src/components/vector_store_management/TestVectorStoreTab.tsx new file mode 100644 index 00000000000..da005491ca6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/vector_store_management/TestVectorStoreTab.tsx @@ -0,0 +1,75 @@ +import React, { useState } from "react"; +import { Card, Select, Typography } from "antd"; +import { VectorStoreTester } from "./VectorStoreTester"; +import { VectorStore } from "./types"; + +const { Text, Title } = Typography; + +interface TestVectorStoreTabProps { + accessToken: string | null; + vectorStores: VectorStore[]; +} + +const TestVectorStoreTab: React.FC = ({ accessToken, vectorStores }) => { + const [selectedVectorStoreId, setSelectedVectorStoreId] = useState( + vectorStores.length > 0 ? vectorStores[0].vector_store_id : undefined + ); + + if (!accessToken) { + return ( + + Access token is required to test vector stores. + + ); + } + + if (vectorStores.length === 0) { + return ( + +
+ No vector stores available. Create one first to test it. +
+
+ ); + } + + return ( +
+ +
+
+ Select Vector Store + Choose a vector store to test search queries against +
+ + +
+
+ + {selectedVectorStoreId && ( + + )} +
+ ); +}; + +export default TestVectorStoreTab; diff --git a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.test.tsx b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.test.tsx index 65d15260c4c..0e2be7f62df 100644 --- a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.test.tsx +++ b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.test.tsx @@ -130,9 +130,9 @@ describe("VectorStoreTable", () => { expect(screen.getByText("Provider")).toBeInTheDocument(); expect(screen.getByText("Created At")).toBeInTheDocument(); expect(screen.getByText("Updated At")).toBeInTheDocument(); - // Check that we have the expected number of header cells (6 data + 1 actions) + // Check that we have the expected number of header cells (7 data + 1 actions) const headers = screen.getAllByRole("columnheader"); - expect(headers).toHaveLength(7); + expect(headers).toHaveLength(8); }); it("should render all vector store rows", () => { @@ -183,7 +183,7 @@ describe("VectorStoreTable", () => { it("should render fallback for missing name", () => { renderComponent(); const fallbackElements = screen.getAllByText("-"); - expect(fallbackElements.length).toBe(2); // One for missing name, one for missing description + expect(fallbackElements.length).toBe(3); // One for missing name, one for missing description, one for missing files }); it("should wrap name in tooltip", () => { @@ -203,7 +203,7 @@ describe("VectorStoreTable", () => { it("should render fallback for missing description", () => { renderComponent(); const fallbackElements = screen.getAllByText("-"); - expect(fallbackElements.length).toBe(2); // One for missing name, one for missing description + expect(fallbackElements.length).toBe(3); // One for missing name, one for missing description, one for missing files }); it("should wrap description in tooltip", () => { @@ -386,7 +386,7 @@ describe("VectorStoreTable", () => { it("should span all columns in empty state", () => { renderComponent({ data: [] }); const emptyCell = screen.getByText("No vector stores found").closest("td"); - expect(emptyCell).toHaveAttribute("colSpan", "7"); // 6 data columns + 1 actions column + expect(emptyCell).toHaveAttribute("colSpan", "8"); // 7 data columns + 1 actions column }); }); @@ -403,7 +403,7 @@ describe("VectorStoreTable", () => { renderComponent({ data: minimalData }); expect(screen.getByText("minimal")).toBeInTheDocument(); - expect(screen.getAllByText("-")).toHaveLength(2); // Name and description fallbacks + expect(screen.getAllByText("-")).toHaveLength(3); // Name, description, and files fallbacks }); it("should handle single vector store", () => { diff --git a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.tsx b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.tsx index 52462c02e98..41e2b63112e 100644 --- a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.tsx +++ b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreTable.tsx @@ -66,6 +66,32 @@ const VectorStoreTable: React.FC = ({ data, onView, onEdi ); }, }, + { + header: "Files", + accessorKey: "vector_store_metadata", + cell: ({ row }) => { + const vectorStore = row.original; + const ingestedFiles = vectorStore.vector_store_metadata?.ingested_files || []; + + if (ingestedFiles.length === 0) { + return -; + } + + const filenames = ingestedFiles + .map((file) => file.filename || file.file_url || "Unknown") + .join(", "); + + const displayText = ingestedFiles.length === 1 + ? ingestedFiles[0].filename || ingestedFiles[0].file_url || "1 file" + : `${ingestedFiles.length} files`; + + return ( + + {displayText} + + ); + }, + }, { header: "Provider", accessorKey: "custom_llm_provider", diff --git a/ui/litellm-dashboard/src/components/vector_store_management/index.tsx b/ui/litellm-dashboard/src/components/vector_store_management/index.tsx index 6d21e861d4a..9cb57b8b9f6 100644 --- a/ui/litellm-dashboard/src/components/vector_store_management/index.tsx +++ b/ui/litellm-dashboard/src/components/vector_store_management/index.tsx @@ -1,5 +1,5 @@ import React, { useState, useEffect } from "react"; -import { Icon, Button as TremorButton, Col, Text, Grid } from "@tremor/react"; +import { Icon, Button as TremorButton, Col, Text, Grid, TabGroup, TabList, Tab, TabPanels, TabPanel } from "@tremor/react"; import { RefreshIcon } from "@heroicons/react/outline"; import { vectorStoreListCall, vectorStoreDeleteCall, credentialListCall, CredentialItem } from "../networking"; import { VectorStore } from "./types"; @@ -7,6 +7,8 @@ import VectorStoreTable from "./VectorStoreTable"; import VectorStoreForm from "./VectorStoreForm"; import DeleteResourceModal from "../common_components/DeleteResourceModal"; import VectorStoreInfoView from "./vector_store_info"; +import CreateVectorStore from "./CreateVectorStore"; +import TestVectorStoreTab from "./TestVectorStoreTab"; import { isAdminRole } from "@/utils/roles"; import NotificationsManager from "../molecules/notifications_manager"; @@ -101,6 +103,12 @@ const VectorStoreManagement: React.FC = ({ accessToken, userID fetchVectorStores(); }; + const handleVectorStoreCreated = (vectorStoreId: string) => { + console.log("Vector store created:", vectorStoreId); + fetchVectorStores(); + // Optionally switch to the manage tab + }; + useEffect(() => { fetchVectorStores(); fetchCredentials(); @@ -134,18 +142,46 @@ const VectorStoreManagement: React.FC = ({ accessToken, userID -

You can use vector stores to store and retrieve LLM embeddings..

+

You can use vector stores to store and retrieve LLM embeddings.

- setIsCreateModalVisible(true)}> - + Add Vector Store - + + + Create Vector Store + Manage Vector Stores + Test Vector Store + - -
- - - + + {/* Tab 1: Create Vector Store */} + + + + + {/* Tab 2: Manage Vector Stores */} + + setIsCreateModalVisible(true)}> + + Add Vector Store + + + + + + + + + + {/* Tab 3: Test Vector Store */} + + + + + {/* Create Vector Store Modal */} ; + vector_store_metadata?: VectorStoreMetadata; created_at: string; updated_at: string; created_by?: string; @@ -41,3 +55,32 @@ export interface VectorStoreListResponse { current_page: number; total_pages: number; } + +// Document ingestion types +export interface DocumentUpload { + uid: string; + name: string; + status: "uploading" | "done" | "error" | "removed"; + size?: number; + type?: string; + originFileObj?: File; +} + +export interface RAGIngestRequest { + file_url?: string; + file_id?: string; + ingest_options: { + vector_store: { + custom_llm_provider: string; + vector_store_id?: string; + }; + }; +} + +export interface RAGIngestResponse { + id: string; + status: "completed" | "processing" | "failed"; + vector_store_id: string; + file_id: string; + error?: string; +} diff --git a/ui/litellm-dashboard/src/components/view_users.tsx b/ui/litellm-dashboard/src/components/view_users.tsx index f8cf8302a5e..8a09c6d1f97 100644 --- a/ui/litellm-dashboard/src/components/view_users.tsx +++ b/ui/litellm-dashboard/src/components/view_users.tsx @@ -2,7 +2,7 @@ import { Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; import React, { useEffect, useState } from "react"; import { Button } from "@tremor/react"; -import BulkEditUserModal from "./bulk_edit_user"; +import BulkEditUserModal from "./BulkEditUsers"; import CreateUser from "./create_user_button"; import EditUserModal from "./edit_user"; import { @@ -286,7 +286,7 @@ const ViewUserDashboard: React.FC = ({ accessToken, toke }, handleDelete, handleResetPassword, - () => {}, // placeholder function, will be overridden in UserDataTable + () => { }, // placeholder function, will be overridden in UserDataTable ); return ( @@ -415,7 +415,7 @@ const ViewUserDashboard: React.FC = ({ accessToken, toke /> setIsBulkEditModalVisible(false)} selectedUsers={selectedUsers} possibleUIRoles={possibleUIRoles}