diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 351c4f6bc48..09b5265191b 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -12,7 +12,10 @@ WORKDIR /app USER root # Install build dependencies -RUN apk add --no-cache gcc python3-dev openssl openssl-dev +RUN apk add --no-cache \ + build-base \ + python3-dev \ + openssl-dev RUN pip install --upgrade pip && \ diff --git a/docs/my-website/docs/mcp.md b/docs/my-website/docs/mcp.md index 408560e5c0a..9a1e25a516c 100644 --- a/docs/my-website/docs/mcp.md +++ b/docs/my-website/docs/mcp.md @@ -657,7 +657,7 @@ LiteLLM Proxy provides two methods for controlling access to specific MCP server ### Method 1: URL-based Namespacing -LiteLLM Proxy supports URL-based namespacing for MCP servers using the format `/mcp/`. This allows you to: +LiteLLM Proxy supports URL-based namespacing for MCP servers using the format `//mcp`. This allows you to: - **Direct URL Access**: Point MCP clients directly to specific servers or access groups via URL - **Simplified Configuration**: Use URLs instead of headers for server selection @@ -666,14 +666,14 @@ LiteLLM Proxy supports URL-based namespacing for MCP servers using the format `/ #### URL Format ``` -/mcp/ +//mcp ``` **Examples:** -- `/mcp/github` - Access tools from the "github" MCP server -- `/mcp/zapier` - Access tools from the "zapier" MCP server -- `/mcp/dev_group` - Access tools from all servers in the "dev_group" access group -- `/mcp/github,zapier` - Access tools from multiple specific servers +- `/github_mcp/mcp` - Access tools from the "github_mcp" MCP server +- `/zapier/mcp` - Access tools from the "zapier" MCP server +- `/dev_group/mcp` - Access tools from all servers in the "dev_group" access group +- `/github_mcp,zapier/mcp` - Access tools from multiple specific servers #### Usage Examples @@ -690,7 +690,7 @@ curl --location 'https://api.openai.com/v1/responses' \ { "type": "mcp", "server_label": "litellm", - "server_url": "/mcp/github", + "server_url": "/github_mcp/mcp", "require_approval": "never", "headers": { "x-litellm-api-key": "Bearer YOUR_LITELLM_API_KEY" @@ -718,7 +718,7 @@ curl --location '/v1/responses' \ { "type": "mcp", "server_label": "litellm", - "server_url": "/mcp/dev_group", + "server_url": "/dev_group/mcp", "require_approval": "never", "headers": { "x-litellm-api-key": "Bearer YOUR_LITELLM_API_KEY" @@ -740,7 +740,7 @@ This example uses URL namespacing to access all servers in the "dev_group" acces { "mcpServers": { "LiteLLM": { - "url": "/mcp/github,zapier", + "url": "/github_mcp,zapier/mcp", "headers": { "x-litellm-api-key": "Bearer $LITELLM_API_KEY" } @@ -862,8 +862,8 @@ This configuration in Cursor IDE settings will limit tool access to only the spe | Feature | Header Namespacing | URL Namespacing | |---------|-------------------|-----------------| -| **Method** | Uses `x-mcp-servers` header | Uses URL path `/mcp/` | -| **Endpoint** | Standard `litellm_proxy` endpoint | Custom `/mcp/` endpoint | +| **Method** | Uses `x-mcp-servers` header | Uses URL path `//mcp` | +| **Endpoint** | Standard `litellm_proxy` endpoint | Custom `//mcp` endpoint | | **Configuration** | Requires additional header | Self-contained in URL | | **Multiple Servers** | Comma-separated in header | Comma-separated in URL path | | **Access Groups** | Supported via header | Supported via URL path | diff --git a/docs/my-website/docs/providers/docker_model_runner.md b/docs/my-website/docs/providers/docker_model_runner.md new file mode 100644 index 00000000000..fcd4c74f8f4 --- /dev/null +++ b/docs/my-website/docs/providers/docker_model_runner.md @@ -0,0 +1,277 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Docker Model Runner + +## Overview + +| Property | Details | +|-------|-------| +| Description | Docker Model Runner allows you to run large language models locally using Docker Desktop. | +| Provider Route on LiteLLM | `docker_model_runner/` | +| Link to Provider Doc | [Docker Model Runner ↗](https://docs.docker.com/ai/model-runner/) | +| Base URL | `http://localhost:22088` | +| Supported Operations | [`/chat/completions`](#sample-usage) | + +
+
+ +https://docs.docker.com/ai/model-runner/ + +**We support ALL Docker Model Runner models, just set `docker_model_runner/` as a prefix when sending completion requests** + +## Quick Start + +Docker Model Runner is a Docker Desktop feature that lets you run AI models locally. It provides better performance than other local solutions while maintaining OpenAI compatibility. + +### Installation + +1. Install [Docker Desktop](https://www.docker.com/products/docker-desktop/) +2. Enable Docker Model Runner in Docker Desktop settings +3. Download your preferred model through Docker Desktop + +## Environment Variables + +```python showLineNumbers title="Environment Variables" +os.environ["DOCKER_MODEL_RUNNER_API_BASE"] = "http://localhost:22088/engines/llama.cpp" # Optional - defaults to this +os.environ["DOCKER_MODEL_RUNNER_API_KEY"] = "dummy-key" # Optional - Docker Model Runner may not require auth for local instances +``` + +**Note:** +- Docker Model Runner typically runs locally and may not require authentication. LiteLLM will use a dummy key by default if no key is provided. +- The API base should include the engine path (e.g., `/engines/llama.cpp`) + +## API Base Structure + +Docker Model Runner uses a unique URL structure: + +``` +http://model-runner.docker.internal/engines/{engine}/v1/chat/completions +``` + +Where `{engine}` is the engine you want to use (typically `llama.cpp`). + +**Important:** Specify the engine in your `api_base` URL, not in the model name: +- ✅ Correct: `api_base="http://localhost:22088/engines/llama.cpp"`, `model="docker_model_runner/llama-3.1"` +- ❌ Incorrect: `api_base="http://localhost:22088"`, `model="docker_model_runner/llama.cpp/llama-3.1"` + +## Usage - LiteLLM Python SDK + +### Non-streaming + +```python showLineNumbers title="Docker Model Runner Non-streaming Completion" +import os +import litellm +from litellm import completion + +# Specify the engine in the api_base URL +os.environ["DOCKER_MODEL_RUNNER_API_BASE"] = "http://localhost:22088/engines/llama.cpp" + +messages = [{"content": "Hello, how are you?", "role": "user"}] + +# Docker Model Runner call +response = completion( + model="docker_model_runner/llama-3.1", + messages=messages +) + +print(response) +``` + +### Streaming + +```python showLineNumbers title="Docker Model Runner Streaming Completion" +import os +import litellm +from litellm import completion + +# Specify the engine in the api_base URL +os.environ["DOCKER_MODEL_RUNNER_API_BASE"] = "http://localhost:22088/engines/llama.cpp" + +messages = [{"content": "Hello, how are you?", "role": "user"}] + +# Docker Model Runner call with streaming +response = completion( + model="docker_model_runner/llama-3.1", + messages=messages, + stream=True +) + +for chunk in response: + print(chunk) +``` + +### Custom API Base and Engine + +```python showLineNumbers title="Custom API Base with Different Engine" +import litellm +from litellm import completion + +messages = [{"content": "Hello, how are you?", "role": "user"}] + +# Specify the engine in the api_base URL +# Using a different host and engine +response = completion( + model="docker_model_runner/llama-3.1", + messages=messages, + api_base="http://model-runner.docker.internal/engines/llama.cpp" +) + +print(response) +``` + +### Using Different Engines + +```python showLineNumbers title="Using a Different Engine" +import litellm +from litellm import completion + +messages = [{"content": "Hello, how are you?", "role": "user"}] + +# To use a different engine, specify it in the api_base +# For example, if Docker Model Runner supports other engines: +response = completion( + model="docker_model_runner/mistral-7b", + messages=messages, + api_base="http://localhost:22088/engines/custom-engine" +) + +print(response) +``` + +## Usage - LiteLLM Proxy + +Add the following to your LiteLLM Proxy configuration file: + +```yaml showLineNumbers title="config.yaml" +model_list: + - model_name: llama-3.1 + litellm_params: + model: docker_model_runner/llama-3.1 + api_base: http://localhost:22088/engines/llama.cpp + + - model_name: mistral-7b + litellm_params: + model: docker_model_runner/mistral-7b + api_base: http://localhost:22088/engines/llama.cpp +``` + +Start your LiteLLM Proxy server: + +```bash showLineNumbers title="Start LiteLLM Proxy" +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + + + + +```python showLineNumbers title="Docker Model Runner via Proxy - Non-streaming" +from openai import OpenAI + +# Initialize client with your proxy URL +client = OpenAI( + base_url="http://localhost:4000", # Your proxy URL + api_key="your-proxy-api-key" # Your proxy API key +) + +# Non-streaming response +response = client.chat.completions.create( + model="llama-3.1", + messages=[{"role": "user", "content": "hello from litellm"}] +) + +print(response.choices[0].message.content) +``` + +```python showLineNumbers title="Docker Model Runner via Proxy - Streaming" +from openai import OpenAI + +# Initialize client with your proxy URL +client = OpenAI( + base_url="http://localhost:4000", # Your proxy URL + api_key="your-proxy-api-key" # Your proxy API key +) + +# Streaming response +response = client.chat.completions.create( + model="llama-3.1", + messages=[{"role": "user", "content": "hello from litellm"}], + stream=True +) + +for chunk in response: + if chunk.choices[0].delta.content is not None: + print(chunk.choices[0].delta.content, end="") +``` + + + + + +```python showLineNumbers title="Docker Model Runner via Proxy - LiteLLM SDK" +import litellm + +# Configure LiteLLM to use your proxy +response = litellm.completion( + model="litellm_proxy/llama-3.1", + messages=[{"role": "user", "content": "hello from litellm"}], + api_base="http://localhost:4000", + api_key="your-proxy-api-key" +) + +print(response.choices[0].message.content) +``` + +```python showLineNumbers title="Docker Model Runner via Proxy - LiteLLM SDK Streaming" +import litellm + +# Configure LiteLLM to use your proxy with streaming +response = litellm.completion( + model="litellm_proxy/llama-3.1", + messages=[{"role": "user", "content": "hello from litellm"}], + api_base="http://localhost:4000", + api_key="your-proxy-api-key", + stream=True +) + +for chunk in response: + if hasattr(chunk.choices[0], 'delta') and chunk.choices[0].delta.content is not None: + print(chunk.choices[0].delta.content, end="") +``` + + + + + +```bash showLineNumbers title="Docker Model Runner via Proxy - cURL" +curl http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer your-proxy-api-key" \ + -d '{ + "model": "llama-3.1", + "messages": [{"role": "user", "content": "hello from litellm"}] + }' +``` + +```bash showLineNumbers title="Docker Model Runner via Proxy - cURL Streaming" +curl http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer your-proxy-api-key" \ + -d '{ + "model": "llama-3.1", + "messages": [{"role": "user", "content": "hello from litellm"}], + "stream": true + }' +``` + + + + +For more detailed information on using the LiteLLM Proxy, see the [LiteLLM Proxy documentation](../providers/litellm_proxy). + +## API Reference + +For detailed API information, see the [Docker Model Runner API Reference](https://docs.docker.com/ai/model-runner/api-reference/). + diff --git a/docs/my-website/docs/providers/gemini.md b/docs/my-website/docs/providers/gemini.md index fd20e907d3b..e04225e1f85 100644 --- a/docs/my-website/docs/providers/gemini.md +++ b/docs/my-website/docs/providers/gemini.md @@ -1308,6 +1308,8 @@ curl --location 'http://localhost:4000/v1/chat/completions' \ 5. **Format**: Thought signatures are stored in `provider_specific_fields.thought_signature` of tool calls in the response, and are automatically included when you append the assistant message to your conversation history. +6. **Chat Completions Clients**: With chat completions clients where you cannot control whether or not the previous assistant message is included as-is (ex langchain's ChatOpenAI), LiteLLM also preserves the thought signature by appending it to the tool call id (`call_123__thought__`) and extracting it back out before sending the outbound request to Gemini. + ## JSON Mode diff --git a/docs/my-website/docs/providers/xai.md b/docs/my-website/docs/providers/xai.md index 49a3640991d..afeecc21528 100644 --- a/docs/my-website/docs/providers/xai.md +++ b/docs/my-website/docs/providers/xai.md @@ -11,6 +11,68 @@ https://docs.x.ai/docs ::: +## Supported Models + + + +**Latest Release** - Grok 4.1 Fast: Optimized for high-performance agentic tool calling with 2M context and prompt caching. + +| Model | Context | Features | +|-------|---------|----------| +| `xai/grok-4-1-fast-reasoning` | 2M tokens | **Reasoning**, Function calling, Vision, Audio, Web search, Caching | +| `xai/grok-4-1-fast-non-reasoning` | 2M tokens | Function calling, Vision, Audio, Web search, Caching | + +**When to use:** +- ✅ **Reasoning model**: Complex analysis, planning, multi-step reasoning problems +- ✅ **Non-reasoning model**: Simple queries, faster responses, lower token usage + +**Example:** +```python +from litellm import completion + +# With reasoning +response = completion( + model="xai/grok-4-1-fast-reasoning", + messages=[{"role": "user", "content": "Analyze this problem step by step..."}] +) + +# Without reasoning +response = completion( + model="xai/grok-4-1-fast-non-reasoning", + messages=[{"role": "user", "content": "What's 2+2?"}] +) +``` + +--- + +### All Available Models + +| Model Family | Model | Context | Features | +|--------------|-------|---------|----------| +| **Grok 4.1** | `xai/grok-4-1-fast-reasoning` | 2M | **Reasoning**, Tools, Vision, Audio, Web search, Caching | +| | `xai/grok-4-1-fast-non-reasoning` | 2M | Tools, Vision, Audio, Web search, Caching | +| **Grok 4** | `xai/grok-4` | 256K | Tools, Web search | +| | `xai/grok-4-0709` | 256K | Tools, Web search | +| | `xai/grok-4-fast-reasoning` | 2M | **Reasoning**, Tools, Web search | +| | `xai/grok-4-fast-non-reasoning` | 2M | Tools, Web search | +| **Grok 3** | `xai/grok-3` | 131K | Tools, Web search | +| | `xai/grok-3-mini` | 131K | Tools, Web search | +| | `xai/grok-3-fast-beta` | 131K | Tools, Web search | +| **Grok Code** | `xai/grok-code-fast` | 256K | **Reasoning**, Tools, Code generation, Caching | +| **Grok 2** | `xai/grok-2` | 131K | Tools, **Vision** | +| | `xai/grok-2-vision-latest` | 32K | Tools, **Vision** | + +**Features:** +- **Reasoning** = Chain-of-thought reasoning with reasoning tokens +- **Tools** = Function calling / Tool use +- **Web search** = Live internet search +- **Vision** = Image understanding +- **Audio** = Audio input support +- **Caching** = Prompt caching for cost savings +- **Code generation** = Optimized for code tasks + +**Pricing:** See [xAI's pricing page](https://docs.x.ai/docs/models) for current rates. + ## API Key ```python # env variable diff --git a/docs/my-website/docs/proxy/guardrails/tool_permission.md b/docs/my-website/docs/proxy/guardrails/tool_permission.md index 9ed05ed46a8..22ecdd2251e 100644 --- a/docs/my-website/docs/proxy/guardrails/tool_permission.md +++ b/docs/my-website/docs/proxy/guardrails/tool_permission.md @@ -46,6 +46,43 @@ guardrails: - `pre_call` Run **before** LLM call, on **input** - `post_call` Run **after** LLM call, on **input & output** +### `on_disallowed_action` behavior + +| Value | What happens | +| --- | --- | +| `block` | The request is immediately rejected. Pre-call checks raise a `400` HTTP error. Post-call checks raise `GuardrailRaisedException`, so the proxy responds with an error instead of the model output. Use when invoking the forbidden tool must halt the workflow. | +| `rewrite` | LiteLLM silently strips disallowed tools from the payload before it reaches the model (pre-call) or rewrites the model response/tool calls after the fact. The guardrail inserts error text into `message.content`/`tool_result` entries so the client learns the tool was blocked while the rest of the completion continues. Use when you want graceful degradation instead of hard failures. | + +### Custom denial message + +Set `violation_message_template` when you want the guardrail to return a branded error (e.g., “this violates our org policy…”). LiteLLM replaces placeholders from the denied tool: + +- `{tool_name}` – the tool/function name (e.g., `Read`) +- `{rule_id}` – the matching rule ID (or `None` when the default action kicks in) +- `{default_message}` – the original LiteLLM message if you need to append it + +Example: + +```yaml +guardrails: + - guardrail_name: "tool-permission-guardrail" + litellm_params: + guardrail: tool_permission + mode: "post_call" + violation_message_template: "this violates our org policy, we don't support executing {tool_name} commands" + rules: + - id: "allow_bash" + tool_name: "Bash" + decision: "allow" + - id: "deny_read" + tool_name: "Read" + decision: "deny" + default_action: "deny" + on_disallowed_action: "block" +``` + +If a request tries to invoke `Read`, the proxy now returns “this violates our org policy, we don't support executing Read commands” instead of the stock error text. Omit the field to keep the default messaging. + ### 2. Start the Proxy ```shell @@ -57,7 +94,7 @@ litellm --config config.yaml --port 4000 -**Block requset** +**Block request (`on_disallowed_action: block`)** ```bash # Test @@ -96,7 +133,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \ -**Rewrite requset** +**Rewrite request (`on_disallowed_action: rewrite`)** ```bash # Test @@ -118,7 +155,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \ }' ``` -**Expected response:** +**Expected response (tool removed, completion continues):** ```json { diff --git a/docs/my-website/docs/tutorials/claude_responses_api.md b/docs/my-website/docs/tutorials/claude_responses_api.md index 0dbb4a2f1e7..aafeccceaf5 100644 --- a/docs/my-website/docs/tutorials/claude_responses_api.md +++ b/docs/my-website/docs/tutorials/claude_responses_api.md @@ -105,7 +105,7 @@ LITELLM_MASTER_KEY gives claude access to all proxy models, whereas a virtual ke Alternatively, use the Anthropic pass-through endpoint: ```bash -export ANTHROPIC_BASE_URL="http://0.0.0.0:4000" +export ANTHROPIC_BASE_URL="http://0.0.0.0:4000/anthropic" export ANTHROPIC_AUTH_TOKEN="$LITELLM_MASTER_KEY" ``` @@ -221,7 +221,6 @@ You can also connect MCP servers to Claude Code via LiteLLM Proxy. Limitations: - Currently, only HTTP MCP servers are supported -- Does not work in Cursor IDE yet. ::: diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 432d2d109eb..376bcfd5357 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -530,13 +530,39 @@ const sidebars = { "providers/bedrock_vector_store", ] }, - "providers/milvus_vector_stores", "providers/litellm_proxy", - "providers/meta_llama", - "providers/mistral", + "providers/ai21", + "providers/aiml", + "providers/aleph_alpha", + "providers/anyscale", + "providers/baseten", + "providers/bytez", + "providers/cerebras", + "providers/clarifai", + "providers/cloudflare_workers", "providers/codestral", "providers/cohere", - "providers/anyscale", + "providers/cometapi", + "providers/compactifai", + "providers/custom_llm_server", + "providers/dashscope", + "providers/databricks", + "providers/datarobot", + "providers/deepgram", + "providers/deepinfra", + "providers/deepseek", + "providers/docker_model_runner", + "providers/elevenlabs", + "providers/fal_ai", + "providers/featherless_ai", + "providers/fireworks_ai", + "providers/friendliai", + "providers/galadriel", + "providers/github", + "providers/github_copilot", + "providers/gradient_ai", + "providers/groq", + "providers/heroku", { type: "category", label: "HuggingFace", @@ -546,10 +572,21 @@ const sidebars = { ] }, "providers/hyperbolic", - "providers/databricks", - "providers/deepgram", - "providers/watsonx", - "providers/predibase", + "providers/infinity", + "providers/jina_ai", + "providers/lambda_ai", + "providers/lemonade", + "providers/llamafile", + "providers/lm_studio", + "providers/meta_llama", + "providers/milvus_vector_stores", + "providers/mistral", + "providers/moonshot", + "providers/morph", + "providers/nebius", + "providers/nlp_cloud", + "providers/novita", + { type: "doc", id: "providers/nscale", label: "Nscale (EU Sovereign)" }, { type: "category", label: "Nvidia NIM", @@ -558,37 +595,13 @@ const sidebars = { "providers/nvidia_nim_rerank", ] }, - { type: "doc", id: "providers/nscale", label: "Nscale (EU Sovereign)" }, - "providers/xai", - "providers/moonshot", - "providers/lm_studio", - "providers/cerebras", - "providers/volcano", - "providers/triton-inference-server", + "providers/oci", "providers/ollama", + "providers/openrouter", + "providers/ovhcloud", "providers/perplexity", - "providers/friendliai", - "providers/galadriel", - "providers/topaz", - "providers/groq", - "providers/deepseek", - "providers/elevenlabs", - "providers/fal_ai", - "providers/fireworks_ai", - "providers/clarifai", - "providers/compactifai", - "providers/lemonade", - "providers/vllm", - "providers/llamafile", - "providers/infinity", - "providers/xinference", - "providers/aiml", - "providers/cloudflare_workers", - "providers/deepinfra", - "providers/github", - "providers/github_copilot", - "providers/ai21", - "providers/nlp_cloud", + "providers/petals", + "providers/predibase", "providers/recraft", "providers/replicate", { @@ -599,32 +612,20 @@ const sidebars = { "providers/runwayml/videos", ] }, + "providers/sambanova", + "providers/snowflake", "providers/togetherai", + "providers/topaz", + "providers/triton-inference-server", "providers/v0", "providers/vercel_ai_gateway", - "providers/morph", - "providers/lambda_ai", - "providers/novita", + "providers/vllm", + "providers/volcano", "providers/voyage", - "providers/jina_ai", - "providers/aleph_alpha", - "providers/baseten", - "providers/openrouter", - "providers/sambanova", - "providers/custom_llm_server", - "providers/petals", - "providers/snowflake", - "providers/gradient_ai", - "providers/featherless_ai", - "providers/nebius", - "providers/dashscope", - "providers/bytez", - "providers/heroku", - "providers/oci", - "providers/datarobot", - "providers/ovhcloud", "providers/wandb_inference", - "providers/cometapi", + "providers/watsonx", + "providers/xai", + "providers/xinference", ], }, { diff --git a/litellm/__init__.py b/litellm/__init__.py index b46a165ed10..51be5ee2e29 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -563,6 +563,7 @@ wandb_models: Set = set(WANDB_MODELS) ovhcloud_models: Set = set() ovhcloud_embedding_models: Set = set() lemonade_models: Set = set() +docker_model_runner_models: Set = set() def is_bedrock_pricing_only_model(key: str) -> bool: @@ -797,6 +798,8 @@ def add_known_models(): ovhcloud_embedding_models.add(key) elif value.get("litellm_provider") == "lemonade": lemonade_models.add(key) + elif value.get("litellm_provider") == "docker_model_runner": + docker_model_runner_models.add(key) add_known_models() @@ -900,6 +903,7 @@ model_list = list( | wandb_models | ovhcloud_models | lemonade_models + | docker_model_runner_models | set(clarifai_models) ) @@ -1350,6 +1354,7 @@ from .llms.nebius.chat.transformation import NebiusConfig from .llms.wandb.chat.transformation import WandbConfig from .llms.dashscope.chat.transformation import DashScopeChatConfig from .llms.moonshot.chat.transformation import MoonshotChatConfig +from .llms.docker_model_runner.chat.transformation import DockerModelRunnerChatConfig from .llms.v0.chat.transformation import V0ChatConfig from .llms.oci.chat.transformation import OCIChatConfig from .llms.morph.chat.transformation import MorphChatConfig diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 5279dd70bc4..838ee95b2b5 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -18,7 +18,6 @@ from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast import httpx from openai.types.batch import BatchRequestCounts -from openai.types.batch import Metadata as BatchMetadata import litellm from litellm._logging import verbose_logger diff --git a/litellm/constants.py b/litellm/constants.py index bc72e93850b..b312a15892b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -381,6 +381,7 @@ LITELLM_CHAT_PROVIDERS = [ "wandb", "ovhcloud", "lemonade", + "docker_model_runner", ] LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS = [ @@ -567,6 +568,7 @@ openai_compatible_providers: List = [ "wandb", "cometapi", "clarifai", + "docker_model_runner", ] openai_text_completion_compatible_providers: List = ( [ # providers that support `/v1/completions` diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index b50d05ed2ec..b52f1b3095e 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -36,6 +36,7 @@ class CustomGuardrail(CustomLogger): default_on: bool = False, mask_request_content: bool = False, mask_response_content: bool = False, + violation_message_template: Optional[str] = None, **kwargs, ): """ @@ -57,12 +58,34 @@ class CustomGuardrail(CustomLogger): self.default_on: bool = default_on self.mask_request_content: bool = mask_request_content self.mask_response_content: bool = mask_response_content + self.violation_message_template: Optional[str] = violation_message_template if supported_event_hooks: ## validate event_hook is in supported_event_hooks self._validate_event_hook(event_hook, supported_event_hooks) super().__init__(**kwargs) + def render_violation_message( + self, default: str, context: Optional[Dict[str, Any]] = None + ) -> str: + """Return a custom violation message if template is configured.""" + + if not self.violation_message_template: + return default + + format_context: Dict[str, Any] = {"default_message": default} + if context: + format_context.update(context) + try: + return self.violation_message_template.format(**format_context) + except Exception as e: + verbose_logger.warning( + "Failed to format violation message template for guardrail %s: %s", + self.guardrail_name, + e, + ) + return default + @staticmethod def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: """ @@ -279,7 +302,7 @@ class CustomGuardrail(CustomLogger): data, self.event_hook ) if result is not None: - return result + return result return True def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool: diff --git a/litellm/integrations/dotprompt/__init__.py b/litellm/integrations/dotprompt/__init__.py index 3af7fbf6dd3..3847c8fa192 100644 --- a/litellm/integrations/dotprompt/__init__.py +++ b/litellm/integrations/dotprompt/__init__.py @@ -25,6 +25,23 @@ def set_global_prompt_directory(directory: str) -> None: litellm.global_prompt_directory = directory # type: ignore +def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict: + """ + Get the prompt data from the dotprompt content. + + The UI stores prompts under `dotprompt_content` in the database. This function parses the content and returns the prompt data in the format expected by the prompt manager. + """ + from .prompt_manager import PromptManager + + # Parse the dotprompt content to extract frontmatter and content + temp_manager = PromptManager() + metadata, content = temp_manager._parse_frontmatter(dotprompt_content) + + # Convert to prompt_data format + return { + "content": content.strip(), + "metadata": metadata + } def prompt_initializer( litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec" @@ -41,6 +58,11 @@ def prompt_initializer( ) prompt_file = getattr(litellm_params, "prompt_file", None) + + # Handle dotprompt_content from database + dotprompt_content = getattr(litellm_params, "dotprompt_content", None) + if dotprompt_content and not prompt_data and not prompt_file: + prompt_data = _get_prompt_data_from_dotprompt_content(dotprompt_content) try: dot_prompt_manager = DotpromptManager( diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index ef0ebe074d7..eefe680217d 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -741,6 +741,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915 ) = litellm.MoonshotChatConfig()._get_openai_compatible_provider_info( api_base, api_key ) + elif custom_llm_provider == "docker_model_runner": + ( + api_base, + dynamic_api_key, + ) = litellm.DockerModelRunnerChatConfig()._get_openai_compatible_provider_info( + api_base, api_key + ) elif custom_llm_provider == "v0": ( api_base, diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 44987bd1d08..0f4a159975a 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -58,6 +58,10 @@ def prompt_injection_detection_default_pt(): BAD_MESSAGE_ERROR_STR = "Invalid Message " +# Separator used to embed Gemini thought signatures in tool call IDs +# See: https://ai.google.dev/gemini-api/docs/thought-signatures +THOUGHT_SIGNATURE_SEPARATOR = "__thought__" + # used to interweave user messages, to ensure user/assistant alternating DEFAULT_USER_CONTINUE_MESSAGE = { "role": "user", @@ -1162,9 +1166,34 @@ def _gemini_tool_call_invoke_helper( return function_call -def _get_thought_signature_from_tool(tool: dict, model: Optional[str] = None) -> Optional[str]: +def _encode_tool_call_id_with_signature( + tool_call_id: str, thought_signature: Optional[str] +) -> str: + """ + Embed thought signature into tool call ID for OpenAI client compatibility. + + Args: + tool_call_id: The tool call ID (e.g., "call_abc123...") + thought_signature: Base64-encoded signature from Gemini response + + Returns: + Tool call ID with embedded signature if present, otherwise original ID + Format: call___thought__ + + See: https://ai.google.dev/gemini-api/docs/thought-signatures + """ + if thought_signature: + return f"{tool_call_id}{THOUGHT_SIGNATURE_SEPARATOR}{thought_signature}" + return tool_call_id + + +def _get_thought_signature_from_tool( + tool: dict, model: Optional[str] = None +) -> Optional[str]: """Extract thought signature from tool call's provider_specific_fields. - + + If not provided try to extract thought signature from tool call id + Checks both tool.provider_specific_fields and tool.function.provider_specific_fields. If no signature is found and model is gemini-3, returns a dummy signature. """ @@ -1174,7 +1203,7 @@ def _get_thought_signature_from_tool(tool: dict, model: Optional[str] = None) -> signature = provider_fields.get("thought_signature") if signature: return signature - + # Then check function's provider_specific_fields function = tool.get("function") if function: @@ -1184,23 +1213,34 @@ def _get_thought_signature_from_tool(tool: dict, model: Optional[str] = None) -> signature = func_provider_fields.get("thought_signature") if signature: return signature - elif hasattr(function, "provider_specific_fields") and function.provider_specific_fields: + elif ( + hasattr(function, "provider_specific_fields") + and function.provider_specific_fields + ): if isinstance(function.provider_specific_fields, dict): signature = function.provider_specific_fields.get("thought_signature") if signature: return signature - + # Check if thought signature is embedded in tool call ID + tool_call_id = tool.get("id") + if tool_call_id and THOUGHT_SIGNATURE_SEPARATOR in tool_call_id: + parts = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1) + if len(parts) == 2: + _, signature = parts + return signature # If no signature found and model is gemini-3, return dummy signature - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + if model and VertexGeminiConfig._is_gemini_3_or_newer(model): return _get_dummy_thought_signature() - return None def _get_dummy_thought_signature() -> str: """Generate a dummy thought signature for models that require it. - + This is used when transferring conversation history from older models (like gemini-2.5-flash) to gemini-3, which requires thought_signature for strict validation. @@ -1258,23 +1298,25 @@ def convert_to_gemini_tool_call_invoke( _parts_list: List[VertexPartType] = [] tool_calls = message.get("tool_calls", None) function_call = message.get("function_call", None) - + if tool_calls is not None: for idx, tool in enumerate(tool_calls): if "function" in tool: - gemini_function_call: Optional[VertexFunctionCall] = ( - _gemini_tool_call_invoke_helper( - function_call_params=tool["function"] - ) + gemini_function_call: Optional[ + VertexFunctionCall + ] = _gemini_tool_call_invoke_helper( + function_call_params=tool["function"] ) if gemini_function_call is not None: part_dict: VertexPartType = { "function_call": gemini_function_call } - thought_signature = _get_thought_signature_from_tool(dict(tool), model=model) + thought_signature = _get_thought_signature_from_tool( + dict(tool), model=model + ) if thought_signature: part_dict["thoughtSignature"] = thought_signature - + _parts_list.append(part_dict) else: # don't silently drop params. Make it clear to user what's happening. raise Exception( @@ -1290,21 +1332,32 @@ def convert_to_gemini_tool_call_invoke( part_dict_function: VertexPartType = { "function_call": gemini_function_call } - + # Extract thought signature from function_call's provider_specific_fields thought_signature = None - provider_fields = function_call.get("provider_specific_fields") if isinstance(function_call, dict) else {} + provider_fields = ( + function_call.get("provider_specific_fields") + if isinstance(function_call, dict) + else {} + ) if isinstance(provider_fields, dict): thought_signature = provider_fields.get("thought_signature") - + # If no signature found and model is gemini-3, use dummy signature - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig - if not thought_signature and model and VertexGeminiConfig._is_gemini_3_or_newer(model): + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + if ( + not thought_signature + and model + and VertexGeminiConfig._is_gemini_3_or_newer(model) + ): thought_signature = _get_dummy_thought_signature() - + if thought_signature: part_dict_function["thoughtSignature"] = thought_signature - + _parts_list.append(part_dict_function) else: # don't silently drop params. Make it clear to user what's happening. raise Exception( @@ -1807,9 +1860,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_element["cache_control"] = ( - _content_element["cache_control"] - ) + _anthropic_content_element[ + "cache_control" + ] = _content_element["cache_control"] user_content.append(_anthropic_content_element) elif m.get("type", "") == "text": m = cast(ChatCompletionTextObject, m) @@ -1847,9 +1900,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_content_text_element["cache_control"] = ( - _content_element["cache_control"] - ) + _anthropic_content_text_element[ + "cache_control" + ] = _content_element["cache_control"] user_content.append(_anthropic_content_text_element) @@ -2615,17 +2668,19 @@ class BedrockImageProcessor: """Handles both sync and async image processing for Bedrock conversations.""" @staticmethod - def _post_call_image_processing(response: httpx.Response, image_url: str = "") -> Tuple[str, str]: + def _post_call_image_processing( + response: httpx.Response, image_url: str = "" + ) -> Tuple[str, str]: # Check the response's content type to ensure it is an image content_type = response.headers.get("content-type") - + # Use helper function to infer content type with fallback logic content_type = infer_content_type_from_url_and_content( url=image_url, content=response.content, current_content_type=content_type, ) - + content_type = _parse_content_type(content_type) # Convert the image content to base64 bytes @@ -2644,7 +2699,9 @@ class BedrockImageProcessor: response = await client.get(image_url, follow_redirects=True) response.raise_for_status() # Raise an exception for HTTP errors - return BedrockImageProcessor._post_call_image_processing(response, image_url) + return BedrockImageProcessor._post_call_image_processing( + response, image_url + ) except Exception as e: raise e @@ -2657,7 +2714,9 @@ class BedrockImageProcessor: response = client.get(image_url, follow_redirects=True) response.raise_for_status() # Raise an exception for HTTP errors - return BedrockImageProcessor._post_call_image_processing(response, image_url) + return BedrockImageProcessor._post_call_image_processing( + response, image_url + ) except Exception as e: raise e @@ -2988,21 +3047,33 @@ def _convert_to_bedrock_tool_call_result( """ - """ - content_str: str = "" + tool_result_content_blocks:List[BedrockToolResultContentBlock] = [] if isinstance(message["content"], str): - content_str = message["content"] + tool_result_content_blocks.append(BedrockToolResultContentBlock(text=message["content"])) elif isinstance(message["content"], List): content_list = message["content"] for content in content_list: if content["type"] == "text": - content_str += content["text"] + tool_result_content_blocks.append(BedrockToolResultContentBlock(text=content["text"])) + elif content["type"] == "image_url": + format: Optional[str] = None + if isinstance(content["image_url"], dict): + image_url = content["image_url"]["url"] + format = content["image_url"].get("format") + else: + image_url = content["image_url"] + _block:BedrockContentBlock = BedrockImageProcessor.process_image_sync( + image_url=image_url, + format=format, + ) + if "image" in _block: + tool_result_content_blocks.append(BedrockToolResultContentBlock(image=_block["image"])) message.get("name", "") id = str(message.get("tool_call_id", str(uuid.uuid4()))) - tool_result_content_block = BedrockToolResultContentBlock(text=content_str) tool_result = BedrockToolResultBlock( - content=[tool_result_content_block], + content=tool_result_content_blocks, toolUseId=id, ) @@ -3914,7 +3985,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 ) elif element["type"] == "text": # AWS Bedrock doesn't allow empty or whitespace-only text content, so use placeholder for empty strings - text_content = element["text"] if element["text"].strip() else "." + text_content = ( + element["text"] if element["text"].strip() else "." + ) assistants_part = BedrockContentBlock(text=text_content) assistants_parts.append(assistants_part) elif element["type"] == "image_url": diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 71429e4191f..b35e86cabd2 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -51,7 +51,11 @@ from litellm.types.llms.openai import ( ChatCompletionToolCallFunctionChunk, ChatCompletionUsageBlock, ) -from litellm.types.utils import ChatCompletionMessageToolCall, Choices, Delta +from litellm.types.utils import ( + ChatCompletionMessageToolCall, + Choices, + Delta, +) from litellm.types.utils import GenericStreamingChunk as GChunk from litellm.types.utils import ( ModelResponse, @@ -1246,18 +1250,168 @@ class AWSEventStreamDecoder: thinking_blocks_list.append(_thinking_block) return thinking_blocks_list + def _initialize_converse_response_id(self, chunk_data: dict): + """Initialize response_id from chunk data if not already set.""" + if self.response_id is None: + if "messageStart" in chunk_data: + conversation_id = chunk_data["messageStart"].get("conversationId") + if conversation_id: + self.response_id = f"chatcmpl-{conversation_id}" + else: + # Fallback to generating a UUID if the first chunk is not messageStart + self.response_id = f"chatcmpl-{uuid.uuid4()}" + + def _handle_converse_start_event( + self, + start_obj: ContentBlockStartEvent, + ) -> Tuple[ + Optional[ChatCompletionToolCallChunk], + dict, + Optional[ + List[ + Union[ + ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock + ] + ] + ], + ]: + """Handle 'start' event in converse chunk parsing.""" + tool_use: Optional[ChatCompletionToolCallChunk] = None + provider_specific_fields: dict = {} + thinking_blocks: Optional[ + List[ + Union[ + ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock + ] + ] + ] = None + + self.content_blocks = [] # reset + if start_obj is not None: + if "toolUse" in start_obj and start_obj["toolUse"] is not None: + ## check tool name was formatted by litellm + _response_tool_name = start_obj["toolUse"]["name"] + response_tool_name = get_bedrock_tool_name( + response_tool_name=_response_tool_name + ) + self.tool_calls_index = ( + 0 + if self.tool_calls_index is None + else self.tool_calls_index + 1 + ) + tool_use = { + "id": start_obj["toolUse"]["toolUseId"], + "type": "function", + "function": { + "name": response_tool_name, + "arguments": "", + }, + "index": self.tool_calls_index, + } + elif ( + "reasoningContent" in start_obj + and start_obj["reasoningContent"] is not None + ): # redacted thinking can be in start object + thinking_blocks = self.translate_thinking_blocks( + start_obj["reasoningContent"] + ) + provider_specific_fields = { + "reasoningContent": start_obj["reasoningContent"], + } + return tool_use, provider_specific_fields, thinking_blocks + + def _handle_converse_delta_event( + self, + delta_obj: ContentBlockDeltaEvent, + index: int, + ) -> Tuple[ + str, + Optional[ChatCompletionToolCallChunk], + dict, + Optional[str], + Optional[ + List[ + Union[ + ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock + ] + ] + ], + ]: + """Handle 'delta' event in converse chunk parsing.""" + text = "" + tool_use: Optional[ChatCompletionToolCallChunk] = None + provider_specific_fields: dict = {} + reasoning_content: Optional[str] = None + thinking_blocks: Optional[ + List[ + Union[ + ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock + ] + ] + ] = None + + self.content_blocks.append(delta_obj) + if "text" in delta_obj: + text = delta_obj["text"] + elif "toolUse" in delta_obj: + tool_use = { + "id": None, + "type": "function", + "function": { + "name": None, + "arguments": delta_obj["toolUse"]["input"], + }, + "index": ( + self.tool_calls_index + if self.tool_calls_index is not None + else index + ), + } + elif "reasoningContent" in delta_obj: + provider_specific_fields = { + "reasoningContent": delta_obj["reasoningContent"], + } + reasoning_content = self.extract_reasoning_content_str( + delta_obj["reasoningContent"] + ) + thinking_blocks = self.translate_thinking_blocks( + delta_obj["reasoningContent"] + ) + if ( + thinking_blocks + and len(thinking_blocks) > 0 + and reasoning_content is None + ): + reasoning_content = "" # set to non-empty string to ensure consistency with Anthropic + return text, tool_use, provider_specific_fields, reasoning_content, thinking_blocks + + def _handle_converse_stop_event( + self, index: int + ) -> Optional[ChatCompletionToolCallChunk]: + """Handle stop/contentBlockIndex event in converse chunk parsing.""" + tool_use: Optional[ChatCompletionToolCallChunk] = None + is_empty = self.check_empty_tool_call_args() + if is_empty: + tool_use = { + "id": None, + "type": "function", + "function": { + "name": None, + "arguments": "{}", + }, + "index": ( + self.tool_calls_index + if self.tool_calls_index is not None + else index + ), + } + return tool_use + def converse_chunk_parser(self, chunk_data: dict) -> ModelResponseStream: try: # Capture the conversationId from the first messageStart event # and use it as the consistent ID for all subsequent chunks. - if self.response_id is None: - if "messageStart" in chunk_data: - conversation_id = chunk_data["messageStart"].get("conversationId") - if conversation_id: - self.response_id = f"chatcmpl-{conversation_id}" - else: - # Fallback to generating a UUID if the first chunk is not messageStart - self.response_id = f"chatcmpl-{uuid.uuid4()}" + self._initialize_converse_response_id(chunk_data) verbose_logger.debug("\n\nRaw Chunk: {}\n\n".format(chunk_data)) text = "" @@ -1277,91 +1431,22 @@ class AWSEventStreamDecoder: index = int(chunk_data.get("contentBlockIndex", 0)) if "start" in chunk_data: start_obj = ContentBlockStartEvent(**chunk_data["start"]) - self.content_blocks = [] # reset - if start_obj is not None: - if "toolUse" in start_obj and start_obj["toolUse"] is not None: - ## check tool name was formatted by litellm - _response_tool_name = start_obj["toolUse"]["name"] - response_tool_name = get_bedrock_tool_name( - response_tool_name=_response_tool_name - ) - self.tool_calls_index = ( - 0 - if self.tool_calls_index is None - else self.tool_calls_index + 1 - ) - tool_use = { - "id": start_obj["toolUse"]["toolUseId"], - "type": "function", - "function": { - "name": response_tool_name, - "arguments": "", - }, - "index": self.tool_calls_index, - } - elif ( - "reasoningContent" in start_obj - and start_obj["reasoningContent"] is not None - ): # redacted thinking can be in start object - thinking_blocks = self.translate_thinking_blocks( - start_obj["reasoningContent"] - ) - provider_specific_fields = { - "reasoningContent": start_obj["reasoningContent"], - } + tool_use, provider_specific_fields, thinking_blocks = ( + self._handle_converse_start_event(start_obj) + ) elif "delta" in chunk_data: delta_obj = ContentBlockDeltaEvent(**chunk_data["delta"]) - self.content_blocks.append(delta_obj) - if "text" in delta_obj: - text = delta_obj["text"] - elif "toolUse" in delta_obj: - tool_use = { - "id": None, - "type": "function", - "function": { - "name": None, - "arguments": delta_obj["toolUse"]["input"], - }, - "index": ( - self.tool_calls_index - if self.tool_calls_index is not None - else index - ), - } - elif "reasoningContent" in delta_obj: - provider_specific_fields = { - "reasoningContent": delta_obj["reasoningContent"], - } - reasoning_content = self.extract_reasoning_content_str( - delta_obj["reasoningContent"] - ) - thinking_blocks = self.translate_thinking_blocks( - delta_obj["reasoningContent"] - ) - if ( - thinking_blocks - and len(thinking_blocks) > 0 - and reasoning_content is None - ): - reasoning_content = "" # set to non-empty string to ensure consistency with Anthropic + ( + text, + tool_use, + provider_specific_fields, + reasoning_content, + thinking_blocks, + ) = self._handle_converse_delta_event(delta_obj, index) elif ( "contentBlockIndex" in chunk_data ): # stop block, no 'start' or 'delta' object - is_empty = self.check_empty_tool_call_args() - if is_empty: - tool_use = { - "id": None, - "type": "function", - "function": { - "name": None, - "arguments": "{}", - }, - "index": ( - self.tool_calls_index - if self.tool_calls_index is not None - else index - ), - } + tool_use = self._handle_converse_stop_event(index) elif "stopReason" in chunk_data: finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop")) elif "usage" in chunk_data: diff --git a/litellm/llms/docker_model_runner/chat/transformation.py b/litellm/llms/docker_model_runner/chat/transformation.py new file mode 100644 index 00000000000..3d84b24a01c --- /dev/null +++ b/litellm/llms/docker_model_runner/chat/transformation.py @@ -0,0 +1,144 @@ +""" +Translates from OpenAI's `/v1/chat/completions` to Docker Model Runner's `/engines/{engine}/v1/chat/completions` + +Docker Model Runner API Reference: https://docs.docker.com/ai/model-runner/api-reference/ +""" + +from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload + +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + handle_messages_with_content_list_to_str_conversion, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllMessageValues + +from ...openai.chat.gpt_transformation import OpenAIGPTConfig + + +class DockerModelRunnerChatConfig(OpenAIGPTConfig): + """ + Configuration for Docker Model Runner API. + + Docker Model Runner uses URLs in the format: /engines/{engine}/v1/chat/completions + The engine name (e.g., "llama.cpp") is part of the API endpoint path. + """ + + @overload + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, + ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: + """ + Docker Model Runner is OpenAI-compatible, so we use standard message transformation. + """ + messages = handle_messages_with_content_list_to_str_conversion(messages) + if is_async: + return super()._transform_messages( + messages=messages, model=model, is_async=True + ) + else: + return super()._transform_messages( + messages=messages, model=model, is_async=False + ) + + def _get_openai_compatible_provider_info( + self, api_base: Optional[str], api_key: Optional[str] + ) -> Tuple[Optional[str], Optional[str]]: + """ + Get API base and key for Docker Model Runner. + + Default API base: http://localhost:22088/engines/llama.cpp + The engine path should be included in the api_base. + """ + api_base = ( + api_base + or get_secret_str("DOCKER_MODEL_RUNNER_API_BASE") + or "http://localhost:22088/engines/llama.cpp" + ) # type: ignore + # Docker Model Runner may not require authentication for local instances + dynamic_api_key = api_key or get_secret_str("DOCKER_MODEL_RUNNER_API_KEY") or "dummy-key" + return api_base, dynamic_api_key + + 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: + """ + Build the complete URL for Docker Model Runner API. + + Docker Model Runner uses URLs in the format: /engines/{engine}/v1/chat/completions + + The engine name should be specified in the api_base: + - api_base="http://model-runner.docker.internal/engines/llama.cpp" + - Default: "http://localhost:22088/engines/llama.cpp" + + Args: + api_base: Base URL for the Docker Model Runner instance including engine path + api_key: API key (may not be required for local instances) + model: Model name (e.g., "llama-3.1") + optional_params: Optional parameters + litellm_params: LiteLLM parameters + stream: Whether streaming is enabled + + Returns: + Complete URL for the API call + """ + if not api_base: + api_base = "http://localhost:22088/engines/llama.cpp" + + # Remove trailing slashes from api_base + api_base = api_base.rstrip("/") + + # Build the URL: {api_base}/v1/chat/completions + # api_base is expected to already contain the engine path + complete_url = f"{api_base}/v1/chat/completions" + + return complete_url + + def get_supported_openai_params(self, model: str) -> list: + """ + Get the supported OpenAI params for Docker Model Runner. + + Docker Model Runner is OpenAI-compatible and supports standard parameters. + """ + return super().get_supported_openai_params(model=model) + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to Docker Model Runner parameters. + + Docker Model Runner is OpenAI-compatible, so most parameters map directly. + """ + supported_openai_params = self.get_supported_openai_params(model) + for param, value in non_default_params.items(): + if param == "max_completion_tokens": + optional_params["max_tokens"] = value + elif param in supported_openai_params: + optional_params[param] = value + + return optional_params + diff --git a/litellm/llms/gemini/videos/transformation.py b/litellm/llms/gemini/videos/transformation.py index d1ae47af269..ce2519e9177 100644 --- a/litellm/llms/gemini/videos/transformation.py +++ b/litellm/llms/gemini/videos/transformation.py @@ -15,17 +15,16 @@ from litellm.images.utils import ImageEditRequestUtils import litellm from litellm.types.llms.gemini import GeminiLongRunningOperationResponse, GeminiVideoGenerationInstance, GeminiVideoGenerationParameters, GeminiVideoGenerationRequest from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS +from litellm.llms.base_llm.videos.transformation import BaseVideoConfig + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj - from ...base_llm.videos.transformation import BaseVideoConfig as _BaseVideoConfig from ...base_llm.chat.transformation import BaseLLMException as _BaseLLMException LiteLLMLoggingObj = _LiteLLMLoggingObj - BaseVideoConfig = _BaseVideoConfig BaseLLMException = _BaseLLMException else: LiteLLMLoggingObj = Any - BaseVideoConfig = Any BaseLLMException = Any diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index a85c37fd9b8..cc96e3415f3 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Union from uuid import uuid4 from litellm._logging import verbose_logger +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.exceptions import AuthenticationError from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.types.llms.openai import ( @@ -273,18 +274,29 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): """ return self._contains_vision_content(input_param) - def _contains_vision_content(self, value: Any) -> bool: + def _contains_vision_content( + self, value: Any, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH + ) -> bool: """ Recursively check if a value contains vision content. Looks for items with type="input_image" in the structure. """ + if depth > max_depth: + verbose_logger.warning( + f"[GitHub Copilot] Max recursion depth {max_depth} reached while checking for vision content" + ) + return False + if value is None: return False # Check arrays if isinstance(value, list): - return any(self._contains_vision_content(item) for item in value) + return any( + self._contains_vision_content(item, depth=depth + 1, max_depth=max_depth) + for item in value + ) # Only check dict/object types if not isinstance(value, dict): @@ -298,7 +310,8 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Check content field recursively if "content" in value and isinstance(value["content"], list): return any( - self._contains_vision_content(item) for item in value["content"] + self._contains_vision_content(item, depth=depth + 1, max_depth=max_depth) + for item in value["content"] ) return False diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index d49d86f8ca4..f1f0a67b9a1 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -81,6 +81,9 @@ from litellm.types.utils import ( TopLogprob, Usage, ) +from litellm.litellm_core_utils.prompt_templates.factory import ( + _encode_tool_call_id_with_signature, +) from litellm.utils import ( CustomStreamWrapper, ModelResponse, @@ -1192,7 +1195,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "function": _function_chunk, "index": cumulative_tool_call_idx, } + # Embed thought signature in ID for OpenAI client compatibility if thought_signature: + _tool_response_chunk[ + "id" + ] = _encode_tool_call_id_with_signature( + _tool_response_chunk["id"], thought_signature + ) _tool_response_chunk["provider_specific_fields"] = { # type: ignore "thought_signature": thought_signature } @@ -1702,7 +1711,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Convert thinking_blocks to reasoning_content for streaming # This ensures reasoning_content is available in streaming responses - if isinstance(model_response, ModelResponseStream) and reasoning_content is None: + if ( + isinstance(model_response, ModelResponseStream) + and reasoning_content is None + ): reasoning_content_parts = [] for block in thinking_blocks: thinking_text = block.get("thinking") diff --git a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py index caf347a1fb4..ad650e38499 100644 --- a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py @@ -10,6 +10,7 @@ from httpx._types import RequestFiles import litellm +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM from litellm.secret_managers.main import get_secret_str @@ -286,11 +287,18 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): return reference_images - def _read_all_bytes(self, image: Any) -> bytes: + def _read_all_bytes( + self, image: Any, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH + ) -> bytes: + if depth > max_depth: + raise ValueError( + f"Max recursion depth {max_depth} reached while reading image bytes for Vertex AI Imagen image edit." + ) + if isinstance(image, (list, tuple)): for item in image: if item is not None: - return self._read_all_bytes(item) + return self._read_all_bytes(item, depth=depth + 1, max_depth=max_depth) raise ValueError("Unsupported image type for Vertex AI Imagen image edit.") if isinstance(image, dict): @@ -302,9 +310,9 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): return base64.b64decode(value) except Exception: continue - return self._read_all_bytes(value) + return self._read_all_bytes(value, depth=depth + 1, max_depth=max_depth) if "path" in image: - return self._read_all_bytes(image["path"]) + return self._read_all_bytes(image["path"], depth=depth + 1, max_depth=max_depth) if isinstance(image, bytes): return image diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b4b4f763860..fb3d4c91710 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -5906,7 +5906,7 @@ "supports_function_calling": true, "supports_tool_choice": true }, - "cerebras/openai/gpt-oss-120b": { + "cerebras/gpt-oss-120b": { "input_cost_per_token": 2.5e-07, "litellm_provider": "cerebras", "max_input_tokens": 131072, @@ -11367,6 +11367,39 @@ "supports_web_search": true, "tpm": 8000000 }, + "gemini-3-pro-image-preview": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 65536, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini-2.5-flash-lite": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_audio_token": 5e-07, @@ -13071,6 +13104,39 @@ "supports_web_search": true, "tpm": 8000000 }, + "gemini/gemini-3-pro-image-preview": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 65536, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini/gemini-2.5-flash-lite": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_audio_token": 5e-07, @@ -19977,6 +20043,53 @@ "supports_tool_choice": true, "supports_vision": true }, + "openrouter/google/gemini-3-pro-preview": { + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "openrouter", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_batches": 6e-06, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true + }, "openrouter/google/gemini-pro-1.5": { "input_cost_per_image": 0.00265, "input_cost_per_token": 2.5e-06, @@ -22556,6 +22669,20 @@ "supports_parallel_function_calling": true, "supports_tool_choice": true }, + "together_ai/zai-org/GLM-4.6": { + "input_cost_per_token": 0.6e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 200000, + "max_output_tokens": 200000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 2.2e-06, + "source": "https://www.together.ai/models/glm-4-6", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "together_ai/moonshotai/Kimi-K2-Instruct-0905": { "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", @@ -24496,6 +24623,20 @@ "output_cost_per_image": 0.039, "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-generation#edit-an-image" }, + "vertex_ai/gemini-3-pro-image-preview": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 65536, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + }, "vertex_ai/imagegeneration@006": { "litellm_provider": "vertex_ai-image-models", "mode": "image_generation", @@ -26038,6 +26179,104 @@ "supports_tool_choice": true, "supports_web_search": true }, + "xai/grok-4-1-fast": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-reasoning": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-reasoning-latest": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-non-reasoning": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-non-reasoning-latest": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "xai/grok-beta": { "input_cost_per_token": 5e-06, "litellm_provider": "xai", diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 61123743c60..4e0a3e258cb 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -647,7 +647,7 @@ if MCP_AVAILABLE: allowed_mcp_server_ids = ( await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) ) - allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( + allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( # type: ignore[attr-defined] allowed_mcp_server_ids ) @@ -1173,7 +1173,7 @@ if MCP_AVAILABLE: ) ) - allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( + allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( # type: ignore[attr-defined] allowed_mcp_server_ids ) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index ac7082edb69..b0f13f6afbc 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -22,7 +22,6 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, - convert_b64_uid_to_unified_uid, get_batch_id_from_unified_batch_id, get_model_id_from_unified_batch_id, get_models_from_unified_file_id, diff --git a/litellm/proxy/example_config_yaml/tool_permission_example.yaml b/litellm/proxy/example_config_yaml/tool_permission_example.yaml index e18425ba383..735b4bb7ed2 100644 --- a/litellm/proxy/example_config_yaml/tool_permission_example.yaml +++ b/litellm/proxy/example_config_yaml/tool_permission_example.yaml @@ -10,6 +10,7 @@ guardrails: guardrail: tool_permission mode: "post_call" default_on: true # Apply to all requests by default + violation_message_template: "this violates our org policy, we don't support executing {tool_name} commands" rules: - id: "allow_bash" tool_name: "Bash" @@ -33,4 +34,4 @@ general_settings: # Optional: Add logging configuration litellm_settings: success_callback: ["langfuse"] - failure_callback: ["langfuse"] \ No newline at end of file + failure_callback: ["langfuse"] diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index 97f8dd76bd4..19060fa9d6d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -120,13 +120,27 @@ class ToolPermissionGuardrail(CustomGuardrail): for rule in self.rules: if self._matches_pattern(tool_name, rule.tool_name): is_allowed = rule.decision == "allow" - message = f"Tool '{tool_name}' {'allowed' if is_allowed else 'denied'} by rule '{rule.id}'" + default_message = f"Tool '{tool_name}' {'allowed' if is_allowed else 'denied'} by rule '{rule.id}'" + message = self.render_violation_message( + default=default_message, + context={ + "tool_name": tool_name, + "rule_id": rule.id, + }, + ) verbose_proxy_logger.debug(message) return is_allowed, rule.id, message # No rule matched, use default action is_allowed = self.default_action == "allow" - message = f"Tool '{tool_name}' {'allowed' if is_allowed else 'denied'} by default action" + default_message = f"Tool '{tool_name}' {'allowed' if is_allowed else 'denied'} by default action" + message = self.render_violation_message( + default=default_message, + context={ + "tool_name": tool_name, + "rule_id": None, + }, + ) verbose_proxy_logger.debug(message) return is_allowed, None, message @@ -449,7 +463,9 @@ class ToolPermissionGuardrail(CustomGuardrail): verbose_proxy_logger.debug("Tool Permission Guardrail: Checking response") # Extract tool_calls from the response - tool_calls = self._extract_tool_calls_from_response(assembled_model_response) + tool_calls = self._extract_tool_calls_from_response( + assembled_model_response + ) if not tool_calls: verbose_proxy_logger.debug( diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 6a5ba22419b..f2083e9c67e 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -135,6 +135,7 @@ def initialize_tool_permission(litellm_params: LitellmParams, guardrail: Guardra default_action=getattr(litellm_params, "default_action", "deny"), on_disallowed_action=getattr(litellm_params, "on_disallowed_action", "block"), default_on=litellm_params.default_on, + violation_message_template=litellm_params.violation_message_template, ) litellm.logging_callback_manager.add_litellm_callback(_tool_permission_callback) return _tool_permission_callback @@ -172,9 +173,12 @@ def initialize_panw_prisma_airs(litellm_params, guardrail): raise ValueError("PANW Prisma AIRS: profile_name is required") _panw_callback = PanwPrismaAirsHandler( - guardrail_name=guardrail.get("guardrail_name", "panw_prisma_airs"), # Use .get() with default + guardrail_name=guardrail.get( + "guardrail_name", "panw_prisma_airs" + ), # Use .get() with default api_key=litellm_params.api_key, - api_base=litellm_params.api_base or "https://service.api.aisecurity.paloaltonetworks.com/v1/scan/sync/request", + api_base=litellm_params.api_base + or "https://service.api.aisecurity.paloaltonetworks.com/v1/scan/sync/request", profile_name=litellm_params.profile_name, default_on=litellm_params.default_on, ) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 66085b69b3d..2ca1c4cc483 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -705,6 +705,7 @@ def _process_keys_for_user_info( keys: Optional[List[LiteLLM_VerificationToken]], all_teams: Optional[Union[List[LiteLLM_TeamTable], List[TeamListResponseObject]]], ): + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID from litellm.proxy.proxy_server import general_settings, litellm_master_key_hash returned_keys = [] @@ -724,6 +725,11 @@ def _process_keys_for_user_info( except Exception: # if using pydantic v1 _key = key.dict() + + # Filter out UI session tokens (team_id="litellm-dashboard") + if _key.get("team_id") == UI_SESSION_TOKEN_TEAM_ID: + continue + if ( "team_id" in _key and _key["team_id"] is not None diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index ddd838a3de6..f2b6abfc262 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -58,7 +58,6 @@ if MCP_AVAILABLE: from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.management_helpers.utils import management_endpoint_wrapper - from litellm.types.mcp_server.mcp_server_manager import MCPInfo def _redact_mcp_credentials( mcp_server: LiteLLM_MCPServerTable, diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index b4677ed374b..ac39e5f4b2c 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -6,7 +6,15 @@ import tempfile from pathlib import Path from typing import Any, Dict, List, Optional, cast -from fastapi import APIRouter, Depends, File, HTTPException, UploadFile +from fastapi import ( + APIRouter, + Depends, + File, + HTTPException, + Request, + Response, + UploadFile, +) from pydantic import BaseModel from litellm._logging import verbose_proxy_logger @@ -20,10 +28,168 @@ from litellm.types.prompts.init_prompts import ( PromptSpec, PromptTemplateBase, ) +from litellm.types.proxy.prompt_endpoints import TestPromptRequest router = APIRouter() +def get_base_prompt_id(prompt_id: str) -> str: + """ + Extract the base prompt ID by stripping the version suffix if present. + + Args: + prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v1" or "jack_success_v1") + + Returns: + Base prompt ID without version suffix (e.g., "jack_success") + + Examples: + >>> get_base_prompt_id("jack_success.v1") + "jack_success" + >>> get_base_prompt_id("jack_success_v1") + "jack_success" + >>> get_base_prompt_id("jack_success") + "jack_success" + """ + # Try dot separator first (.v) + if ".v" in prompt_id: + return prompt_id.split(".v")[0] + # Try underscore separator (_v) + if "_v" in prompt_id: + return prompt_id.split("_v")[0] + return prompt_id + + +def get_version_number(prompt_id: str) -> int: + """ + Extract the version number from a versioned prompt ID. + + Args: + prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v2" or "jack_success_v2") + + Returns: + Version number (defaults to 1 if no version suffix or invalid format) + + Examples: + >>> get_version_number("jack_success.v2") + 2 + >>> get_version_number("jack_success_v2") + 2 + >>> get_version_number("jack_success") + 1 + """ + # Try dot separator first (.v) + if ".v" in prompt_id: + version_str = prompt_id.split(".v")[1] + try: + return int(version_str) + except ValueError: + pass + + # Try underscore separator (_v) + if "_v" in prompt_id: + version_str = prompt_id.split("_v")[1] + try: + return int(version_str) + except ValueError: + pass + + return 1 + + +def construct_versioned_prompt_id(prompt_id: str, version: Optional[int] = None) -> str: + """ + Construct a versioned prompt ID from a base prompt_id and version number. + + Args: + prompt_id: Base prompt ID (e.g., "jack_success") + version: Version number (if None, returns the base prompt_id unchanged) + + Returns: + Versioned prompt ID (e.g., "jack_success.v4") + + Examples: + >>> construct_versioned_prompt_id("jack_success", 4) + "jack_success.v4" + >>> construct_versioned_prompt_id("jack_success", None) + "jack_success" + >>> construct_versioned_prompt_id("jack_success.v2", 4) + "jack_success.v4" + """ + if version is None: + return prompt_id + + # Strip any existing version suffix first + base_id = get_base_prompt_id(prompt_id) + return f"{base_id}.v{version}" + + +def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Dict[str, Any]) -> str: + """ + Find the latest version of a prompt from available prompt IDs. + + Args: + prompt_id: Base prompt ID or versioned prompt ID (e.g., "jack_success" or "jack_success.v2") + all_prompt_ids: Dictionary of all available prompt IDs (keys are prompt IDs) + + Returns: + The prompt ID with the highest version number, or the original prompt_id if no versions exist + + Examples: + >>> all_ids = {"jack.v1": {}, "jack.v2": {}, "jack.v3": {}} + >>> get_latest_version_prompt_id("jack", all_ids) + "jack.v3" + >>> get_latest_version_prompt_id("jack.v1", all_ids) + "jack.v3" + >>> all_ids = {"simple": {}} + >>> get_latest_version_prompt_id("simple", all_ids) + "simple" + """ + base_id = get_base_prompt_id(prompt_id=prompt_id) + + # Find all versions of this prompt + matching_versions = [] + for stored_prompt_id in all_prompt_ids.keys(): + if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id: + version_num = get_version_number(prompt_id=stored_prompt_id) + matching_versions.append((version_num, stored_prompt_id)) + + # Use the highest version number + if matching_versions: + matching_versions.sort(reverse=True) + return matching_versions[0][1] + else: + # No versioned prompts found, use the base ID as-is + return prompt_id + + +def get_latest_prompt_versions(prompts: List[PromptSpec]) -> List[PromptSpec]: + """ + Filter a list of prompts to return only the latest version of each unique prompt. + + Args: + prompts: List of PromptSpec objects + + Returns: + List of PromptSpec objects with only the latest version of each prompt + """ + latest_prompts: Dict[str, PromptSpec] = {} + + for prompt in prompts: + base_id = get_base_prompt_id(prompt_id=prompt.prompt_id) + version = get_version_number(prompt_id=prompt.prompt_id) + + # Keep the prompt with the highest version number + if base_id not in latest_prompts: + latest_prompts[base_id] = prompt + else: + existing_version = get_version_number(prompt_id=latest_prompts[base_id].prompt_id) + if version > existing_version: + latest_prompts[base_id] = prompt + + return list(latest_prompts.values()) + + async def get_next_version_for_prompt(prisma_client, prompt_id: str) -> int: """ Get the next version number for a prompt. @@ -150,25 +316,140 @@ async def list_prompts( if key_metadata is not None: prompts = cast(Optional[List[str]], key_metadata.get("prompts", None)) if prompts is not None: - return ListPromptsResponse( - prompts=[ - IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt] - for prompt in prompts - if prompt in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS - ] - ) + prompt_list = [] + for prompt_id in prompts: + if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS: + original_prompt = IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id] + # Create a copy with base prompt_id (without version suffix) + prompt_copy = PromptSpec( + prompt_id=get_base_prompt_id(prompt_id=original_prompt.prompt_id), + litellm_params=original_prompt.litellm_params, + prompt_info=original_prompt.prompt_info, + created_at=original_prompt.created_at, + updated_at=original_prompt.updated_at, + ) + prompt_list.append(prompt_copy) + return ListPromptsResponse(prompts=prompt_list) # check if user is proxy admin - show all prompts if user_api_key_dict.user_role is not None and ( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value ): - return ListPromptsResponse( - prompts=list(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()) - ) + # Get all prompts and filter to show only the latest version of each + all_prompts = list(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()) + latest_prompts = get_latest_prompt_versions(prompts=all_prompts) + # Create copies with base prompt_id (without version suffix) for display + prompts_for_display = [] + for original_prompt in latest_prompts: + prompt_copy = PromptSpec( + prompt_id=get_base_prompt_id(prompt_id=original_prompt.prompt_id), + litellm_params=original_prompt.litellm_params, + prompt_info=original_prompt.prompt_info, + created_at=original_prompt.created_at, + updated_at=original_prompt.updated_at, + ) + prompts_for_display.append(prompt_copy) + return ListPromptsResponse(prompts=prompts_for_display) else: return ListPromptsResponse(prompts=[]) +@router.get( + "/prompts/{prompt_id}/versions", + tags=["Prompt Management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ListPromptsResponse, +) +async def get_prompt_versions( + prompt_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get all versions of a specific prompt by base prompt ID + + 👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management) + + Example Request: + ```bash + curl -X GET "http://localhost:4000/prompts/jack_success/versions" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "prompts": [ + { + "prompt_id": "jack_success.v1", + "litellm_params": {...}, + "prompt_info": {"prompt_type": "db"}, + "created_at": "2023-11-09T12:34:56.789Z", + "updated_at": "2023-11-09T12:34:56.789Z" + }, + { + "prompt_id": "jack_success.v2", + "litellm_params": {...}, + "prompt_info": {"prompt_type": "db"}, + "created_at": "2023-11-09T13:45:12.345Z", + "updated_at": "2023-11-09T13:45:12.345Z" + } + ] + } + ``` + """ + from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY + + # Only allow proxy admins to view version history + if user_api_key_dict.user_role is None or ( + user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN + and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + ): + raise HTTPException( + status_code=403, detail="Only proxy admins can view prompt versions" + ) + + # Strip version suffix if provided (e.g., "jack_success.v1" -> "jack_success") + base_prompt_id = get_base_prompt_id(prompt_id=prompt_id) + + # Get all prompts and filter by base_prompt_id + all_prompts = list(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()) + prompt_versions = [ + prompt for prompt in all_prompts + if get_base_prompt_id(prompt_id=prompt.prompt_id) == base_prompt_id + ] + + if not prompt_versions: + raise HTTPException( + status_code=404, detail=f"No versions found for prompt ID {base_prompt_id}" + ) + + # Create response with explicit version field for each prompt + versioned_prompts = [] + for prompt in prompt_versions: + # Extract version number from the root prompt_id which has version suffix + # (e.g., "jack-sparrow.v3" -> 3) + version_number = get_version_number(prompt_id=prompt.prompt_id) + + # Strip version from prompt_id for clean display + base_prompt_id = get_base_prompt_id(prompt_id=prompt.prompt_id) + + # Create a copy with explicit version field and clean prompt_id + versioned_prompt = PromptSpec( + prompt_id=base_prompt_id, # Clean ID without version (e.g., "jack-sparrow") + litellm_params=prompt.litellm_params, + prompt_info=prompt.prompt_info, + created_at=prompt.created_at, + updated_at=prompt.updated_at, + version=version_number, # Explicit version field (e.g., 3) + ) + versioned_prompts.append(versioned_prompt) + + # Sort by version number (descending - newest first) + versioned_prompts.sort(key=lambda p: p.version or 1, reverse=True) + + return ListPromptsResponse(prompts=versioned_prompts) + + @router.get( "/prompts/{prompt_id}", tags=["Prompt Management"], @@ -235,10 +516,34 @@ async def get_prompt_info( detail=f"You are not authorized to access this prompt. Your role - {user_api_key_dict.user_role}, Your key's prompts - {prompts}", ) + # Try to get prompt directly first prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id) + + # If not found, try to find the latest version + if prompt_spec is None: + latest_prompt_id = get_latest_version_prompt_id( + prompt_id=prompt_id, + all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS + ) + prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id) + if prompt_spec is None: raise HTTPException(status_code=400, detail=f"Prompt {prompt_id} not found") + # Extract version number from the prompt_id + version_number = get_version_number(prompt_id=prompt_spec.prompt_id) + + # Create a copy of the prompt spec with the base prompt ID (stripped of version) + # and explicit version field for consistency with list_prompts and versions endpoints + prompt_spec_response = PromptSpec( + prompt_id=get_base_prompt_id(prompt_id=prompt_spec.prompt_id), + litellm_params=prompt_spec.litellm_params, # This preserves the versioned ID + prompt_info=prompt_spec.prompt_info, + created_at=prompt_spec.created_at, + updated_at=prompt_spec.updated_at, + version=version_number, # Explicit version field + ) + # Get prompt content from the callback prompt_template: Optional[PromptTemplateBase] = None try: @@ -269,7 +574,7 @@ async def get_prompt_info( # Create response with content return PromptInfoResponse( - prompt_spec=prompt_spec, + prompt_spec=prompt_spec_response, raw_prompt_template=prompt_template, ) @@ -398,8 +703,6 @@ async def update_prompt( }' ``` """ - from datetime import datetime - from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY from litellm.proxy.proxy_server import prisma_client @@ -418,19 +721,21 @@ async def update_prompt( ) try: + # Strip version suffix from prompt_id if present (e.g., "jack_success.v1" -> "jack_success") + base_prompt_id = get_base_prompt_id(prompt_id=prompt_id) + # Check if any version exists existing_prompts = await prisma_client.db.litellm_prompttable.find_many( - where={"prompt_id": request.prompt_id} + where={"prompt_id": base_prompt_id} ) if not existing_prompts: raise HTTPException( - status_code=404, detail=f"Prompt with ID {request.prompt_id} not found" + status_code=404, detail=f"Prompt with ID {base_prompt_id} not found" ) # Check if it's a config prompt - base_prompt_id = request.prompt_id - existing_in_memory = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(base_prompt_id) + existing_in_memory = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id) if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config": raise HTTPException( status_code=400, @@ -439,13 +744,13 @@ async def update_prompt( # Get next version number (UPDATE creates a new version) new_version = await get_next_version_for_prompt( - prisma_client=prisma_client, prompt_id=request.prompt_id + prisma_client=prisma_client, prompt_id=base_prompt_id ) # Store new version in db prompt_db_entry = await prisma_client.db.litellm_prompttable.create( data={ - "prompt_id": request.prompt_id, + "prompt_id": base_prompt_id, "version": new_version, "litellm_params": request.litellm_params.model_dump_json(), "prompt_info": ( @@ -521,8 +826,19 @@ async def delete_prompt( ) try: - # Check if prompt exists + # Try to get prompt directly first existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id) + + # If not found, try to find the latest version + if existing_prompt is None: + latest_prompt_id = get_latest_version_prompt_id( + prompt_id=prompt_id, + all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS + ) + existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id) + # Use the resolved prompt_id for deletion + prompt_id = latest_prompt_id + if existing_prompt is None: raise HTTPException( status_code=404, detail=f"Prompt with ID {prompt_id} not found" @@ -667,6 +983,154 @@ async def patch_prompt( raise HTTPException(status_code=500, detail=str(e)) +@router.post( + "/prompts/test", + tags=["Prompt Management"], + dependencies=[Depends(user_api_key_auth)], +) +async def test_prompt( + request: TestPromptRequest, + fastapi_request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Test a prompt by rendering it with variables and executing an LLM call. + + This endpoint allows testing prompts before saving them to the database. + The response is always streamed. + + 👉 [Prompt docs](https://docs.litellm.ai/docs/proxy/prompt_management) + + Example Request: + ```bash + curl -X POST "http://localhost:4000/prompts/test" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "dotprompt_content": "---\\nmodel: gpt-4o\\ntemperature: 0.7\\n---\\n\\nUser: Hello {{name}}", + "prompt_variables": { + "name": "World" + } + }' + ``` + """ + from pydantic import BaseModel + + from litellm.integrations.dotprompt.dotprompt_manager import DotpromptManager + from litellm.integrations.dotprompt.prompt_manager import ( + PromptManager, + PromptTemplate, + ) + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + try: + # Parse the dotprompt content and create PromptTemplate + prompt_manager = PromptManager() + frontmatter, template_content = prompt_manager._parse_frontmatter( + content=request.dotprompt_content + ) + + # Create PromptTemplate to leverage existing parameter extraction logic + template = PromptTemplate( + content=template_content, + metadata=frontmatter, + template_id="test_prompt" + ) + + # Extract model from template + if not template.model: + raise HTTPException( + status_code=400, + detail="Model is required in dotprompt metadata" + ) + + # Always render the template to extract system messages and other metadata + variables = request.prompt_variables or {} + rendered_content = prompt_manager.jinja_env.from_string( + template_content + ).render(**variables) + + # Convert rendered content to messages using DotpromptManager's method + dotprompt_manager = DotpromptManager() + rendered_messages = dotprompt_manager._convert_to_messages( + rendered_content=rendered_content + ) + + if not rendered_messages: + raise HTTPException( + status_code=400, + detail="No messages found in rendered prompt" + ) + + # If conversation history is provided, use it but preserve system messages + if request.conversation_history: + # Extract system messages from rendered prompt + system_messages = [msg for msg in rendered_messages if msg.get("role") == "system"] + # Use conversation history for user/assistant messages + messages = system_messages + request.conversation_history + else: + messages = rendered_messages # type: ignore[assignment] + + # Use PromptTemplate's optional_params which already extracts all parameters + optional_params = template.optional_params.copy() + + # Always stream the response + optional_params["stream"] = True + + # Build request data for chat completion + data = { + "model": template.model, + "messages": messages, + } + data.update(optional_params) + + # Use ProxyBaseLLMRequestProcessing to go through all proxy logic + base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) + result = await base_llm_response_processor.base_process_llm_request( + request=fastapi_request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="acompletion", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=None, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + + if isinstance(result, BaseModel): + return result.model_dump(exclude_none=True, exclude_unset=True) + else: + return result + + except HTTPException as e: + raise e + except Exception as e: + verbose_proxy_logger.exception(f"Error testing prompt: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + @router.post( "/utils/dotprompt_json_converter", tags=["prompts", "utils"], diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 2e58d8554c6..014bcdc1670 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -26,10 +26,10 @@ search_tools: litellm_params: search_provider: perplexity api_key: os.environ/PERPLEXITYAI_API_KEY - - search_tool_name: exa-search + - search_tool_name: firecrawl-search litellm_params: - search_provider: exa_ai - api_key: os.environ/EXA_API_KEY + search_provider: firecrawl + api_key: os.environ/FIRECRAWL_API_KEY litellm_settings: diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 3b7abbf4f0e..71c91fac6f7 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -9,7 +9,6 @@ from litellm.proxy.public_endpoints.provider_create_metadata import ( ) from litellm.types.agents import AgentCard from litellm.types.mcp import MCPPublicServer -from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.proxy.management_endpoints.model_management_endpoints import ( ModelGroupInfoProxy, ) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 9747329cd4f..5ec8aecfef0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -894,6 +894,7 @@ class ProxyLogging: Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None) ) prompt_id = data.get("prompt_id", None) + prompt_version = data.get("prompt_version", None) ## PROMPT TEMPLATE CHECK ## if ( @@ -901,12 +902,28 @@ class ProxyLogging: and prompt_id is not None and (call_type == "completion" or call_type == "acompletion") ): + from litellm.proxy.prompts.prompt_endpoints import ( + construct_versioned_prompt_id, + get_latest_version_prompt_id, + ) from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY + # If no version is specified, find the latest version + if prompt_version is None: + lookup_prompt_id = get_latest_version_prompt_id( + prompt_id=prompt_id, + all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS, + ) + else: + # Construct versioned prompt_id if prompt_version is provided + lookup_prompt_id = construct_versioned_prompt_id( + prompt_id=prompt_id, version=prompt_version + ) + custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id( - prompt_id + lookup_prompt_id ) - prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id) + prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id) litellm_prompt_id: Optional[str] = None if prompt_spec is not None: litellm_prompt_id = prompt_spec.litellm_params.prompt_id diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index e0890aaef57..f8f5154e2d6 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -110,7 +110,7 @@ class LiteLLM_Proxy_MCP_Handler: allowed_mcp_server_ids = ( await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) ) - allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( + allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( # type: ignore[attr-defined] allowed_mcp_server_ids ) diff --git a/litellm/router.py b/litellm/router.py index 998e19739b7..841391653da 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -835,8 +835,8 @@ class Router: litellm.acancel_batch, call_type="acancel_batch" ) - def _initialize_specialized_endpoints(self): - """Helper to initialize specialized router endpoints (vector store, OCR, search, video, container).""" + def _initialize_vector_store_endpoints(self): + """Initialize vector store endpoints.""" from litellm.vector_stores.main import acreate, asearch, create, search self.avector_store_search = self.factory_function( @@ -852,6 +852,8 @@ class Router: create, call_type="vector_store_create" ) + def _initialize_vector_store_file_endpoints(self): + """Initialize vector store file endpoints.""" from litellm.vector_store_files.main import ( acreate as avector_store_file_create_fn, ) @@ -921,6 +923,8 @@ class Router: vector_store_file_delete_fn, call_type="vector_store_file_delete" ) + def _initialize_google_genai_endpoints(self): + """Initialize Google GenAI endpoints.""" from litellm.google_genai import ( agenerate_content, agenerate_content_stream, @@ -941,6 +945,8 @@ class Router: generate_content_stream, call_type="generate_content_stream" ) + def _initialize_ocr_search_endpoints(self): + """Initialize OCR and search endpoints.""" from litellm.ocr import aocr, ocr self.aocr = self.factory_function(aocr, call_type="aocr") @@ -951,6 +957,8 @@ class Router: self.asearch = self.factory_function(asearch, call_type="asearch") self.search = self.factory_function(search, call_type="search") + def _initialize_video_endpoints(self): + """Initialize video endpoints.""" from litellm.videos import ( avideo_content, avideo_generation, @@ -989,6 +997,8 @@ class Router: ) self.video_remix = self.factory_function(video_remix, call_type="video_remix") + def _initialize_container_endpoints(self): + """Initialize container endpoints.""" from litellm.containers import ( acreate_container, adelete_container, @@ -1025,6 +1035,15 @@ class Router: delete_container, call_type="delete_container" ) + def _initialize_specialized_endpoints(self): + """Helper to initialize specialized router endpoints (vector store, OCR, search, video, container).""" + self._initialize_vector_store_endpoints() + self._initialize_vector_store_file_endpoints() + self._initialize_google_genai_endpoints() + self._initialize_ocr_search_endpoints() + self._initialize_video_endpoints() + self._initialize_container_endpoints() + def initialize_router_endpoints(self): self._initialize_core_endpoints() self._initialize_specialized_endpoints() diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index cae9623b44b..f2b9d71cca6 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -16,7 +16,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( ) - """ Pydantic object defining how to set guardrails on litellm proxy @@ -51,7 +50,7 @@ class SupportedGuardrailIntegrations(Enum): OPENAI_MODERATION = "openai_moderation" NOMA = "noma" TOOL_PERMISSION = "tool_permission" - ZSCALER_AI_GUARD = "zscaler_ai_guard" + ZSCALER_AI_GUARD = "zscaler_ai_guard" JAVELIN = "javelin" ENKRYPTAI = "enkryptai" IBM_GUARDRAILS = "ibm_guardrails" @@ -432,7 +431,7 @@ class ZscalerAIGuardConfigModel(BaseModel): policy_id: Optional[int] = Field( default=None, - description="Policy ID for Zscaler AI Guard. Can also be set via ZSCALER_AI_GUARD_POLICY_ID environment variable" + description="Policy ID for Zscaler AI Guard. Can also be set via ZSCALER_AI_GUARD_POLICY_ID environment variable", ) send_user_api_key_alias: Optional[bool] = Field( default=False, description="Whether to send user_API_key_alias in headers" @@ -444,6 +443,7 @@ class ZscalerAIGuardConfigModel(BaseModel): default=False, description="Whether to send user_API_key_team_id in headers" ) + class JavelinGuardrailConfigModel(BaseModel): """Configuration parameters for the Javelin guardrail""" @@ -479,7 +479,8 @@ class BlockedWord(BaseModel): description="Action to take when keyword is detected (BLOCK or MASK)" ) description: Optional[str] = Field( - default=None, description="Optional description explaining why this keyword is sensitive" + default=None, + description="Optional description explaining why this keyword is sensitive", ) @@ -491,15 +492,15 @@ class ContentFilterPattern(BaseModel): ) pattern_name: Optional[str] = Field( default=None, - description="Name of prebuilt pattern (e.g., 'us_ssn', 'credit_card'). Required if pattern_type is 'prebuilt'" + description="Name of prebuilt pattern (e.g., 'us_ssn', 'credit_card'). Required if pattern_type is 'prebuilt'", ) pattern: Optional[str] = Field( default=None, - description="Custom regex pattern. Required if pattern_type is 'regex'" + description="Custom regex pattern. Required if pattern_type is 'regex'", ) name: Optional[str] = Field( default=None, - description="Name for this pattern (used in logging and error messages)" + description="Name for this pattern (used in logging and error messages)", ) action: ContentFilterAction = Field( description="Action to take when pattern matches (BLOCK or MASK)" @@ -511,15 +512,13 @@ class ContentFilterConfigModel(BaseModel): patterns: Optional[List[ContentFilterPattern]] = Field( default=None, - description="List of patterns (prebuilt or custom regex) to detect" + description="List of patterns (prebuilt or custom regex) to detect", ) blocked_words: Optional[List[BlockedWord]] = Field( - default=None, - description="List of blocked words with individual actions" + default=None, description="List of blocked words with individual actions" ) blocked_words_file: Optional[str] = Field( - default=None, - description="Path to YAML file containing blocked_words list" + default=None, description="Path to YAML file containing blocked_words list" ) @@ -575,6 +574,11 @@ class BaseLitellmParams(BaseModel): # works for new and patch update guardrails description="Optional field if guardrail requires a 'model' parameter", ) + violation_message_template: Optional[str] = Field( + default=None, + description="Custom message when a guardrail blocks an action. Supports placeholders like {tool_name}, {rule_id}, and {default_message}.", + ) + # Model Armor params template_id: Optional[str] = Field( default=None, description="The ID of your Model Armor template" @@ -613,7 +617,7 @@ class LitellmParams( GraySwanGuardrailConfigModel, NomaGuardrailConfigModel, ToolPermissionGuardrailConfigModel, - ZscalerAIGuardConfigModel, + ZscalerAIGuardConfigModel, JavelinGuardrailConfigModel, ContentFilterConfigModel, BaseLitellmParams, @@ -671,10 +675,12 @@ class GuardrailEventHooks(str, Enum): class DynamicGuardrailParams(TypedDict): extra_body: Dict[str, Any] + class GUARDRAIL_DEFINITION_LOCATION(str, Enum): DB = "db" CONFIG = "config" + class GuardrailInfoResponse(BaseModel): guardrail_id: Optional[str] = None guardrail_name: str @@ -682,7 +688,9 @@ class GuardrailInfoResponse(BaseModel): guardrail_info: Optional[Dict] = None created_at: Optional[datetime] = None updated_at: Optional[datetime] = None - guardrail_definition_location: GUARDRAIL_DEFINITION_LOCATION = GUARDRAIL_DEFINITION_LOCATION.CONFIG + guardrail_definition_location: GUARDRAIL_DEFINITION_LOCATION = ( + GUARDRAIL_DEFINITION_LOCATION.CONFIG + ) def __init__(self, **kwargs): super().__init__(**kwargs) diff --git a/litellm/types/prompts/init_prompts.py b/litellm/types/prompts/init_prompts.py index 102ace93a52..e046d6a45a0 100644 --- a/litellm/types/prompts/init_prompts.py +++ b/litellm/types/prompts/init_prompts.py @@ -37,6 +37,7 @@ class PromptSpec(BaseModel): prompt_info: PromptInfo created_at: Optional[datetime] = None updated_at: Optional[datetime] = None + version: Optional[int] = None # Version number for version history def __init__(self, **data): if "prompt_info" not in data: diff --git a/litellm/types/proxy/prompt_endpoints.py b/litellm/types/proxy/prompt_endpoints.py new file mode 100644 index 00000000000..620a565b0a2 --- /dev/null +++ b/litellm/types/proxy/prompt_endpoints.py @@ -0,0 +1,10 @@ +from typing import Any, Dict, List, Optional + +from pydantic import BaseModel + + +class TestPromptRequest(BaseModel): + dotprompt_content: str + prompt_variables: Optional[Dict[str, Any]] = None + conversation_history: Optional[List[Dict[str, str]]] = None + diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 9bef6ebef7d..a751adf542a 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -302,6 +302,10 @@ class CallTypes(str, Enum): avector_store_file_update = "avector_store_file_update" vector_store_file_delete = "vector_store_file_delete" avector_store_file_delete = "avector_store_file_delete" + vector_store_create = "vector_store_create" + avector_store_create = "avector_store_create" + vector_store_search = "vector_store_search" + avector_store_search = "avector_store_search" ######################################################### # Container Call Types @@ -375,8 +379,10 @@ CallTypesLiteral = Literal[ "agenerate_content_stream", "ocr", "aocr", - "avector_store_search", + "vector_store_create", + "avector_store_create", "vector_store_search", + "avector_store_search", "vector_store_file_create", "avector_store_file_create", "vector_store_file_list", @@ -2472,6 +2478,7 @@ all_litellm_params = ( "use_litellm_proxy", "prompt_label", "shared_session", + "search_tool_name", ] + list(StandardCallbackDynamicParams.__annotations__.keys()) + list(CustomPricingLiteLLMParams.model_fields.keys()) @@ -2587,6 +2594,7 @@ class LlmProviders(str, Enum): EMPOWER = "empower" GITHUB = "github" COMPACTIFAI = "compactifai" + DOCKER_MODEL_RUNNER = "docker_model_runner" CUSTOM = "custom" LITELLM_PROXY = "litellm_proxy" HOSTED_VLLM = "hosted_vllm" @@ -2722,7 +2730,7 @@ class LiteLLMFineTuningJob(FineTuningJob): class LiteLLMBatch(Batch): _hidden_params: dict = {} - usage: Optional[Usage] = None + usage: Optional[Usage] = None # type: ignore[assignment] def __contains__(self, key): # Define custom behavior for the 'in' operator diff --git a/litellm/utils.py b/litellm/utils.py index 8e6809c511b..474f5d57b12 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7205,6 +7205,8 @@ class ProviderConfigManager: return litellm.DashScopeChatConfig() elif litellm.LlmProviders.MOONSHOT == provider: return litellm.MoonshotChatConfig() + elif litellm.LlmProviders.DOCKER_MODEL_RUNNER == provider: + return litellm.DockerModelRunnerChatConfig() elif litellm.LlmProviders.V0 == provider: return litellm.V0ChatConfig() elif litellm.LlmProviders.MORPH == provider: @@ -7758,7 +7760,9 @@ class ProviderConfigManager: return LiteLLMProxyImageEditConfig() elif LlmProviders.VERTEX_AI == provider: - from litellm.llms.vertex_ai.image_edit import get_vertex_ai_image_edit_config + from litellm.llms.vertex_ai.image_edit import ( + get_vertex_ai_image_edit_config, + ) return get_vertex_ai_image_edit_config(model) return None diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b4b4f763860..fb3d4c91710 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -5906,7 +5906,7 @@ "supports_function_calling": true, "supports_tool_choice": true }, - "cerebras/openai/gpt-oss-120b": { + "cerebras/gpt-oss-120b": { "input_cost_per_token": 2.5e-07, "litellm_provider": "cerebras", "max_input_tokens": 131072, @@ -11367,6 +11367,39 @@ "supports_web_search": true, "tpm": 8000000 }, + "gemini-3-pro-image-preview": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 65536, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini-2.5-flash-lite": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_audio_token": 5e-07, @@ -13071,6 +13104,39 @@ "supports_web_search": true, "tpm": 8000000 }, + "gemini/gemini-3-pro-image-preview": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 65536, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini/gemini-2.5-flash-lite": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_audio_token": 5e-07, @@ -19977,6 +20043,53 @@ "supports_tool_choice": true, "supports_vision": true }, + "openrouter/google/gemini-3-pro-preview": { + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "openrouter", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_200k_tokens": 1.8e-05, + "output_cost_per_token_batches": 6e-06, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true + }, "openrouter/google/gemini-pro-1.5": { "input_cost_per_image": 0.00265, "input_cost_per_token": 2.5e-06, @@ -22556,6 +22669,20 @@ "supports_parallel_function_calling": true, "supports_tool_choice": true }, + "together_ai/zai-org/GLM-4.6": { + "input_cost_per_token": 0.6e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 200000, + "max_output_tokens": 200000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 2.2e-06, + "source": "https://www.together.ai/models/glm-4-6", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "together_ai/moonshotai/Kimi-K2-Instruct-0905": { "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", @@ -24496,6 +24623,20 @@ "output_cost_per_image": 0.039, "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/image-generation#edit-an-image" }, + "vertex_ai/gemini-3-pro-image-preview": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 65536, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + }, "vertex_ai/imagegeneration@006": { "litellm_provider": "vertex_ai-image-models", "mode": "image_generation", @@ -26038,6 +26179,104 @@ "supports_tool_choice": true, "supports_web_search": true }, + "xai/grok-4-1-fast": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-reasoning": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-reasoning-latest": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-non-reasoning": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "xai/grok-4-1-fast-non-reasoning-latest": { + "cache_read_input_token_cost": 0.05e-06, + "input_cost_per_token": 0.2e-06, + "input_cost_per_token_above_128k_tokens": 0.4e-06, + "litellm_provider": "xai", + "max_input_tokens": 2e6, + "max_output_tokens": 2e6, + "max_tokens": 2e6, + "mode": "chat", + "output_cost_per_token": 0.5e-06, + "output_cost_per_token_above_128k_tokens": 1e-06, + "source": "https://docs.x.ai/docs/models/grok-4-1-fast-non-reasoning", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "xai/grok-beta": { "input_cost_per_token": 5e-06, "litellm_provider": "xai", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 4d9218d6095..368d12b6605 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -1036,6 +1036,22 @@ "rerank": false } }, + "docker_model_runner": { + "display_name": "Docker Model Runner (`docker_model_runner`)", + "url": "https://docs.litellm.ai/docs/providers/docker_model_runner", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": true, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false + } + }, "morph": { "display_name": "Morph (`morph`)", "url": "https://docs.litellm.ai/docs/providers/morph", diff --git a/requirements.txt b/requirements.txt index 45d7d529ea1..3a426d83e31 100644 --- a/requirements.txt +++ b/requirements.txt @@ -55,7 +55,7 @@ jinja2==3.1.6 # for prompt templates aiohttp==3.12.14 # for network calls aioboto3==13.4.0 # for async sagemaker calls tenacity==8.5.0 # for retrying requests, when litellm.num_retries set -pydantic==2.10.2 # proxy + openai req. +pydantic>=2.11,<3 # proxy + openai req. + mcp jsonschema==4.22.0 # validating json schema websockets==13.1.0 # for realtime API soundfile==0.12.1 # for audio file processing diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index fd98fa4fce1..cac245f8d5b 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -30,6 +30,8 @@ IGNORE_FUNCTIONS = [ "_fix_enum_empty_strings", # max depth set., "get_access_token", # max depth set., "_redact_base64", # max depth set. + "_contains_vision_content", # max depth set. + "_read_all_bytes", # max depth set. ] diff --git a/tests/image_gen_tests/test_fal_ai_image_generation.py b/tests/image_gen_tests/test_fal_ai_image_generation.py index 2db1aeefead..44e6c34cea0 100644 --- a/tests/image_gen_tests/test_fal_ai_image_generation.py +++ b/tests/image_gen_tests/test_fal_ai_image_generation.py @@ -1,6 +1,7 @@ import asyncio import os import sys +from unittest.mock import MagicMock, patch import pytest @@ -11,63 +12,83 @@ from litellm import aimage_generation @pytest.mark.parametrize( - "model", + "model,expected_endpoint", [ - "fal_ai/fal-ai/flux-pro/v1.1-ultra", - "fal_ai/fal-ai/flux-pro/v1.1", - "fal_ai/fal-ai/flux/schnell", - "fal_ai/fal-ai/bytedance/seedream/v3/text-to-image", - "fal_ai/fal-ai/bytedance/dreamina/v3.1/text-to-image", - "fal_ai/fal-ai/recraft/v3/text-to-image", - "fal_ai/fal-ai/ideogram/v3", - "fal_ai/bria/text-to-image/3.2", - "fal_ai/fal-ai/stable-diffusion-v35-medium" + ("fal_ai/fal-ai/flux-pro/v1.1-ultra", "fal-ai/flux-pro/v1.1-ultra"), + ("fal_ai/fal-ai/stable-diffusion-v35-medium", "fal-ai/stable-diffusion-v35-medium"), ], ) @pytest.mark.asyncio -async def test_fal_ai_image_generation_basic(model): +async def test_fal_ai_image_generation_basic(model, expected_endpoint): """ - Test basic image generation for various Fal AI models. + Test that fal_ai image generation constructs correct request body and URL. - Tests that each model can: - - Accept a basic text prompt - - Return a valid response with image data - - Handle the response properly through litellm + Validates: + - Correct API endpoint URL construction + - Proper request body format with prompt + - Correct Authorization header format """ - try: - litellm.set_verbose = True + captured_url = None + captured_json_data = None + captured_headers = None + + def capture_post_call(*args, **kwargs): + nonlocal captured_url, captured_json_data, captured_headers + + captured_url = args[0] if args else kwargs.get("url") + captured_json_data = kwargs.get("json") + captured_headers = kwargs.get("headers") + + # Mock response with fal.ai format + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "images": [ + { + "url": "https://example.com/generated-image.png", + "width": 1024, + "height": 768, + "content_type": "image/jpeg" + } + ], + "seed": 42 + } + + return mock_response + + with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post: + mock_post.side_effect = capture_post_call + + test_api_key = "test-fal-ai-key-12345" + test_prompt = "A cute baby sea otter" response = await aimage_generation( model=model, - prompt="A cute baby sea otter", + prompt=test_prompt, + api_key=test_api_key, ) - print(f"\nResponse from {model}:") - print(f" Number of images: {len(response.data)}") - print(f" First image URL: {response.data[0].url if response.data else 'None'}") + # Validate response + assert response is not None + assert hasattr(response, "data") + assert response.data is not None + assert len(response.data) > 0 - # Basic assertions - assert response is not None, f"Response should not be None for {model}" - assert hasattr(response, "data"), f"Response should have data attribute for {model}" - assert len(response.data) > 0, f"Response should have at least one image for {model}" + # Validate URL + assert captured_url is not None + assert "fal.run" in captured_url + assert expected_endpoint in captured_url + print(f"Validated URL: {captured_url}") - # Check that we got a URL or b64_json - first_image = response.data[0] - assert ( - first_image.url is not None or first_image.b64_json is not None - ), f"Image should have either url or b64_json for {model}" + # Validate headers + assert captured_headers is not None + assert "Authorization" in captured_headers + assert captured_headers["Authorization"] == f"Key {test_api_key}" + print(f"Validated headers: {captured_headers}") - print(f"✓ Test passed for {model}") - - except litellm.RateLimitError as e: - pytest.skip(f"Rate limit error for {model}: {str(e)}") - except litellm.ContentPolicyViolationError as e: - pytest.skip(f"Content policy violation for {model}: {str(e)}") - except litellm.InternalServerError as e: - pytest.skip(f"Internal server error for {model}: {str(e)}") - except Exception as e: - if "Your task failed as a result of our safety system" in str(e): - pytest.skip(f"Safety system rejection for {model}") - else: - pytest.fail(f"Test failed for {model}: {str(e)}") + # Validate request body + assert captured_json_data is not None + assert captured_json_data["prompt"] == test_prompt + print(f"Validated request body: {captured_json_data}") diff --git a/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts b/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts index 8971a84649e..f691389de53 100644 --- a/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts +++ b/tests/proxy_admin_ui_tests/e2e_ui_tests/login_to_ui.spec.ts @@ -26,7 +26,7 @@ test("admin login test", async ({ page }) => { await loginButton.click(); const tabs = [ "Virtual Keys", - "Test Key", + "Playground", "Models", "Usage", "Teams", diff --git a/tests/proxy_unit_tests/test_prompt_test_endpoint.py b/tests/proxy_unit_tests/test_prompt_test_endpoint.py new file mode 100644 index 00000000000..327f60e3d7b --- /dev/null +++ b/tests/proxy_unit_tests/test_prompt_test_endpoint.py @@ -0,0 +1,134 @@ +""" +Test /prompts/test endpoint for testing prompts before saving +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch +from fastapi import HTTPException + + +class TestPromptTestEndpoint: + """ + Tests the /prompts/test endpoint that allows testing prompts with variables + """ + + @pytest.mark.asyncio + async def test_parse_dotprompt_with_variables(self): + """ + Test that dotprompt content is parsed and variables are rendered correctly + """ + from litellm.integrations.dotprompt.prompt_manager import PromptManager + + dotprompt_content = """--- +model: gpt-4o +temperature: 0.7 +max_tokens: 100 +--- + +User: Hello {{name}}, how are you?""" + + # Parse the dotprompt + prompt_manager = PromptManager() + frontmatter, template_content = prompt_manager._parse_frontmatter( + content=dotprompt_content + ) + + assert frontmatter["model"] == "gpt-4o" + assert frontmatter["temperature"] == 0.7 + assert frontmatter["max_tokens"] == 100 + assert "{{name}}" in template_content + + # Render with variables + from jinja2 import Environment + + jinja_env = Environment( + variable_start_string="{{", + variable_end_string="}}", + ) + jinja_template = jinja_env.from_string(template_content) + rendered = jinja_template.render(name="World") + + assert "Hello World" in rendered + assert "{{name}}" not in rendered + + @pytest.mark.asyncio + async def test_convert_to_messages_format(self): + """ + Test that rendered prompt is converted to OpenAI messages format + """ + import re + + rendered_content = """System: You are a helpful assistant. + +User: Hello World, how are you?""" + + messages = [] + role_pattern = r"^(System|User|Assistant|Developer):\s*(.*?)(?=\n(?:System|User|Assistant|Developer):|$)" + matches = list( + re.finditer( + pattern=role_pattern, + string=rendered_content.strip(), + flags=re.MULTILINE | re.DOTALL, + ) + ) + + for match in matches: + role = match.group(1).lower() + content = match.group(2).strip() + + if role == "developer": + role = "system" + + if content: + messages.append({"role": role, "content": content}) + + assert len(messages) == 2 + assert messages[0]["role"] == "system" + assert "helpful assistant" in messages[0]["content"] + assert messages[1]["role"] == "user" + assert "Hello World" in messages[1]["content"] + + @pytest.mark.asyncio + async def test_single_message_without_role(self): + """ + Test that content without role markers is treated as a user message + """ + import re + + rendered_content = "Just a plain message without any role markers" + + messages = [] + role_pattern = r"^(System|User|Assistant|Developer):\s*(.*?)(?=\n(?:System|User|Assistant|Developer):|$)" + matches = list( + re.finditer( + pattern=role_pattern, + string=rendered_content.strip(), + flags=re.MULTILINE | re.DOTALL, + ) + ) + + if not matches: + messages.append({"role": "user", "content": rendered_content.strip()}) + + assert len(messages) == 1 + assert messages[0]["role"] == "user" + assert messages[0]["content"] == rendered_content + + @pytest.mark.asyncio + async def test_missing_model_raises_error(self): + """ + Test that missing model in frontmatter raises an error + """ + from litellm.integrations.dotprompt.prompt_manager import PromptManager + + dotprompt_content = """--- +temperature: 0.7 +--- + +User: Hello""" + + prompt_manager = PromptManager() + frontmatter, _ = prompt_manager._parse_frontmatter(content=dotprompt_content) + + model = frontmatter.get("model") + assert model is None diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index 04c582f6013..68931cebc9a 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -872,3 +872,201 @@ def test_initialize_specialized_endpoints(): for endpoint in specialized_endpoints: assert hasattr(router, endpoint) assert callable(getattr(router, endpoint)) + + +def test_initialize_vector_store_endpoints(): + """ + Test that _initialize_vector_store_endpoints correctly sets up vector store endpoints. + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test-model", + "api_key": "fake-api-key", + }, + } + ] + ) + + router._initialize_vector_store_endpoints() + + vector_store_endpoints = [ + "avector_store_search", + "avector_store_create", + "vector_store_search", + "vector_store_create", + ] + + for endpoint in vector_store_endpoints: + assert hasattr(router, endpoint) + assert callable(getattr(router, endpoint)) + + +def test_initialize_vector_store_file_endpoints(): + """ + Test that _initialize_vector_store_file_endpoints correctly sets up vector store file endpoints. + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test-model", + "api_key": "fake-api-key", + }, + } + ] + ) + + router._initialize_vector_store_file_endpoints() + + vector_store_file_endpoints = [ + "avector_store_file_create", + "vector_store_file_create", + "avector_store_file_list", + "vector_store_file_list", + "avector_store_file_retrieve", + "vector_store_file_retrieve", + "avector_store_file_content", + "vector_store_file_content", + "avector_store_file_update", + "vector_store_file_update", + "avector_store_file_delete", + "vector_store_file_delete", + ] + + for endpoint in vector_store_file_endpoints: + assert hasattr(router, endpoint) + assert callable(getattr(router, endpoint)) + + +def test_initialize_google_genai_endpoints(): + """ + Test that _initialize_google_genai_endpoints correctly sets up Google GenAI endpoints. + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test-model", + "api_key": "fake-api-key", + }, + } + ] + ) + + router._initialize_google_genai_endpoints() + + google_genai_endpoints = [ + "agenerate_content", + "generate_content", + "agenerate_content_stream", + "generate_content_stream", + ] + + for endpoint in google_genai_endpoints: + assert hasattr(router, endpoint) + assert callable(getattr(router, endpoint)) + + +def test_initialize_ocr_search_endpoints(): + """ + Test that _initialize_ocr_search_endpoints correctly sets up OCR and search endpoints. + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test-model", + "api_key": "fake-api-key", + }, + } + ] + ) + + router._initialize_ocr_search_endpoints() + + ocr_search_endpoints = [ + "aocr", + "ocr", + "asearch", + "search", + ] + + for endpoint in ocr_search_endpoints: + assert hasattr(router, endpoint) + assert callable(getattr(router, endpoint)) + + +def test_initialize_video_endpoints(): + """ + Test that _initialize_video_endpoints correctly sets up video endpoints. + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test-model", + "api_key": "fake-api-key", + }, + } + ] + ) + + router._initialize_video_endpoints() + + video_endpoints = [ + "avideo_generation", + "video_generation", + "avideo_list", + "video_list", + "avideo_status", + "video_status", + "avideo_content", + "video_content", + "avideo_remix", + "video_remix", + ] + + for endpoint in video_endpoints: + assert hasattr(router, endpoint) + assert callable(getattr(router, endpoint)) + + +def test_initialize_container_endpoints(): + """ + Test that _initialize_container_endpoints correctly sets up container endpoints. + """ + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/test-model", + "api_key": "fake-api-key", + }, + } + ] + ) + + router._initialize_container_endpoints() + + container_endpoints = [ + "acreate_container", + "create_container", + "alist_containers", + "list_containers", + "aretrieve_container", + "retrieve_container", + "adelete_container", + "delete_container", + ] + + for endpoint in container_endpoints: + assert hasattr(router, endpoint) + assert callable(getattr(router, endpoint)) diff --git a/tests/search_tests/test_search_tool_name_filtering.py b/tests/search_tests/test_search_tool_name_filtering.py new file mode 100644 index 00000000000..2cfe1177b43 --- /dev/null +++ b/tests/search_tests/test_search_tool_name_filtering.py @@ -0,0 +1,50 @@ +""" +Test that search_tool_name is properly filtered out from search requests. + +The search_tool_name parameter is used internally by LiteLLM to identify +which search tool configuration to use, but should not be sent to external +search provider APIs. +""" +import sys +import os + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.types.utils import all_litellm_params +from litellm.utils import filter_out_litellm_params + + +def test_search_tool_name_in_all_litellm_params(): + """ + Test that search_tool_name is in all_litellm_params. + + If missing, it gets passed to provider APIs causing errors. + """ + assert "search_tool_name" in all_litellm_params + + +def test_filter_out_search_tool_name(): + """ + Test that filter_out_litellm_params correctly filters search_tool_name. + """ + kwargs = { + "query": "latest ai developments", + "max_results": 5, + "scrapeOptions": {"formats": ["markdown"]}, + "search_tool_name": "firecrawl-search", + "metadata": {"user": "test"}, + "litellm_call_id": "test-123" + } + + filtered = filter_out_litellm_params(kwargs=kwargs) + + assert "search_tool_name" not in filtered + assert "metadata" not in filtered + assert "litellm_call_id" not in filtered + + assert "query" in filtered + assert "max_results" in filtered + assert "scrapeOptions" in filtered + assert filtered["query"] == "latest ai developments" + assert filtered["max_results"] == 5 + diff --git a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py index 9d7e2b3d382..6aa65db5ea4 100644 --- a/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py +++ b/tests/test_litellm/integrations/dotprompt/test_prompt_manager.py @@ -553,6 +553,7 @@ async def test_dotprompt_auto_detection_with_model_only(): without needing to specify model="dotprompt/gpt-4". """ from litellm.integrations.dotprompt import DotpromptManager + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler prompt_dir = Path(__file__).parent dotprompt_manager = DotpromptManager(prompt_directory=str(prompt_dir)) @@ -563,49 +564,26 @@ async def test_dotprompt_auto_detection_with_model_only(): try: # Mock the HTTP handler to avoid actual API calls - with patch("litellm.llms.custom_httpx.llm_http_handler.AsyncHTTPHandler.post") as mock_post: - mock_response_data = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="Hello!"), - index=0, - finish_reason="stop", - ) - ] - ).model_dump() - - # Create a proper mock response - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.text = json.dumps(mock_response_data) - mock_response.headers = {"Content-Type": "application/json"} - mock_response.json.return_value = mock_response_data - - mock_post.return_value = mock_response - + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=MagicMock()) as mock_post: # Call with model="gpt-4" (no "dotprompt/" prefix) and prompt_id await litellm.acompletion( model="gpt-4", prompt_id="chat_prompt", prompt_variables={"user_message": "Hello world"}, messages=[{"role": "user", "content": "This will be ignored"}], + client=client, ) mock_post.assert_called_once() - # Get request body from the call (it's passed as 'data' parameter as JSON string) - data_str = mock_post.call_args.kwargs.get("data", "{}") - request_body = json.loads(data_str) - - print(f"Request body: {json.dumps(request_body, indent=2)}") + # Get request body from the call + request_body = mock_post.call_args.kwargs.get("json") or json.loads(mock_post.call_args.kwargs.get("data", "{}")) # Verify the prompt was auto-detected and used # The chat_prompt.prompt has metadata: model: gpt-4, temperature: 0.7, max_tokens: 150 assert request_body["model"] == "gpt-4" - # Note: OpenAI API might strip out temperature/max_tokens if they're not in the request - # The key test is that the messages were transformed - # Verify the messages were transformed using the prompt template # chat_prompt template: "User: {{user_message}}" messages = request_body["messages"] @@ -614,7 +592,6 @@ async def test_dotprompt_auto_detection_with_model_only(): # The first message should be from the prompt template with the variable substituted # Template is: "User: {{user_message}}" with user_message="Hello world" first_message_content = messages[0]["content"] - print(f"First message content: {first_message_content}") assert "Hello world" in first_message_content finally: @@ -639,41 +616,20 @@ async def test_dotprompt_with_prompt_version(): litellm.callbacks = [dotprompt_manager] try: - # Mock the HTTP handler to avoid actual API calls - with patch("litellm.llms.custom_httpx.llm_http_handler.AsyncHTTPHandler.post") as mock_post: - mock_response_data = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="Hello!"), - index=0, - finish_reason="stop", - ) - ] - ).model_dump() - - # Create a proper mock response - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.text = json.dumps(mock_response_data) - mock_response.headers = {"Content-Type": "application/json"} - mock_response.json.return_value = mock_response_data - - mock_post.return_value = mock_response - - # Test version 1 + # Test version 1 + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=MagicMock()) as mock_post: await litellm.acompletion( model="gpt-3.5-turbo", prompt_id="chat_prompt", prompt_version=1, prompt_variables={"user_message": "Test v1"}, messages=[], + client=client, ) - assert mock_post.call_count >= 1 - data_str = mock_post.call_args.kwargs.get("data", "{}") - request_body = json.loads(data_str) - - print(f"Version 1 request body: {json.dumps(request_body, indent=2)}") + mock_post.assert_called_once() + request_body = mock_post.call_args.kwargs.get("json") or json.loads(mock_post.call_args.kwargs.get("data", "{}")) # Verify version 1 prompt was used # chat_prompt.v1.prompt has: model: gpt-3.5-turbo, temperature: 0.5, max_tokens: 100 @@ -683,47 +639,23 @@ async def test_dotprompt_with_prompt_version(): messages = request_body["messages"] assert len(messages) >= 1 first_message_content = messages[0]["content"] - print(f"Version 1 message: {first_message_content}") assert "Version 1:" in first_message_content assert "Test v1" in first_message_content - - # Reset mock for version 2 test - mock_post.reset_mock() # Test version 2 - with patch("litellm.llms.custom_httpx.llm_http_handler.AsyncHTTPHandler.post") as mock_post: - mock_response_data = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="Hello!"), - index=0, - finish_reason="stop", - ) - ] - ).model_dump() - - # Create a proper mock response - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.text = json.dumps(mock_response_data) - mock_response.headers = {"Content-Type": "application/json"} - mock_response.json.return_value = mock_response_data - - mock_post.return_value = mock_response - + client = AsyncHTTPHandler() + with patch.object(client, "post", return_value=MagicMock()) as mock_post: await litellm.acompletion( model="gpt-4", prompt_id="chat_prompt", prompt_version=2, prompt_variables={"user_message": "Test v2"}, messages=[], + client=client, ) mock_post.assert_called_once() - data_str = mock_post.call_args.kwargs.get("data", "{}") - request_body = json.loads(data_str) - - print(f"Version 2 request body: {json.dumps(request_body, indent=2)}") + request_body = mock_post.call_args.kwargs.get("json") or json.loads(mock_post.call_args.kwargs.get("data", "{}")) # Verify version 2 prompt was used # chat_prompt.v2.prompt has: model: gpt-4, temperature: 0.9, max_tokens: 200 @@ -733,7 +665,6 @@ async def test_dotprompt_with_prompt_version(): messages = request_body["messages"] assert len(messages) >= 1 first_message_content = messages[0]["content"] - print(f"Version 2 message: {first_message_content}") assert "Version 2:" in first_message_content assert "Test v2" in first_message_content diff --git a/tests/test_litellm/llms/azure/test_azure_common_utils.py b/tests/test_litellm/llms/azure/test_azure_common_utils.py index 1c81f12e0fe..0e5ddb391ee 100644 --- a/tests/test_litellm/llms/azure/test_azure_common_utils.py +++ b/tests/test_litellm/llms/azure/test_azure_common_utils.py @@ -438,6 +438,8 @@ def test_select_azure_base_url_called(setup_mocks): "allm_passthrough_route", "llm_passthrough_route", "asearch", + "avector_store_create", + "avector_store_search", ] ], ) diff --git a/tests/test_litellm/llms/docker_model_runner/test_docker_model_runner_chat_transformation.py b/tests/test_litellm/llms/docker_model_runner/test_docker_model_runner_chat_transformation.py new file mode 100644 index 00000000000..9cd76c3ef6a --- /dev/null +++ b/tests/test_litellm/llms/docker_model_runner/test_docker_model_runner_chat_transformation.py @@ -0,0 +1,172 @@ +""" +Unit tests for Docker Model Runner configuration. + +This test validates that litellm.completion correctly routes requests to Docker Model Runner +with the proper URL structure and request body. +""" + +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) + +import json +from unittest.mock import Mock, patch + +import pytest + +import litellm +from litellm import completion + + +class TestDockerModelRunnerIntegration: + """Integration test for Docker Model Runner""" + + @pytest.mark.asyncio + async def test_completion_hits_correct_url_and_body(self): + """ + Test that litellm.completion with docker_model_runner provider: + 1. Hits the correct URL: {api_base}/v1/chat/completions where api_base includes engine path + 2. Sends the correct request body with messages and parameters + """ + with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post: + # Mock the response + mock_response = Mock() + mock_response.json.return_value = { + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "llama-3.1", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 20, + "total_tokens": 30 + } + } + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_post.return_value = mock_response + + # Make the completion call with engine in api_base + response = completion( + model="docker_model_runner/llama-3.1", + messages=[{"role": "user", "content": "Hello, how are you?"}], + api_base="http://localhost:22088/engines/llama.cpp", + temperature=0.7, + max_tokens=100 + ) + + # Verify the URL was correct + assert mock_post.called + call_args = mock_post.call_args + url = call_args[1]["url"] + print("URL For request", url) + print("request body for request", json.dumps(call_args[1]["data"], indent=4)) + + # Should hit {api_base}/v1/chat/completions where api_base includes engine + assert "/engines/llama.cpp/v1/chat/completions" in url + assert "http://localhost:22088" in url + + # Verify the request body + request_data = call_args[1]["data"] + if isinstance(request_data, str): + request_data = json.loads(request_data) + + # Check messages + assert "messages" in request_data + assert len(request_data["messages"]) == 1 + assert request_data["messages"][0]["role"] == "user" + assert request_data["messages"][0]["content"] == "Hello, how are you?" + + # Check parameters + assert request_data["temperature"] == 0.7 + assert request_data["max_tokens"] == 100 + + # Verify response + assert response.choices[0].message.content == "Hello! How can I help you today?" + + @pytest.mark.asyncio + async def test_completion_with_custom_engine_and_host(self): + """ + Test that litellm.completion works with custom engine and host: + 1. Uses model-runner.docker.internal as host + 2. Specifies a different engine in the api_base + 3. Model name is sent in the request body + """ + with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post: + # Mock the response + mock_response = Mock() + mock_response.json.return_value = { + "id": "chatcmpl-456", + "object": "chat.completion", + "created": 1677652288, + "model": "mistral-7b", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Bonjour! How can I assist you?" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 15, + "completion_tokens": 25, + "total_tokens": 40 + } + } + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_post.return_value = mock_response + + # Make the completion call with custom engine and host + response = completion( + model="docker_model_runner/mistral-7b", + messages=[{"role": "user", "content": "Hello!"}], + api_base="http://model-runner.docker.internal/engines/custom-engine", + temperature=0.5, + max_tokens=200 + ) + + # Verify the URL was correct + assert mock_post.called + call_args = mock_post.call_args + url = call_args[1]["url"] + print("URL For request", url) + print("request body for request", json.dumps(call_args[1]["data"], indent=4)) + + # Should hit the custom host and engine + assert "model-runner.docker.internal" in url + assert "/engines/custom-engine/v1/chat/completions" in url + + # Verify the request body contains the model name + request_data = call_args[1]["data"] + if isinstance(request_data, str): + request_data = json.loads(request_data) + + # Check that model name is in the request body + assert request_data["model"] == "mistral-7b" + + # Check messages + assert "messages" in request_data + assert len(request_data["messages"]) == 1 + assert request_data["messages"][0]["role"] == "user" + assert request_data["messages"][0]["content"] == "Hello!" + + # Check parameters + assert request_data["temperature"] == 0.5 + assert request_data["max_tokens"] == 200 + + # Verify response + assert response.choices[0].message.content == "Bonjour! How can I assist you?" + diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py b/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py new file mode 100644 index 00000000000..68c0f3bdbfc --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py @@ -0,0 +1,264 @@ +""" +Tests for embedding thought signatures in tool call IDs for OpenAI client compatibility. + +When using OpenAI clients (instead of LiteLLM SDK), provider_specific_fields are not preserved. +This test suite validates that thought signatures can be embedded in tool call IDs and extracted +when converting back to Gemini format. +""" + +import pytest +from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, +) +from litellm.litellm_core_utils.prompt_templates.factory import ( + THOUGHT_SIGNATURE_SEPARATOR, + convert_to_gemini_tool_call_invoke, + _encode_tool_call_id_with_signature, + _get_thought_signature_from_tool, +) +from litellm.types.llms.vertex_ai import HttpxPartType + + +def test_encode_decode_tool_call_id_with_signature(): + """Test that thought signatures can be encoded in and decoded from tool call IDs""" + base_id = "call_abc123" + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Test encoding + encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature) + assert THOUGHT_SIGNATURE_SEPARATOR in encoded_id + assert encoded_id.startswith(base_id) + + # Test decoding using factory function with realistic tool call structure + tool = { + "id": encoded_id, + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + } + + extracted_signature = _get_thought_signature_from_tool(tool) + assert extracted_signature == test_signature + + # Verify base ID is preserved + decoded_base_id = encoded_id.split(THOUGHT_SIGNATURE_SEPARATOR)[0] + assert decoded_base_id == base_id + + +def test_encode_tool_call_id_without_signature(): + """Test that IDs without signatures are returned unchanged""" + base_id = "call_abc123def456" + + # Encode without signature + encoded_id = _encode_tool_call_id_with_signature(base_id, None) + assert encoded_id == base_id + assert THOUGHT_SIGNATURE_SEPARATOR not in encoded_id + + # Decode ID without signature using factory function + tool_obj = {"id": base_id, "type": "function"} + decoded_signature = _get_thought_signature_from_tool(tool_obj) + assert decoded_signature is None + + +def test_tool_call_id_includes_signature_in_response(): + """Test that tool call IDs in responses include embedded thought signatures""" + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + parts_with_signature = [ + HttpxPartType( + functionCall={ + "name": "get_current_temperature", + "args": {"location": "Paris"}, + }, + thoughtSignature=test_signature, + ) + ] + + function, tools, _ = VertexGeminiConfig._transform_parts( + parts=parts_with_signature, + cumulative_tool_call_idx=0, + is_function_call=False, + ) + + # Verify tool call ID includes thought signature + assert tools is not None + assert len(tools) == 1 + tool_call_id = tools[0]["id"] + assert THOUGHT_SIGNATURE_SEPARATOR in tool_call_id + + # Verify we can decode it using the factory function + tool_obj = {"id": tool_call_id, "type": "function"} + decoded_sig = _get_thought_signature_from_tool(tool_obj) + assert decoded_sig == test_signature + + +def test_get_thought_signature_backward_compatibility(): + """Test that provider_specific_fields still works (backward compatibility)""" + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Test with provider_specific_fields (LiteLLM SDK scenario) + tool = { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + "provider_specific_fields": {"thought_signature": test_signature}, + } + + extracted_signature = _get_thought_signature_from_tool(tool) + assert extracted_signature == test_signature + + +def test_get_thought_signature_prioritizes_provider_fields(): + """Test that provider_specific_fields takes priority over tool call ID""" + signature_in_fields = "signature_from_fields" + signature_in_id = "signature_from_id" + + encoded_id = _encode_tool_call_id_with_signature("call_abc123", signature_in_id) + + tool = { + "id": encoded_id, + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + "provider_specific_fields": {"thought_signature": signature_in_fields}, + } + + extracted_signature = _get_thought_signature_from_tool(tool) + # Should prioritize provider_specific_fields + assert extracted_signature == signature_in_fields + + +def test_convert_to_gemini_with_embedded_signature(): + """Test that convert_to_gemini_tool_call_invoke extracts signatures from tool call IDs""" + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Create tool call ID with embedded signature (as OpenAI client would send) + base_id = "call_abc123" + encoded_id = _encode_tool_call_id_with_signature(base_id, test_signature) + + # Assistant message as sent by OpenAI client (no provider_specific_fields) + assistant_message = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": encoded_id, # ID has signature embedded + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + } + ], + } + + gemini_parts = convert_to_gemini_tool_call_invoke(assistant_message) + + # Verify thought signature is extracted and sent to Gemini + assert len(gemini_parts) == 1 + assert "function_call" in gemini_parts[0] + assert "thoughtSignature" in gemini_parts[0] + assert gemini_parts[0]["thoughtSignature"] == test_signature + + +def test_openai_client_e2e_flow(): + """ + End-to-end test simulating OpenAI client usage: + 1. LiteLLM receives response from Gemini with thought signature + 2. LiteLLM embeds signature in tool call ID + 3. OpenAI client sends message back with same tool call ID + 4. LiteLLM extracts signature from ID and sends to Gemini + """ + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Step 1: Gemini returns function call with thought signature + gemini_parts = [ + HttpxPartType( + functionCall={ + "name": "get_current_temperature", + "args": {"location": "Paris"}, + }, + thoughtSignature=test_signature, + ) + ] + + # Step 2: LiteLLM transforms to OpenAI format with embedded signature + function, tools, _ = VertexGeminiConfig._transform_parts( + parts=gemini_parts, + cumulative_tool_call_idx=0, + is_function_call=False, + ) + + assert tools is not None + assert len(tools) == 1 + tool_call_id = tools[0]["id"] + assert THOUGHT_SIGNATURE_SEPARATOR in tool_call_id + + # Step 3: OpenAI client sends back assistant message (preserves tool_call_id) + openai_assistant_message = { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": tool_call_id, # Preserved from response + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + } + ], + } + + # Step 4: LiteLLM converts back to Gemini format, extracting signature + gemini_parts_converted = convert_to_gemini_tool_call_invoke( + openai_assistant_message + ) + + # Verify signature is preserved through the round trip + assert len(gemini_parts_converted) == 1 + assert "thoughtSignature" in gemini_parts_converted[0] + assert gemini_parts_converted[0]["thoughtSignature"] == test_signature + + +def test_parallel_tool_calls_with_signatures(): + """Test that parallel tool calls preserve signatures correctly""" + signature1 = "signature_for_first_call" + # Only first call has signature (Gemini behavior for parallel calls) + + gemini_parts = [ + HttpxPartType( + functionCall={"name": "get_temperature", "args": {"location": "Paris"}}, + thoughtSignature=signature1, + ), + HttpxPartType( + functionCall={"name": "get_temperature", "args": {"location": "London"}}, + # No signature for second parallel call + ), + ] + + function, tools, _ = VertexGeminiConfig._transform_parts( + parts=gemini_parts, + cumulative_tool_call_idx=0, + is_function_call=False, + ) + + assert tools is not None + assert len(tools) == 2 + + # First tool call has signature in ID + assert THOUGHT_SIGNATURE_SEPARATOR in tools[0]["id"] + sig1 = _get_thought_signature_from_tool({"id": tools[0]["id"], "type": "function"}) + assert sig1 == signature1 + + # Second tool call has no signature in ID + assert THOUGHT_SIGNATURE_SEPARATOR not in tools[1]["id"] + sig2 = _get_thought_signature_from_tool({"id": tools[1]["id"], "type": "function"}) + assert sig2 is None diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py index a9d87398217..8c88b22f60e 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py @@ -129,6 +129,24 @@ class TestToolPermissionGuardrail: assert rule_id is None assert "default" in (msg or "") + def test_check_tool_permission_custom_template(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="custom-template", + rules=self.test_rules, + default_action="deny", + violation_message_template="custom {tool_name} {rule_id} :: {default_message}", + ) + + _, rule_id, message = guardrail._check_tool_permission("Read") + assert rule_id == "deny_read" + assert message.startswith("custom Read deny_read") + assert "Tool 'Read' denied" in message + + _, rule_id, message = guardrail._check_tool_permission("UnknownTool") + assert rule_id is None + assert message.startswith("custom UnknownTool None") + assert "Tool 'UnknownTool' denied by default action" in message + def test_extract_tool_calls_openai_format(self): tool_call = { "id": "call_123", @@ -224,6 +242,39 @@ class TestToolPermissionGuardrail: ) assert excinfo.value.status_code == 400 + @pytest.mark.asyncio + async def test_async_pre_call_hook_uses_custom_template(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="custom-template", + rules=self.test_rules, + default_action="deny", + on_disallowed_action="block", + violation_message_template="blocked {tool_name} by policy", + ) + + data = { + "tools": [ + {"type": "function", "function": {"name": "Read"}}, + ] + } + user_api_key_dict = UserAPIKeyAuth() + cache = DualCache(default_in_memory_ttl=1) + + with patch.object(guardrail, "should_run_guardrail", return_value=True): + with pytest.raises(HTTPException) as excinfo: + await guardrail.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="completion", + ) + + assert excinfo.value.status_code == 400 + assert ( + excinfo.value.detail.get("detection_message") + == "blocked Read by policy" + ) + @pytest.mark.asyncio async def test_async_pre_call_hook_rewrite_mode(self): guardrail = ToolPermissionGuardrail( 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 266056bcdd2..a112ae046ab 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 @@ -690,3 +690,125 @@ async def test_check_duplicate_user_email_case_insensitive(mocker): await _check_duplicate_user_email( None, mock_prisma_client ) # Should not raise exception + + +def test_process_keys_for_user_info_filters_dashboard_keys(monkeypatch): + """ + Test that _process_keys_for_user_info filters out keys with team_id='litellm-dashboard' + + UI session tokens (team_id='litellm-dashboard') should be excluded from user info responses + to prevent confusion, as these are automatically created during dashboard login. + """ + from unittest.mock import MagicMock + + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _process_keys_for_user_info, + ) + + # Create mock keys with different team_ids + mock_key_dashboard = MagicMock() + mock_key_dashboard.model_dump.return_value = { + "token": "sk-dashboard-token", + "team_id": UI_SESSION_TOKEN_TEAM_ID, + "user_id": "test-user", + "key_alias": "dashboard-session-key", + } + + mock_key_regular = MagicMock() + mock_key_regular.model_dump.return_value = { + "token": "sk-regular-token", + "team_id": "regular-team", + "user_id": "test-user", + "key_alias": "regular-key", + } + + mock_key_no_team = MagicMock() + mock_key_no_team.model_dump.return_value = { + "token": "sk-no-team-token", + "team_id": None, + "user_id": "test-user", + "key_alias": "no-team-key", + } + + keys = [mock_key_dashboard, mock_key_regular, mock_key_no_team] + + # Mock general_settings and litellm_master_key_hash (they're imported from proxy_server) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {}, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.litellm_master_key_hash", + "different-hash", + ) + + # Call the function + result = _process_keys_for_user_info(keys=keys, all_teams=None) + + # Verify that dashboard key is filtered out + assert len(result) == 2, "Should return 2 keys (dashboard key filtered out)" + + # Verify dashboard key is not in results + result_team_ids = [key.get("team_id") for key in result] + assert UI_SESSION_TOKEN_TEAM_ID not in result_team_ids, "Dashboard key should be filtered out" + + # Verify regular keys are included + assert "regular-team" in result_team_ids, "Regular team key should be included" + assert None in result_team_ids, "No-team key should be included" + + # Verify the correct keys are returned + result_tokens = [key.get("token") for key in result] + assert "sk-regular-token" in result_tokens, "Regular key should be included" + assert "sk-no-team-token" in result_tokens, "No-team key should be included" + assert "sk-dashboard-token" not in result_tokens, "Dashboard key should not be included" + + +def test_process_keys_for_user_info_handles_none_keys(monkeypatch): + """ + Test that _process_keys_for_user_info handles None keys gracefully + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _process_keys_for_user_info, + ) + + # Mock general_settings and litellm_master_key_hash (they're imported from proxy_server) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {}, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.litellm_master_key_hash", + "different-hash", + ) + + # Call with None keys + result = _process_keys_for_user_info(keys=None, all_teams=None) + + # Should return empty list + assert result == [], "Should return empty list when keys is None" + + +def test_process_keys_for_user_info_handles_empty_keys(monkeypatch): + """ + Test that _process_keys_for_user_info handles empty keys list + """ + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _process_keys_for_user_info, + ) + + # Mock general_settings and litellm_master_key_hash (they're imported from proxy_server) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {}, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.litellm_master_key_hash", + "different-hash", + ) + + # Call with empty list + result = _process_keys_for_user_info(keys=[], all_teams=None) + + # Should return empty list + assert result == [], "Should return empty list when keys is empty" diff --git a/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py b/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py new file mode 100644 index 00000000000..6c2e5fa7667 --- /dev/null +++ b/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py @@ -0,0 +1,307 @@ +""" +Test prompt endpoints for version filtering and history +""" + +from unittest.mock import MagicMock + +import pytest + +from litellm.types.prompts.init_prompts import ( + PromptInfo, + PromptLiteLLMParams, + PromptSpec, +) + + +class TestPromptVersioning: + """ + Test prompt versioning functionality + """ + + def test_get_latest_prompt_versions(self): + """ + Test that get_latest_prompt_versions returns only the latest version of each prompt + """ + from litellm.proxy.prompts.prompt_endpoints import get_latest_prompt_versions + + # Create mock prompts with different versions + prompts = [ + PromptSpec( + prompt_id="jack.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v1 content" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + PromptSpec( + prompt_id="jack.v2", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v2 content" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + PromptSpec( + prompt_id="jane.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jane", + prompt_integration="dotprompt", + dotprompt_content="jane v1" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + PromptSpec( + prompt_id="jack.v3", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v3 content" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + ] + + # Get latest versions + latest = get_latest_prompt_versions(prompts=prompts) + + # Should return 2 prompts (jack.v3 and jane.v1) + assert len(latest) == 2 + + # Find jack and jane in results + jack_prompt = next((p for p in latest if "jack" in p.prompt_id), None) + jane_prompt = next((p for p in latest if "jane" in p.prompt_id), None) + + assert jack_prompt is not None + assert jack_prompt.prompt_id == "jack.v3" + assert jack_prompt.litellm_params.dotprompt_content == "v3 content" + + assert jane_prompt is not None + assert jane_prompt.prompt_id == "jane.v1" + + def test_get_version_number(self): + """ + Test that get_version_number correctly extracts version numbers + """ + from litellm.proxy.prompts.prompt_endpoints import get_version_number + + assert get_version_number(prompt_id="jack.v1") == 1 + assert get_version_number(prompt_id="jack.v2") == 2 + assert get_version_number(prompt_id="jack.v10") == 10 + assert get_version_number(prompt_id="jack") == 1 + assert get_version_number(prompt_id="jack.vinvalid") == 1 + + def test_get_base_prompt_id(self): + """ + Test that get_base_prompt_id correctly strips version suffixes + """ + from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id + + assert get_base_prompt_id(prompt_id="jack.v1") == "jack" + assert get_base_prompt_id(prompt_id="jack.v2") == "jack" + assert get_base_prompt_id(prompt_id="jack") == "jack" + assert get_base_prompt_id(prompt_id="my_prompt.v10") == "my_prompt" + + def test_get_latest_version_prompt_id(self): + """ + Test that get_latest_version_prompt_id returns the highest version + """ + from litellm.proxy.prompts.prompt_endpoints import get_latest_version_prompt_id + + # Mock prompt IDs dictionary + all_prompt_ids = { + "jack.v1": {}, + "jack.v2": {}, + "jack.v3": {}, + "jane.v1": {}, + "simple_prompt": {}, + } + + # Test with base prompt ID - should return latest version + assert get_latest_version_prompt_id( + prompt_id="jack", + all_prompt_ids=all_prompt_ids + ) == "jack.v3" + + # Test with versioned prompt ID - should still return latest version + assert get_latest_version_prompt_id( + prompt_id="jack.v1", + all_prompt_ids=all_prompt_ids + ) == "jack.v3" + + # Test with single version + assert get_latest_version_prompt_id( + prompt_id="jane", + all_prompt_ids=all_prompt_ids + ) == "jane.v1" + + # Test with non-versioned prompt + assert get_latest_version_prompt_id( + prompt_id="simple_prompt", + all_prompt_ids=all_prompt_ids + ) == "simple_prompt" + + # Test with non-existent prompt + assert get_latest_version_prompt_id( + prompt_id="nonexistent", + all_prompt_ids=all_prompt_ids + ) == "nonexistent" + + def test_construct_versioned_prompt_id(self): + """ + Test that construct_versioned_prompt_id correctly builds versioned IDs + """ + from litellm.proxy.prompts.prompt_endpoints import construct_versioned_prompt_id + + # Test with base prompt ID and version + assert construct_versioned_prompt_id( + prompt_id="jack_success", + version=4 + ) == "jack_success.v4" + + # Test with None version - should return base ID unchanged + assert construct_versioned_prompt_id( + prompt_id="jack_success", + version=None + ) == "jack_success" + + # Test with existing versioned ID - should replace version + assert construct_versioned_prompt_id( + prompt_id="jack_success.v2", + version=4 + ) == "jack_success.v4" + + # Test with hyphenated prompt ID + assert construct_versioned_prompt_id( + prompt_id="my-prompt", + version=1 + ) == "my-prompt.v1" + + # Test with double-digit version + assert construct_versioned_prompt_id( + prompt_id="test_prompt", + version=10 + ) == "test_prompt.v10" + + +class TestPromptVersionsEndpoint: + """ + Test the /prompts/{prompt_id}/versions endpoint + """ + + @pytest.mark.asyncio + async def test_get_prompt_versions_returns_all_versions(self): + """ + Test that get_prompt_versions returns all versions of a prompt sorted by version number + """ + from unittest.mock import MagicMock, patch + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.prompts.prompt_endpoints import get_prompt_versions + + # Mock user with admin role + mock_user = UserAPIKeyAuth( + api_key="test_key", + user_role=LitellmUserRoles.PROXY_ADMIN + ) + + # Create mock prompt registry with multiple versions + mock_prompts = { + "jack.v1": PromptSpec( + prompt_id="jack.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v1" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + "jack.v2": PromptSpec( + prompt_id="jack.v2", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v2" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + "jack.v3": PromptSpec( + prompt_id="jack.v3", + litellm_params=PromptLiteLLMParams( + prompt_id="jack", + prompt_integration="dotprompt", + dotprompt_content="v3" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + "jane.v1": PromptSpec( + prompt_id="jane.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="jane", + prompt_integration="dotprompt", + dotprompt_content="jane" + ), + prompt_info=PromptInfo(prompt_type="db"), + ), + } + + # Mock the IN_MEMORY_PROMPT_REGISTRY at the import location + with patch("litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY") as mock_registry: + mock_registry.IN_MEMORY_PROMPTS = mock_prompts + + # Test with base prompt ID + response = await get_prompt_versions( + prompt_id="jack", + user_api_key_dict=mock_user + ) + + # Should return 3 versions of jack, sorted newest first + assert len(response.prompts) == 3 + assert response.prompts[0].prompt_id == "jack" + assert response.prompts[0].version == 3 + assert response.prompts[1].prompt_id == "jack" + assert response.prompts[1].version == 2 + assert response.prompts[2].prompt_id == "jack" + assert response.prompts[2].version == 1 + + # Test with versioned prompt ID (should strip version) + response = await get_prompt_versions( + prompt_id="jack.v1", + user_api_key_dict=mock_user + ) + + assert len(response.prompts) == 3 + assert response.prompts[0].prompt_id == "jack" + assert response.prompts[0].version == 3 + + @pytest.mark.asyncio + async def test_get_prompt_versions_not_found(self): + """ + Test that get_prompt_versions raises 404 when prompt doesn't exist + """ + from unittest.mock import patch + + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.prompts.prompt_endpoints import get_prompt_versions + + mock_user = UserAPIKeyAuth( + api_key="test_key", + user_role=LitellmUserRoles.PROXY_ADMIN + ) + + with patch("litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY") as mock_registry: + mock_registry.IN_MEMORY_PROMPTS = {} + + with pytest.raises(HTTPException) as exc_info: + await get_prompt_versions( + prompt_id="nonexistent", + user_api_key_dict=mock_user + ) + + assert exc_info.value.status_code == 404 + assert "No versions found" in exc_info.value.detail + diff --git a/tests/test_litellm/test_video_generation.py b/tests/test_litellm/test_video_generation.py index be7400d9b7f..1a66cd29c71 100644 --- a/tests/test_litellm/test_video_generation.py +++ b/tests/test_litellm/test_video_generation.py @@ -14,6 +14,7 @@ import litellm from litellm.types.videos.main import VideoObject, VideoResponse from litellm.videos.main import video_generation, avideo_generation, video_status, avideo_status from litellm.llms.openai.videos.transformation import OpenAIVideoConfig +from litellm.llms.gemini.videos.transformation import GeminiVideoConfig from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.cost_calculator import default_video_cost_calculator from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging @@ -813,5 +814,12 @@ def test_openai_video_config_has_async_transform(): cfg = OpenAIVideoConfig() assert callable(getattr(cfg, "async_transform_video_content_response", None)) + +def test_gemini_video_config_has_async_transform(): + """Ensure GeminiVideoConfig exposes async_transform_video_content_response at runtime.""" + cfg = GeminiVideoConfig() + assert callable(getattr(cfg, "async_transform_video_content_response", None)) + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx index 8eb985e7ba8..4b4f1ab676b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx @@ -406,7 +406,7 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] code={`{ "mcpServers": { "Zapier_MCP": { - "server_url": "${proxyBaseUrl}/mcp", + "url": "${proxyBaseUrl}/mcp", "headers": { "x-litellm-api-key": "Bearer YOUR_LITELLM_API_KEY", "x-mcp-servers": ["Zapier_MCP,dev"] diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index 9e402b85c23..8fad64e094e 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -132,3 +132,188 @@ describe("daily activity helpers", () => { expect(urlWithTeams.searchParams.get("exclude_team_ids")).toBe("litellm-dashboard"); }); }); + +describe("UI config and public endpoints", () => { + const originalFetch = global.fetch; + + const setupMockFetch = (responses: Array<{ url: string; data: any }>) => { + const mockFetch = vi.fn().mockImplementation((url: string) => { + const response = responses.find((r) => url.includes(r.url)); + if (response) { + return Promise.resolve({ + ok: true, + json: vi.fn().mockResolvedValue(response.data), + } as any); + } + return Promise.resolve({ + ok: true, + json: vi.fn().mockResolvedValue({}), + } as any); + }); + global.fetch = mockFetch as any; + return mockFetch; + }; + + beforeEach(() => { + vi.clearAllMocks(); + }); + + afterEach(() => { + global.fetch = originalFetch; + }); + + it("should use proxyBaseURL and server_root_path for /public/providers/fields when server_root_path is defined", async () => { + const uiConfig = { + server_root_path: "/api/v1", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([ + { url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }, + { url: "/public/providers/fields", data: [] }, + ]); + + // First call getUiConfig to set up proxyBaseUrl + await Networking.getUiConfig(); + + // Then call the public endpoint + await Networking.getProviderCreateMetadata(); + + expect(mockFetch).toHaveBeenCalledTimes(2); + const publicEndpointCall = mockFetch.mock.calls.find((call) => + (call[0] as string).includes("/public/providers/fields"), + ); + expect(publicEndpointCall).toBeDefined(); + const calledUrl = publicEndpointCall![0] as string; + expect(calledUrl).toBe("https://example.com/api/v1/public/providers/fields"); + }); + + it("should use proxyBaseURL and server_root_path for /public/model_hub/info when server_root_path is defined", async () => { + const uiConfig = { + server_root_path: "/api/v1", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([ + { url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }, + { url: "/public/model_hub/info", data: {} }, + ]); + + await Networking.getUiConfig(); + await Networking.getPublicModelHubInfo(); + + expect(mockFetch).toHaveBeenCalledTimes(2); + const publicEndpointCall = mockFetch.mock.calls.find((call) => + (call[0] as string).includes("/public/model_hub/info"), + ); + expect(publicEndpointCall).toBeDefined(); + const calledUrl = publicEndpointCall![0] as string; + expect(calledUrl).toBe("https://example.com/api/v1/public/model_hub/info"); + }); + + it("should use proxyBaseURL and server_root_path for /public/model_hub when server_root_path is defined", async () => { + const uiConfig = { + server_root_path: "/api/v1", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([ + { url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }, + { url: "/public/model_hub", data: [] }, + ]); + + await Networking.getUiConfig(); + await Networking.modelHubPublicModelsCall(); + + expect(mockFetch).toHaveBeenCalledTimes(2); + const publicEndpointCall = mockFetch.mock.calls.find( + (call) => (call[0] as string).includes("/public/model_hub") && !(call[0] as string).includes("/info"), + ); + expect(publicEndpointCall).toBeDefined(); + const calledUrl = publicEndpointCall![0] as string; + expect(calledUrl).toBe("https://example.com/api/v1/public/model_hub"); + }); + + it("should use proxyBaseURL and server_root_path for /public/agent_hub when server_root_path is defined", async () => { + const uiConfig = { + server_root_path: "/api/v1", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([ + { url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }, + { url: "/public/agent_hub", data: [] }, + ]); + + await Networking.getUiConfig(); + await Networking.agentHubPublicModelsCall(); + + expect(mockFetch).toHaveBeenCalledTimes(2); + const publicEndpointCall = mockFetch.mock.calls.find((call) => (call[0] as string).includes("/public/agent_hub")); + expect(publicEndpointCall).toBeDefined(); + const calledUrl = publicEndpointCall![0] as string; + expect(calledUrl).toBe("https://example.com/api/v1/public/agent_hub"); + }); + + it("should use proxyBaseURL and server_root_path for /public/mcp_hub when server_root_path is defined", async () => { + const uiConfig = { + server_root_path: "/api/v1", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([ + { url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }, + { url: "/public/mcp_hub", data: [] }, + ]); + + await Networking.getUiConfig(); + await Networking.mcpHubPublicServersCall(); + + expect(mockFetch).toHaveBeenCalledTimes(2); + const publicEndpointCall = mockFetch.mock.calls.find((call) => (call[0] as string).includes("/public/mcp_hub")); + expect(publicEndpointCall).toBeDefined(); + const calledUrl = publicEndpointCall![0] as string; + expect(calledUrl).toBe("https://example.com/api/v1/public/mcp_hub"); + }); + + it("should not include server_root_path when it is root path", async () => { + const uiConfig = { + server_root_path: "/", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([ + { url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }, + { url: "/public/providers/fields", data: [] }, + ]); + + await Networking.getUiConfig(); + await Networking.getProviderCreateMetadata(); + + expect(mockFetch).toHaveBeenCalledTimes(2); + const publicEndpointCall = mockFetch.mock.calls.find((call) => + (call[0] as string).includes("/public/providers/fields"), + ); + expect(publicEndpointCall).toBeDefined(); + const calledUrl = publicEndpointCall![0] as string; + expect(calledUrl).toBe("https://example.com/public/providers/fields"); + }); + + it("should return UI config from getUiConfig", async () => { + const uiConfig = { + server_root_path: "/api/v1", + proxy_base_url: "https://example.com", + }; + + const mockFetch = setupMockFetch([{ url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }]); + + const result = await Networking.getUiConfig(); + + expect(mockFetch).toHaveBeenCalledOnce(); + expect(result).toEqual(uiConfig); + const configCall = mockFetch.mock.calls.find((call) => + (call[0] as string).includes("/litellm/.well-known/litellm-ui-config"), + ); + expect(configCall).toBeDefined(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 065d665e134..990312ee198 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -124,6 +124,7 @@ export interface PromptSpec { prompt_info: PromptInfo; created_at?: string; updated_at?: string; + version?: number; // Explicit version number for version history } export interface PromptTemplateBase { @@ -217,7 +218,7 @@ const handleError = async (errorData: string | any) => { if (currentTime - lastErrorTime > 60000) { // 60000 milliseconds = 60 seconds // Convert errorData to string if it isn't already - const errorString = typeof errorData === 'string' ? errorData : JSON.stringify(errorData); + const errorString = typeof errorData === "string" ? errorData : JSON.stringify(errorData); if (errorString.includes("Authentication Error - Expired Key")) { NotificationsManager.info("UI Session Expired. Logging out."); lastErrorTime = currentTime; @@ -238,7 +239,7 @@ export const getProviderCreateMetadata = async (): Promise * Fetch provider credential field metadata from the proxy's public endpoint. * This is used by the UI to dynamically render provider-specific credential fields. */ - const url = defaultProxyBaseUrl ? `${defaultProxyBaseUrl}/public/providers/fields` : `/public/providers/fields`; + const url = proxyBaseUrl ? `${proxyBaseUrl}/public/providers/fields` : `/public/providers/fields`; const response = await fetch(url, { method: "GET", }); @@ -295,7 +296,7 @@ export const getUiConfig = async () => { }; export const getPublicModelHubInfo = async () => { - const url = defaultProxyBaseUrl ? `${defaultProxyBaseUrl}/public/model_hub/info` : `/public/model_hub/info`; + const url = proxyBaseUrl ? `${proxyBaseUrl}/public/model_hub/info` : `/public/model_hub/info`; const response = await fetch(url); const jsonData: PublicModelHubInfo = await response.json(); return jsonData; @@ -5239,6 +5240,35 @@ export const getPromptInfo = async (accessToken: string, promptId: string): Prom } }; +export const getPromptVersions = async (accessToken: string, promptId: string): Promise => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/${promptId}/versions` : `/prompts/${promptId}/versions`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + // Don't throw global error for 404 (no versions found) as we might want to handle it gracefully + if (response.status !== 404) { + handleError(errorMessage); + } + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to get prompt versions:", error); + throw error; + } +}; + export const createPromptCall = async (accessToken: string, promptData: any) => { try { const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts` : `/prompts`; @@ -6720,7 +6750,6 @@ export const getGuardrailProviderSpecificParams = async (accessToken: string) => } }; - export const getAgentsList = async (accessToken: string) => { try { const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/agents` : `/v1/agents`; @@ -6838,7 +6867,6 @@ export const patchAgentCall = async ( } }; - export const updateGuardrailCall = async ( accessToken: string, guardrailId: string, diff --git a/ui/litellm-dashboard/src/components/prompts.tsx b/ui/litellm-dashboard/src/components/prompts.tsx index a8cdb923f32..24df7e78ab0 100644 --- a/ui/litellm-dashboard/src/components/prompts.tsx +++ b/ui/litellm-dashboard/src/components/prompts.tsx @@ -21,6 +21,7 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { const [selectedPromptId, setSelectedPromptId] = useState(null); const [isAddModalVisible, setIsAddModalVisible] = useState(false); const [showEditorView, setShowEditorView] = useState(false); + const [editPromptData, setEditPromptData] = useState(null); const [isDeleting, setIsDeleting] = useState(false); const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string } | null>(null); @@ -55,6 +56,12 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { if (selectedPromptId) { setSelectedPromptId(null); } + setEditPromptData(null); + setShowEditorView(true); + }; + + const handleEditPrompt = (promptData: any) => { + setEditPromptData(promptData); setShowEditorView(true); }; @@ -71,10 +78,14 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { const handleCloseEditor = () => { setShowEditorView(false); + setEditPromptData(null); }; const handleSuccess = () => { fetchPrompts(); + setShowEditorView(false); + setEditPromptData(null); + setSelectedPromptId(null); }; const handleDeleteClick = (promptId: string, promptName: string) => { @@ -109,6 +120,7 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { onClose={handleCloseEditor} onSuccess={handleSuccess} accessToken={accessToken} + initialPromptData={editPromptData} /> ) : selectedPromptId ? ( = ({ accessToken, userRole }) => { accessToken={accessToken} isAdmin={isAdmin} onDelete={fetchPrompts} + onEdit={handleEditPrompt} /> ) : ( <> diff --git a/ui/litellm-dashboard/src/components/prompts/prompt_editor_view/ConversationPanel.tsx b/ui/litellm-dashboard/src/components/prompts/prompt_editor_view/ConversationPanel.tsx deleted file mode 100644 index a827349d6e1..00000000000 --- a/ui/litellm-dashboard/src/components/prompts/prompt_editor_view/ConversationPanel.tsx +++ /dev/null @@ -1,21 +0,0 @@ -import React from "react"; -import { MessageSquareIcon } from "lucide-react"; - -const ConversationPanel: React.FC = () => { - return ( -
-
-
-
- -
-

Your conversation will appear here

-

Save the prompt to test it

-
-
-
- ); -}; - -export default ConversationPanel; - diff --git a/ui/litellm-dashboard/src/components/prompts/prompt_editor_view/PromptCodeSnippets.tsx b/ui/litellm-dashboard/src/components/prompts/prompt_editor_view/PromptCodeSnippets.tsx new file mode 100644 index 00000000000..ed7caf60de8 --- /dev/null +++ b/ui/litellm-dashboard/src/components/prompts/prompt_editor_view/PromptCodeSnippets.tsx @@ -0,0 +1,284 @@ +import React, { useState } from "react"; +import { Modal, Select, Button as AntdButton, Tabs } from "antd"; +import { CodeOutlined } from "@ant-design/icons"; +import { Button as TremorButton, Text } from "@tremor/react"; +import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; +import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; +import NotificationsManager from "../../molecules/notifications_manager"; + +interface PromptCodeSnippetsProps { + promptId: string; + model: string; + promptVariables?: Record; + accessToken: string | null; + version?: string; + proxySettings?: { + PROXY_BASE_URL?: string; + LITELLM_UI_API_DOC_BASE_URL?: string | null; + }; +} + +const PromptCodeSnippets: React.FC = ({ + promptId, + model, + promptVariables = {}, + accessToken, + version = "1", + proxySettings, +}) => { + const [isModalVisible, setIsModalVisible] = useState(false); + const [selectedLanguage, setSelectedLanguage] = useState<"curl" | "python" | "javascript">("curl"); + const [selectedTab, setSelectedTab] = useState("basic"); + const [generatedCode, setGeneratedCode] = useState(""); + + const showModal = () => { + setIsModalVisible(true); + }; + + const handleCancel = () => { + setIsModalVisible(false); + }; + + // Determine base URL with priority: LITELLM_UI_API_DOC_BASE_URL > PROXY_BASE_URL > window.location.origin + let apiBase = window.location.origin; + const customDocBaseUrl = proxySettings?.LITELLM_UI_API_DOC_BASE_URL; + if (customDocBaseUrl && customDocBaseUrl.trim()) { + apiBase = customDocBaseUrl; + } else if (proxySettings?.PROXY_BASE_URL) { + apiBase = proxySettings.PROXY_BASE_URL; + } + + const effectiveApiKey = accessToken || "sk-1234"; + + // Generate code based on selected language and tab + const generateCode = () => { + const hasVariables = Object.keys(promptVariables).length > 0; + + if (selectedLanguage === "curl") { + if (selectedTab === "basic") { + return `curl -X POST '${apiBase}/chat/completions' \\ + -H 'Content-Type: application/json' \\ + -H 'Authorization: Bearer ${effectiveApiKey}' \\ + -d '{ + "model": "${model}", + "prompt_id": "${promptId}"${hasVariables ? `, + "prompt_variables": ${JSON.stringify(promptVariables, null, 6).replace(/\n/g, '\n ')}` : ''} + }' | jq`; + } else if (selectedTab === "messages") { + return `curl -X POST '${apiBase}/chat/completions' \\ + -H 'Content-Type: application/json' \\ + -H 'Authorization: Bearer ${effectiveApiKey}' \\ + -d '{ + "model": "${model}", + "prompt_id": "${promptId}"${hasVariables ? `, + "prompt_variables": ${JSON.stringify(promptVariables, null, 6).replace(/\n/g, '\n ')}` : ''}, + "messages": [ + { + "role": "user", + "content": "hi" + } + ] + }' | jq`; + } else { + return `curl -X POST '${apiBase}/chat/completions' \\ + -H 'Content-Type: application/json' \\ + -H 'Authorization: Bearer ${effectiveApiKey}' \\ + -d '{ + "model": "${model}", + "prompt_id": "${promptId}", + "prompt_version": ${version}, + "messages": [ + { + "role": "user", + "content": "Who are u" + } + ] + }' | jq`; + } + } else if (selectedLanguage === "python") { + const importCode = `import openai + +client = openai.OpenAI( + api_key="${effectiveApiKey}", + base_url="${apiBase}" +) +`; + if (selectedTab === "basic") { + return `${importCode} +response = client.chat.completions.create( + model="${model}", + extra_body={ + "prompt_id": "${promptId}"${hasVariables ? `, + "prompt_variables": ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, '\n ')}` : ''} + } +) + +print(response)`; + } else if (selectedTab === "messages") { + return `${importCode} +response = client.chat.completions.create( + model="${model}", + messages=[ + {"role": "user", "content": "hi"} + ], + extra_body={ + "prompt_id": "${promptId}"${hasVariables ? `, + "prompt_variables": ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, '\n ')}` : ''} + } +) + +print(response)`; + } else { + return `${importCode} +response = client.chat.completions.create( + model="${model}", + messages=[ + {"role": "user", "content": "Who are u"} + ], + extra_body={ + "prompt_id": "${promptId}", + "prompt_version": ${version} + } +) + +print(response)`; + } + } else { + // JavaScript/Node.js + const importCode = `import OpenAI from 'openai'; + +const client = new OpenAI({ + apiKey: "${effectiveApiKey}", + baseURL: "${apiBase}" +}); +`; + if (selectedTab === "basic") { + return `${importCode} +async function main() { + const response = await client.chat.completions.create({ + model: "${model}", + ${hasVariables ? `prompt_id: "${promptId}", + prompt_variables: ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, '\n ')}` : `prompt_id: "${promptId}"`} + }); + + console.log(response); +} + +main();`; + } else if (selectedTab === "messages") { + return `${importCode} +async function main() { + const response = await client.chat.completions.create({ + model: "${model}", + messages: [ + { role: "user", content: "hi" } + ], + ${hasVariables ? `prompt_id: "${promptId}", + prompt_variables: ${JSON.stringify(promptVariables, null, 8).replace(/\n/g, '\n ')}` : `prompt_id: "${promptId}"`} + }); + + console.log(response); +} + +main();`; + } else { + return `${importCode} +async function main() { + const response = await client.chat.completions.create({ + model: "${model}", + messages: [ + { role: "user", content: "Who are u" } + ], + prompt_id: "${promptId}", + prompt_version: ${version} + }); + + console.log(response); +} + +main();`; + } + } + }; + + // Update generated code when language, tab or props change + React.useEffect(() => { + if (isModalVisible) { + setGeneratedCode(generateCode()); + } + }, [isModalVisible, selectedLanguage, selectedTab, promptId, model, promptVariables]); + + return ( + <> + + Get Code + + + +
+
+ Language +