Merge remote-tracking branch 'origin' into litellm_ui_callback_fix

This commit is contained in:
yuneng-jiang 2025-11-21 16:27:49 -08:00
commit 5dad3c9708
82 changed files with 5501 additions and 550 deletions

View file

@ -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 && \

View file

@ -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/<servers or access groups>`. This allows you to:
LiteLLM Proxy supports URL-based namespacing for MCP servers using the format `/<servers or access groups>/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
```
<your-litellm-proxy-base-url>/mcp/<server_alias_or_access_group>
<your-litellm-proxy-base-url>/<server_alias_or_access_group>/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": "<your-litellm-proxy-base-url>/mcp/github",
"server_url": "<your-litellm-proxy-base-url>/github_mcp/mcp",
"require_approval": "never",
"headers": {
"x-litellm-api-key": "Bearer YOUR_LITELLM_API_KEY"
@ -718,7 +718,7 @@ curl --location '<your-litellm-proxy-base-url>/v1/responses' \
{
"type": "mcp",
"server_label": "litellm",
"server_url": "<your-litellm-proxy-base-url>/mcp/dev_group",
"server_url": "<your-litellm-proxy-base-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": "<your-litellm-proxy-base-url>/mcp/github,zapier",
"url": "<your-litellm-proxy-base-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/<servers>` |
| **Endpoint** | Standard `litellm_proxy` endpoint | Custom `/mcp/<servers>` endpoint |
| **Method** | Uses `x-mcp-servers` header | Uses URL path `/<servers>/mcp` |
| **Endpoint** | Standard `litellm_proxy` endpoint | Custom `/<servers>/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 |

View file

@ -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) |
<br />
<br />
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
```
<Tabs>
<TabItem value="openai-sdk" label="OpenAI SDK">
```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="")
```
</TabItem>
<TabItem value="litellm-sdk" label="LiteLLM SDK">
```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="")
```
</TabItem>
<TabItem value="curl" label="cURL">
```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
}'
```
</TabItem>
</Tabs>
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/).

View file

@ -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__<thought-signature>`) and extracting it back out before sending the outbound request to Gemini.
## JSON Mode
<Tabs>

View file

@ -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

View file

@ -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
<Tabs>
<TabItem value="block" label="Block Request">
**Block requset**
**Block request (`on_disallowed_action: block`)**
```bash
# Test
@ -96,7 +133,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \
</TabItem>
<TabItem value="rewrite" label="Rewrite Request">
**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
{

View file

@ -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.
:::

View file

@ -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",
],
},
{

View file

@ -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

View file

@ -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

View file

@ -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`

View file

@ -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:

View file

@ -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(

View file

@ -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,

View file

@ -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_<uuid>__thought__<base64_signature>
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":

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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",

View file

@ -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
)

View file

@ -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,

View file

@ -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"]
failure_callback: ["langfuse"]

View file

@ -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(

View file

@ -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,
)

View file

@ -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

View file

@ -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,

View file

@ -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 <your_api_key>"
```
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 <your_api_key>" \\
-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"],

View file

@ -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:

View file

@ -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,
)

View file

@ -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

View file

@ -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
)

View file

@ -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()

View file

@ -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)

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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",

View file

@ -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",

View file

@ -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

View file

@ -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.
]

View file

@ -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}")

View file

@ -26,7 +26,7 @@ test("admin login test", async ({ page }) => {
await loginButton.click();
const tabs = [
"Virtual Keys",
"Test Key",
"Playground",
"Models",
"Usage",
"Teams",

View file

@ -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

View file

@ -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))

View file

@ -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

View file

@ -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

View file

@ -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",
]
],
)

View file

@ -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?"

View file

@ -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

View file

@ -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(

View file

@ -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"

View file

@ -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

View file

@ -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__])

View file

@ -406,7 +406,7 @@ const MCPConnect: React.FC<MCPConnectProps> = ({ 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"]

View file

@ -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();
});
});

View file

@ -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<ProviderCreateInfo[]>
* 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<ListPromptsResponse> => {
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,

View file

@ -21,6 +21,7 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
const [selectedPromptId, setSelectedPromptId] = useState<string | null>(null);
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
const [showEditorView, setShowEditorView] = useState(false);
const [editPromptData, setEditPromptData] = useState<any>(null);
const [isDeleting, setIsDeleting] = useState(false);
const [promptToDelete, setPromptToDelete] = useState<{ id: string; name: string } | null>(null);
@ -55,6 +56,12 @@ const PromptsPanel: React.FC<PromptsProps> = ({ 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<PromptsProps> = ({ 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<PromptsProps> = ({ accessToken, userRole }) => {
onClose={handleCloseEditor}
onSuccess={handleSuccess}
accessToken={accessToken}
initialPromptData={editPromptData}
/>
) : selectedPromptId ? (
<PromptInfoView
@ -117,6 +129,7 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
accessToken={accessToken}
isAdmin={isAdmin}
onDelete={fetchPrompts}
onEdit={handleEditPrompt}
/>
) : (
<>

View file

@ -1,21 +0,0 @@
import React from "react";
import { MessageSquareIcon } from "lucide-react";
const ConversationPanel: React.FC = () => {
return (
<div className="flex-1 bg-white flex flex-col">
<div className="flex-1 flex items-center justify-center text-gray-400">
<div className="text-center">
<div className="w-12 h-12 mx-auto mb-3 bg-gray-100 rounded-full flex items-center justify-center">
<MessageSquareIcon size={24} className="text-gray-400" />
</div>
<p className="text-sm">Your conversation will appear here</p>
<p className="text-xs text-gray-500 mt-2">Save the prompt to test it</p>
</div>
</div>
</div>
);
};
export default ConversationPanel;

View file

@ -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<string, string>;
accessToken: string | null;
version?: string;
proxySettings?: {
PROXY_BASE_URL?: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
};
}
const PromptCodeSnippets: React.FC<PromptCodeSnippetsProps> = ({
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 (
<>
<TremorButton
variant="secondary"
icon={CodeOutlined}
onClick={showModal}
>
Get Code
</TremorButton>
<Modal
title="Generated Code"
open={isModalVisible}
onCancel={handleCancel}
footer={null}
width={800}
>
<div className="flex justify-between items-center mb-4">
<div>
<Text className="font-medium block mb-1 text-gray-700">Language</Text>
<Select
value={selectedLanguage}
onChange={(value) => setSelectedLanguage(value as "curl" | "python" | "javascript")}
style={{ width: 180 }}
options={[
{ value: "curl", label: "cURL" },
{ value: "python", label: "Python (OpenAI SDK)" },
{ value: "javascript", label: "JavaScript (OpenAI SDK)" },
]}
/>
</div>
<AntdButton
onClick={() => {
navigator.clipboard.writeText(generatedCode);
NotificationsManager.success("Copied to clipboard!");
}}
>
Copy to Clipboard
</AntdButton>
</div>
<Tabs
activeKey={selectedTab}
onChange={setSelectedTab}
items={[
{ label: "Basic", key: "basic" },
{ label: "With Messages", key: "messages" },
{ label: "With Version", key: "version" },
]}
/>
<SyntaxHighlighter
language={selectedLanguage === "curl" ? "bash" : selectedLanguage === "python" ? "python" : "javascript"}
style={coy as any}
wrapLines={true}
wrapLongLines={true}
className="rounded-md mt-0"
customStyle={{
maxHeight: "60vh",
overflowY: "auto",
marginTop: 0,
borderTopLeftRadius: 0,
borderTopRightRadius: 0,
}}
>
{generatedCode}
</SyntaxHighlighter>
</Modal>
</>
);
};
export default PromptCodeSnippets;

View file

@ -1,7 +1,8 @@
import React from "react";
import { Button as TremorButton } from "@tremor/react";
import { Input } from "antd";
import { ArrowLeftIcon, SaveIcon } from "lucide-react";
import { ArrowLeftIcon, SaveIcon, ClockIcon } from "lucide-react";
import PromptCodeSnippets from "./PromptCodeSnippets";
interface PromptEditorHeaderProps {
promptName: string;
@ -9,6 +10,16 @@ interface PromptEditorHeaderProps {
onBack: () => void;
onSave: () => void;
isSaving: boolean;
editMode?: boolean;
onShowHistory?: () => void;
version?: string | null;
promptModel?: string;
promptVariables?: Record<string, string>;
accessToken: string | null;
proxySettings?: {
PROXY_BASE_URL?: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
};
}
const PromptEditorHeader: React.FC<PromptEditorHeaderProps> = ({
@ -17,6 +28,13 @@ const PromptEditorHeader: React.FC<PromptEditorHeaderProps> = ({
onBack,
onSave,
isSaving,
editMode = false,
onShowHistory,
version,
promptModel = "gpt-4o",
promptVariables = {},
accessToken,
proxySettings,
}) => {
return (
<div className="bg-white border-b border-gray-200 px-6 py-3 flex items-center justify-between">
@ -30,17 +48,39 @@ const PromptEditorHeader: React.FC<PromptEditorHeaderProps> = ({
className="text-base font-medium border-none shadow-none"
style={{ width: "200px" }}
/>
{version && (
<span className="px-2 py-0.5 text-xs bg-blue-100 text-blue-700 rounded font-medium">
{version}
</span>
)}
<span className="px-2 py-0.5 text-xs bg-gray-100 text-gray-600 rounded">Draft</span>
<span className="text-xs text-gray-400">Unsaved changes</span>
</div>
<div className="flex items-center space-x-2">
<PromptCodeSnippets
promptId={promptName}
model={promptModel}
promptVariables={promptVariables}
accessToken={accessToken}
version={version?.replace('v', '') || "1"}
proxySettings={proxySettings}
/>
{editMode && onShowHistory && (
<TremorButton
icon={ClockIcon}
variant="secondary"
onClick={onShowHistory}
>
History
</TremorButton>
)}
<TremorButton
icon={SaveIcon}
onClick={onSave}
loading={isSaving}
disabled={isSaving}
>
Save
{editMode ? "Update" : "Save"}
</TremorButton>
</div>
</div>

View file

@ -0,0 +1,140 @@
import { Drawer, List, Skeleton, Tag, Typography } from "antd";
import React, { useEffect, useState } from "react";
import { getPromptVersions, PromptSpec } from "../../networking";
const { Text } = Typography;
interface VersionHistorySidePanelProps {
isOpen: boolean;
onClose: () => void;
accessToken: string | null;
promptId: string;
activeVersionId?: string;
onSelectVersion?: (version: PromptSpec) => void;
}
const VersionHistorySidePanel: React.FC<VersionHistorySidePanelProps> = ({
isOpen,
onClose,
accessToken,
promptId,
activeVersionId,
onSelectVersion,
}) => {
const [versions, setVersions] = useState<PromptSpec[]>([]);
const [loading, setLoading] = useState(false);
useEffect(() => {
if (isOpen && accessToken && promptId) {
fetchVersions();
}
}, [isOpen, accessToken, promptId]);
const fetchVersions = async () => {
setLoading(true);
try {
// Strip .v suffix if present to get base ID for querying all versions
const basePromptId = promptId.includes(".v") ? promptId.split(".v")[0] : promptId;
const response = await getPromptVersions(accessToken!, basePromptId);
setVersions(response.prompts);
} catch (error) {
console.error("Error fetching prompt versions:", error);
} finally {
setLoading(false);
}
};
const getVersionNumber = (prompt: PromptSpec) => {
// Use explicit version field if available, otherwise try to extract from litellm_params.prompt_id
if (prompt.version) {
return `v${prompt.version}`;
}
// Fallback: try to extract from litellm_params.prompt_id
const versionedId = (prompt.litellm_params as any)?.prompt_id || prompt.prompt_id;
if (versionedId.includes(".v")) {
return `v${versionedId.split(".v")[1]}`;
}
if (versionedId.includes("_v")) {
return `v${versionedId.split("_v")[1]}`;
}
return "v1";
};
const formatDate = (dateString?: string) => {
if (!dateString) return "-";
return new Date(dateString).toLocaleString();
};
return (
<Drawer
title="Version History"
placement="right"
onClose={onClose}
open={isOpen}
width={400}
mask={false} // Allow interacting with the main editor while drawer is open
maskClosable={false}
>
{loading ? (
<Skeleton active paragraph={{ rows: 4 }} />
) : versions.length === 0 ? (
<div className="text-center py-8 text-gray-500">No version history available.</div>
) : (
<List
dataSource={versions}
renderItem={(item, index) => {
// Use version field for comparison since all items have the same prompt_id
const itemVersionNum = item.version || parseInt(getVersionNumber(item).replace('v', ''));
// Extract version number from activeVersionId (may have .vX suffix)
let activeVersionNum: number | null = null;
if (activeVersionId) {
if (activeVersionId.includes('.v')) {
activeVersionNum = parseInt(activeVersionId.split('.v')[1]);
} else if (activeVersionId.includes('_v')) {
activeVersionNum = parseInt(activeVersionId.split('_v')[1]);
}
}
// Default to latest (first item) if no activeVersionId
const isSelected = activeVersionNum ? itemVersionNum === activeVersionNum : index === 0;
return (
<div
key={`${item.prompt_id}-v${item.version || itemVersionNum}`}
className={`mb-4 p-4 rounded-lg border cursor-pointer transition-all hover:shadow-md ${
isSelected ? "border-blue-500 bg-blue-50" : "border-gray-200 bg-white hover:border-blue-300"
}`}
onClick={() => onSelectVersion?.(item)}
>
<div className="flex justify-between items-start mb-2">
<div className="flex items-center gap-2">
<Tag className="m-0">
{getVersionNumber(item)}
</Tag>
{index === 0 && <Tag color="blue" className="m-0">Latest</Tag>}
</div>
{isSelected && (
<Tag color="green" className="m-0">
Active
</Tag>
)}
</div>
<div className="flex flex-col gap-1">
<Text className="text-sm text-gray-600 font-medium">{formatDate(item.created_at)}</Text>
<Text type="secondary" className="text-xs">
{item.prompt_info?.prompt_type === "db" ? "Saved to Database" : "Config Prompt"}
</Text>
</div>
</div>
);
}}
/>
)}
</Drawer>
);
};
export default VersionHistorySidePanel;

View file

@ -0,0 +1,22 @@
import React from "react";
import { RobotOutlined } from "@ant-design/icons";
interface EmptyStateProps {
hasVariables: boolean;
}
const EmptyState: React.FC<EmptyStateProps> = ({ hasVariables }) => {
return (
<div className="h-full flex flex-col items-center justify-center text-gray-400">
<RobotOutlined style={{ fontSize: "48px", marginBottom: "16px" }} />
<span className="text-base">
{hasVariables
? "Fill in the variables above, then type a message to start testing"
: "Type a message below to start testing your prompt"}
</span>
</div>
);
};
export default EmptyState;

View file

@ -0,0 +1,115 @@
import React from "react";
import { RobotOutlined, UserOutlined } from "@ant-design/icons";
import ReactMarkdown from "react-markdown";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { coy } from "react-syntax-highlighter/dist/esm/styles/prism";
import ResponseMetrics from "../../../playground/chat_ui/ResponseMetrics";
import { Message } from "./types";
interface MessageBubbleProps {
message: Message;
}
const MessageBubble: React.FC<MessageBubbleProps> = ({ message }) => {
return (
<div className={`mb-4 flex ${message.role === "user" ? "justify-end" : "justify-start"}`}>
<div
className="max-w-[85%] rounded-lg shadow-sm p-3.5 px-4"
style={{
backgroundColor: message.role === "user" ? "#f0f8ff" : "#ffffff",
border: message.role === "user" ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
}}
>
<div className="flex items-center gap-2 mb-1.5">
<div
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
style={{
backgroundColor: message.role === "user" ? "#e6f0fa" : "#f5f5f5",
}}
>
{message.role === "user" ? (
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
) : (
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
)}
</div>
<strong className="text-sm capitalize">{message.role}</strong>
{message.role === "assistant" && message.model && (
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
{message.model}
</span>
)}
</div>
<div
className="whitespace-pre-wrap break-words max-w-full message-content"
style={{
wordWrap: "break-word",
overflowWrap: "break-word",
wordBreak: "break-word",
hyphens: "auto",
}}
>
{message.role === "assistant" ? (
<ReactMarkdown
components={{
code({
node,
inline,
className,
children,
...props
}: React.ComponentPropsWithoutRef<"code"> & {
inline?: boolean;
node?: any;
}) {
const match = /language-(\w+)/.exec(className || "");
return !inline && match ? (
<SyntaxHighlighter
style={coy as any}
language={match[1]}
PreTag="div"
className="rounded-md my-2"
wrapLines={true}
wrapLongLines={true}
{...props}
>
{String(children).replace(/\n$/, "")}
</SyntaxHighlighter>
) : (
<code
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
style={{ wordBreak: "break-word" }}
{...props}
>
{children}
</code>
);
},
pre: ({ node, ...props }) => (
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
),
}}
>
{message.content}
</ReactMarkdown>
) : (
<div className="whitespace-pre-wrap">{message.content}</div>
)}
{message.role === "assistant" &&
(message.timeToFirstToken || message.totalLatency || message.usage) && (
<ResponseMetrics
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
usage={message.usage}
/>
)}
</div>
</div>
</div>
);
};
export default MessageBubble;

View file

@ -0,0 +1,71 @@
import React from "react";
import { ArrowUpOutlined } from "@ant-design/icons";
import { Button as TremorButton } from "@tremor/react";
import { Input } from "antd";
const { TextArea } = Input;
interface MessageInputProps {
inputMessage: string;
isLoading: boolean;
isDisabled: boolean;
onInputChange: (value: string) => void;
onSend: () => void;
onKeyDown: (event: React.KeyboardEvent<HTMLTextAreaElement>) => void;
onCancel: () => void;
}
const MessageInput: React.FC<MessageInputProps> = ({
inputMessage,
isLoading,
isDisabled,
onInputChange,
onSend,
onKeyDown,
onCancel,
}) => {
return (
<div className="flex items-center gap-2">
<div className="flex items-center flex-1 bg-white border border-gray-300 rounded-xl px-3 py-1 min-h-[44px]">
<TextArea
value={inputMessage}
onChange={(e) => onInputChange(e.target.value)}
onKeyDown={onKeyDown}
placeholder="Type your message... (Shift+Enter for new line)"
disabled={isLoading}
className="flex-1"
autoSize={{ minRows: 1, maxRows: 4 }}
style={{
resize: "none",
border: "none",
boxShadow: "none",
background: "transparent",
padding: "4px 0",
fontSize: "14px",
lineHeight: "20px",
}}
/>
<TremorButton
onClick={onSend}
disabled={isDisabled}
className="flex-shrink-0 ml-2 !w-8 !h-8 !min-w-8 !p-0 !rounded-full !bg-blue-600 hover:!bg-blue-700 disabled:!bg-gray-300 !border-none !text-white disabled:!text-gray-500 !flex !items-center !justify-center"
>
<ArrowUpOutlined style={{ fontSize: "14px" }} />
</TremorButton>
</div>
{isLoading && (
<TremorButton
onClick={onCancel}
className="bg-red-50 hover:bg-red-100 text-red-600 border-red-200"
>
Cancel
</TremorButton>
)}
</div>
);
};
export default MessageInput;

View file

@ -0,0 +1,42 @@
import React from "react";
import { LoadingOutlined } from "@ant-design/icons";
import { Spin } from "antd";
import EmptyState from "./EmptyState";
import MessageBubble from "./MessageBubble";
import { Message } from "./types";
interface MessageListProps {
messages: Message[];
isLoading: boolean;
hasVariables: boolean;
messagesEndRef: React.RefObject<HTMLDivElement>;
}
const MessageList: React.FC<MessageListProps> = ({
messages,
isLoading,
hasVariables,
messagesEndRef,
}) => {
const antIcon = <LoadingOutlined style={{ fontSize: 24 }} spin />;
return (
<div className="flex-1 overflow-y-auto p-4 pb-0">
{messages.length === 0 && <EmptyState hasVariables={hasVariables} />}
{messages.map((message, index) => (
<MessageBubble key={index} message={message} />
))}
{isLoading && (
<div className="flex justify-center items-center my-4">
<Spin indicator={antIcon} />
</div>
)}
<div ref={messagesEndRef} style={{ height: "1px" }} />
</div>
);
};
export default MessageList;

View file

@ -0,0 +1,44 @@
import React from "react";
import { Input } from "antd";
interface VariableInputProps {
extractedVariables: string[];
variables: Record<string, string>;
onVariableChange: (varName: string, value: string) => void;
}
const VariableInput: React.FC<VariableInputProps> = ({
extractedVariables,
variables,
onVariableChange,
}) => {
if (extractedVariables.length === 0) {
return null;
}
return (
<div className="p-4 border-b border-gray-200 bg-blue-50">
<h3 className="text-sm font-semibold text-gray-700 mb-3">
Fill in template variables to start testing
</h3>
<div className="space-y-2">
{extractedVariables.map((varName) => (
<div key={varName}>
<label className="block text-xs text-gray-600 mb-1 font-medium">
{"{{"}{varName}{"}}"}
</label>
<Input
value={variables[varName] || ""}
onChange={(e) => onVariableChange(varName, e.target.value)}
placeholder={`Enter value for ${varName}`}
size="small"
/>
</div>
))}
</div>
</div>
);
};
export default VariableInput;

View file

@ -0,0 +1,38 @@
import React from "react";
interface VariableWarningProps {
extractedVariables: string[];
variables: Record<string, string>;
}
const VariableWarning: React.FC<VariableWarningProps> = ({
extractedVariables,
variables,
}) => {
const missingVariables = extractedVariables.filter(
(varName) => !variables[varName] || variables[varName].trim() === ""
);
if (missingVariables.length === 0) {
return null;
}
return (
<div className="mb-3 p-3 bg-yellow-50 border border-yellow-200 rounded-lg">
<div className="flex items-start gap-2">
<span className="text-yellow-600 text-sm">⚠️</span>
<div className="flex-1">
<p className="text-sm text-yellow-800 font-medium mb-1">
Please fill in all template variables above
</p>
<p className="text-xs text-yellow-700">
Missing: {missingVariables.map((varName) => `{{${varName}}}`).join(", ")}
</p>
</div>
</div>
</div>
);
};
export default VariableWarning;

View file

@ -0,0 +1,78 @@
import React from "react";
import { ClearOutlined } from "@ant-design/icons";
import { Button as TremorButton } from "@tremor/react";
import { ConversationPanelProps } from "./types";
import { useConversation } from "./useConversation";
import VariableInput from "./VariableInput";
import MessageList from "./MessageList";
import VariableWarning from "./VariableWarning";
import MessageInput from "./MessageInput";
const ConversationPanel: React.FC<ConversationPanelProps> = ({ prompt, accessToken }) => {
const {
isLoading,
messages,
inputMessage,
variables,
variablesFilled,
extractedVariables,
allVariablesFilled,
messagesEndRef,
setInputMessage,
handleSendMessage,
handleCancelRequest,
handleClearConversation,
handleKeyDown,
handleVariableChange,
} = useConversation(prompt, accessToken);
return (
<div className="flex flex-col h-full bg-white">
{!variablesFilled && (
<VariableInput
extractedVariables={extractedVariables}
variables={variables}
onVariableChange={handleVariableChange}
/>
)}
{messages.length > 0 && (
<div className="p-3 border-b border-gray-200 bg-white flex justify-end">
<TremorButton
onClick={handleClearConversation}
className="bg-gray-100 hover:bg-gray-200 text-gray-700 border-gray-300"
icon={ClearOutlined}
>
Clear Chat
</TremorButton>
</div>
)}
<MessageList
messages={messages}
isLoading={isLoading}
hasVariables={extractedVariables.length > 0}
messagesEndRef={messagesEndRef}
/>
<div className="p-4 border-t border-gray-200 bg-white">
<VariableWarning extractedVariables={extractedVariables} variables={variables} />
<MessageInput
inputMessage={inputMessage}
isLoading={isLoading}
isDisabled={
isLoading || !inputMessage.trim() || (extractedVariables.length > 0 && !allVariablesFilled)
}
onInputChange={setInputMessage}
onSend={handleSendMessage}
onKeyDown={handleKeyDown}
onCancel={handleCancelRequest}
/>
</div>
</div>
);
};
export default ConversationPanel;

View file

@ -0,0 +1,16 @@
import { TokenUsage } from "../../../playground/chat_ui/ResponseMetrics";
export interface Message {
role: string;
content: string;
model?: string;
timeToFirstToken?: number;
totalLatency?: number;
usage?: TokenUsage;
}
export interface ConversationPanelProps {
prompt: any;
accessToken: string | null;
}

View file

@ -0,0 +1,242 @@
import { useState, useRef, useEffect } from "react";
import NotificationsManager from "../../../molecules/notifications_manager";
import { TokenUsage } from "../../../playground/chat_ui/ResponseMetrics";
import { Message } from "./types";
import { convertToDotPrompt, extractVariables } from "../utils";
import { getProxyBaseUrl } from "../../../networking";
export const useConversation = (prompt: any, accessToken: string | null) => {
const [isLoading, setIsLoading] = useState(false);
const [messages, setMessages] = useState<Message[]>([]);
const [inputMessage, setInputMessage] = useState("");
const [variables, setVariables] = useState<Record<string, string>>({});
const [variablesFilled, setVariablesFilled] = useState(false);
const [abortController, setAbortController] = useState<AbortController | null>(null);
const messagesEndRef = useRef<HTMLDivElement>(null);
const extractedVariables = extractVariables(prompt);
const allVariablesFilled = extractedVariables.every(
(varName) => variables[varName] && variables[varName].trim() !== "",
);
const scrollToBottom = () => {
if (messagesEndRef.current) {
setTimeout(() => {
messagesEndRef.current?.scrollIntoView({
behavior: "smooth",
block: "end",
});
}, 100);
}
};
useEffect(() => {
scrollToBottom();
}, [messages]);
const handleSendMessage = async () => {
if (!accessToken) {
NotificationsManager.fromBackend("Access token is required");
return;
}
if (extractedVariables.length > 0 && !allVariablesFilled) {
NotificationsManager.fromBackend("Please fill in all template variables");
return;
}
if (!inputMessage.trim()) {
return;
}
if (!variablesFilled && extractedVariables.length > 0) {
setVariablesFilled(true);
}
const userMessage: Message = { role: "user", content: inputMessage };
setMessages((prev) => [...prev, userMessage]);
setInputMessage("");
const controller = new AbortController();
setAbortController(controller);
setIsLoading(true);
const startTime = Date.now();
let timeToFirstToken: number | undefined;
try {
const dotpromptContent = convertToDotPrompt(prompt);
const proxyBaseUrl = getProxyBaseUrl();
const requestBody: any = {
dotprompt_content: dotpromptContent,
};
if (messages.length === 0) {
requestBody.prompt_variables = variables;
} else {
requestBody.conversation_history = [
...messages.map((msg) => ({
role: msg.role,
content: msg.content,
})),
{
role: "user",
content: inputMessage,
},
];
}
const response = await fetch(`${proxyBaseUrl}/prompts/test`, {
method: "POST",
headers: {
Authorization: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(requestBody),
signal: controller.signal,
});
if (!response.ok) {
const errorText = await response.text();
throw new Error(`HTTP error! status: ${response.status}, ${errorText}`);
}
if (!response.body) {
throw new Error("No response body");
}
const reader = response.body.getReader();
const decoder = new TextDecoder();
let assistantMessage = "";
let model: string | undefined;
let usage: TokenUsage | undefined;
setMessages((prev) => [...prev, { role: "assistant", content: "" }]);
// eslint-disable-next-line no-constant-condition
while (true) {
const { done, value } = await reader.read();
if (done) break;
const chunk = decoder.decode(value);
const lines = chunk.split("\n");
for (const line of lines) {
if (line.startsWith("data: ")) {
const data = line.slice(6);
if (data === "[DONE]") {
continue;
}
try {
const parsed = JSON.parse(data);
if (!model && parsed.model) {
model = parsed.model;
}
if (parsed.usage) {
usage = parsed.usage;
}
const content = parsed.choices?.[0]?.delta?.content;
if (content) {
if (!timeToFirstToken) {
timeToFirstToken = Date.now() - startTime;
}
assistantMessage += content;
setMessages((prev) => {
const newMessages = [...prev];
newMessages[newMessages.length - 1] = {
role: "assistant",
content: assistantMessage,
model,
timeToFirstToken,
};
return newMessages;
});
}
} catch (e) {
console.error("Error parsing chunk:", e);
}
}
}
}
const totalLatency = Date.now() - startTime;
setMessages((prev) => {
const newMessages = [...prev];
newMessages[newMessages.length - 1] = {
...newMessages[newMessages.length - 1],
totalLatency,
usage,
};
return newMessages;
});
} catch (error: any) {
if (error.name === "AbortError") {
console.log("Request was cancelled");
} else {
console.error("Error testing prompt:", error);
setMessages((prev) => {
const lastMsg = prev[prev.length - 1];
if (lastMsg && lastMsg.role === "assistant" && lastMsg.content === "") {
return [...prev.slice(0, -1), { role: "assistant", content: `Error: ${error.message}` }];
}
return [...prev, { role: "assistant", content: `Error: ${error.message}` }];
});
}
} finally {
setIsLoading(false);
setAbortController(null);
}
};
const handleCancelRequest = () => {
if (abortController) {
abortController.abort();
setAbortController(null);
setIsLoading(false);
NotificationsManager.info("Request cancelled");
}
};
const handleClearConversation = () => {
setMessages([]);
setVariablesFilled(false);
NotificationsManager.success("Chat history cleared.");
};
const handleKeyDown = (event: React.KeyboardEvent<HTMLTextAreaElement>) => {
if (event.key === "Enter" && !event.shiftKey) {
event.preventDefault();
handleSendMessage();
}
};
const handleVariableChange = (varName: string, value: string) => {
setVariables({ ...variables, [varName]: value });
};
return {
// State
isLoading,
messages,
inputMessage,
variables,
variablesFilled,
extractedVariables,
allVariablesFilled,
messagesEndRef,
// Actions
setInputMessage,
handleSendMessage,
handleCancelRequest,
handleClearConversation,
handleKeyDown,
handleVariableChange,
};
};

View file

@ -1,35 +1,72 @@
import React, { useState } from "react";
import ToolModal from "../tool_modal";
import NotificationsManager from "../../molecules/notifications_manager";
import { createPromptCall } from "../../networking";
import { createPromptCall, updatePromptCall } from "../../networking";
import { PromptType, PromptEditorViewProps, Tool } from "./types";
import { convertToDotPrompt } from "./utils";
import { convertToDotPrompt, parseExistingPrompt } from "./utils";
import PromptEditorHeader from "./PromptEditorHeader";
import ModelConfigCard from "./ModelConfigCard";
import ToolsCard from "./ToolsCard";
import DeveloperMessageCard from "./DeveloperMessageCard";
import PromptMessagesCard from "./PromptMessagesCard";
import ConversationPanel from "./ConversationPanel";
import ConversationPanel from "./conversation_panel";
import PublishModal from "./PublishModal";
import DotpromptViewTab from "./DotpromptViewTab";
import VersionHistorySidePanel from "./VersionHistorySidePanel";
const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess, accessToken }) => {
const [prompt, setPrompt] = useState<PromptType>({
name: "New prompt",
model: "gpt-4o",
config: {
temperature: 1,
max_tokens: 1000,
},
tools: [],
developerMessage: "",
messages: [
{
role: "user",
content: "Enter task specifics. Use {{template_variables}} for dynamic inputs",
const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess, accessToken, initialPromptData }) => {
const getInitialPrompt = (): PromptType => {
if (initialPromptData) {
try {
return parseExistingPrompt(initialPromptData);
} catch (error) {
console.error("Error parsing existing prompt:", error);
NotificationsManager.fromBackend("Failed to parse prompt data");
}
}
return {
name: "New prompt",
model: "gpt-4o",
config: {
temperature: 1,
max_tokens: 1000,
},
],
});
tools: [],
developerMessage: "",
messages: [
{
role: "user",
content: "Enter task specifics. Use {{template_variables}} for dynamic inputs",
},
],
};
};
const [prompt, setPrompt] = useState<PromptType>(getInitialPrompt());
const [editMode, setEditMode] = useState<boolean>(!!initialPromptData);
const [showHistoryModal, setShowHistoryModal] = useState(false);
// Construct versioned ID from prompt_id and version field
const getInitialVersionId = () => {
if (!initialPromptData?.prompt_spec) return undefined;
const baseId = initialPromptData.prompt_spec.prompt_id;
const version = initialPromptData.prompt_spec.version ||
(initialPromptData.prompt_spec.litellm_params as any)?.prompt_id;
// If version is a number, construct versioned ID
if (typeof version === 'number') {
return `${baseId}.v${version}`;
}
// If version is a string with version suffix, use it
if (typeof version === 'string' && (version.includes('.v') || version.includes('_v'))) {
return version;
}
return baseId;
};
const [activeVersionId, setActiveVersionId] = useState<string | undefined>(getInitialVersionId());
const [showToolModal, setShowToolModal] = useState(false);
const [showNameModal, setShowNameModal] = useState(false);
@ -124,6 +161,20 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
setShowToolModal(true);
};
const handleLoadVersion = (versionData: any) => {
try {
const loadedPrompt = parseExistingPrompt({ prompt_spec: versionData });
setPrompt(loadedPrompt);
// Store the version number or construct versioned ID for tracking
const versionNum = versionData.version || 1;
setActiveVersionId(`${versionData.prompt_id}.v${versionNum}`);
// NotificationsManager.success(`Loaded version v${versionNum}`);
} catch (error) {
console.error("Error loading version:", error);
NotificationsManager.fromBackend("Failed to load prompt version");
}
};
const handleSaveClick = () => {
if (!prompt.name || prompt.name.trim() === "" || prompt.name === "New prompt") {
setShowNameModal(true);
@ -160,19 +211,53 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
},
};
await createPromptCall(accessToken, promptData);
NotificationsManager.success("Prompt created successfully!");
if (editMode && initialPromptData?.prompt_spec?.prompt_id) {
await updatePromptCall(accessToken, initialPromptData.prompt_spec.prompt_id, promptData);
NotificationsManager.success("Prompt updated successfully!");
} else {
await createPromptCall(accessToken, promptData);
NotificationsManager.success("Prompt created successfully!");
}
onSuccess();
onClose();
} catch (error) {
console.error("Error saving prompt:", error);
NotificationsManager.fromBackend("Failed to save prompt");
NotificationsManager.fromBackend(editMode ? "Failed to update prompt" : "Failed to save prompt");
} finally {
setIsSaving(false);
setShowNameModal(false);
}
};
const getVersionNumber = (pid?: string) => {
if (!pid) return null;
if (pid.includes(".v")) {
return `v${pid.split(".v")[1]}`;
}
return null;
};
const currentVersion = getVersionNumber(activeVersionId);
// Extract template variables from prompt content for code examples
const extractTemplateVariables = (): Record<string, string> => {
const variables: Record<string, string> = {};
const allContent = [
prompt.developerMessage,
...prompt.messages.map(m => m.content)
].join(' ');
const variableRegex = /\{\{(\w+)\}\}/g;
let match;
while ((match = variableRegex.exec(allContent)) !== null) {
const varName = match[1];
if (!variables[varName]) {
variables[varName] = `example_${varName}`;
}
}
return variables;
};
return (
<div className="flex h-full bg-white">
<div className="flex-1 flex flex-col">
@ -182,10 +267,16 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
onBack={onClose}
onSave={handleSaveClick}
isSaving={isSaving}
editMode={editMode}
onShowHistory={() => setShowHistoryModal(true)}
version={currentVersion}
promptModel={prompt.model}
promptVariables={extractTemplateVariables()}
accessToken={accessToken}
/>
<div className="flex-1 flex overflow-hidden">
<div className="w-1/2 overflow-y-auto bg-white border-r border-gray-200">
<div className="w-1/2 overflow-y-auto bg-white border-r border-gray-200 flex-shrink-0">
<div className="border-b border-gray-200 bg-white px-6 py-4 flex items-center gap-3">
<ModelConfigCard
model={prompt.model}
@ -210,9 +301,7 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
<div className="ml-auto inline-flex items-center bg-gray-200 rounded-full p-0.5">
<button
className={`px-3 py-1 text-xs font-medium rounded-full transition-colors ${
viewMode === "pretty"
? "bg-white text-gray-900 shadow-sm"
: "text-gray-600"
viewMode === "pretty" ? "bg-white text-gray-900 shadow-sm" : "text-gray-600"
}`}
onClick={() => setViewMode("pretty")}
>
@ -220,9 +309,7 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
</button>
<button
className={`px-3 py-1 text-xs font-medium rounded-full transition-colors ${
viewMode === "dotprompt"
? "bg-white text-gray-900 shadow-sm"
: "text-gray-600"
viewMode === "dotprompt" ? "bg-white text-gray-900 shadow-sm" : "text-gray-600"
}`}
onClick={() => setViewMode("dotprompt")}
>
@ -258,7 +345,9 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
)}
</div>
<ConversationPanel />
<div className="w-1/2 flex-shrink-0">
<ConversationPanel prompt={prompt} accessToken={accessToken} />
</div>
</div>
</div>
@ -282,9 +371,17 @@ const PromptEditorView: React.FC<PromptEditorViewProps> = ({ onClose, onSuccess,
}}
/>
)}
<VersionHistorySidePanel
isOpen={showHistoryModal}
onClose={() => setShowHistoryModal(false)}
accessToken={accessToken}
promptId={initialPromptData?.prompt_spec?.prompt_id || prompt.name}
activeVersionId={activeVersionId}
onSelectVersion={handleLoadVersion}
/>
</div>
);
};
export default PromptEditorView;

View file

@ -26,5 +26,6 @@ export interface PromptEditorViewProps {
onClose: () => void;
onSuccess: () => void;
accessToken: string | null;
initialPromptData?: any;
}

View file

@ -1,4 +1,4 @@
import { PromptType } from "./types";
import { PromptType, Message, Tool } from "./types";
export const extractVariables = (prompt: PromptType): string[] => {
const variableSet = new Set<string>();
@ -74,3 +74,105 @@ export const convertToDotPrompt = (prompt: PromptType): string => {
return result.trim();
};
export const parseExistingPrompt = (apiResponse: any): PromptType => {
// Extract dotprompt_content from litellm_params
const dotpromptContent = apiResponse?.prompt_spec?.litellm_params?.dotprompt_content || "";
if (!dotpromptContent) {
throw new Error("No dotprompt_content found in API response");
}
// Split into frontmatter and content
const parts = dotpromptContent.split("---");
if (parts.length < 3) {
throw new Error("Invalid dotprompt format");
}
// Parse YAML frontmatter (parts[1])
const frontmatter = parts[1];
const content = parts.slice(2).join("---").trim();
// Extract metadata from frontmatter
const metadata: any = {};
frontmatter.split("\n").forEach((line: string) => {
const trimmedLine = line.trim();
if (trimmedLine && !trimmedLine.startsWith("input:") && !trimmedLine.startsWith("output:") && !trimmedLine.startsWith("schema:") && !trimmedLine.startsWith("format:")) {
const colonIndex = trimmedLine.indexOf(":");
if (colonIndex > 0) {
const key = trimmedLine.substring(0, colonIndex).trim();
const value = trimmedLine.substring(colonIndex + 1).trim();
if (key === "temperature" || key === "max_tokens" || key === "top_p") {
metadata[key] = parseFloat(value);
} else if (key === "model") {
metadata[key] = value;
}
}
}
});
// Parse content to extract developer message and user messages
let developerMessage = "";
const messages: Message[] = [];
const lines = content.split("\n");
let currentRole: "user" | "assistant" | null = null;
let currentContent = "";
for (const line of lines) {
if (line.startsWith("Developer:")) {
developerMessage = line.substring("Developer:".length).trim();
} else if (line.startsWith("User:")) {
if (currentRole && currentContent) {
messages.push({ role: currentRole, content: currentContent.trim() });
}
currentRole = "user";
currentContent = line.substring("User:".length).trim();
} else if (line.startsWith("Assistant:")) {
if (currentRole && currentContent) {
messages.push({ role: currentRole, content: currentContent.trim() });
}
currentRole = "assistant";
currentContent = line.substring("Assistant:".length).trim();
} else if (line.trim() && currentRole) {
currentContent += "\n" + line.trim();
}
}
// Add the last message
if (currentRole && currentContent) {
messages.push({ role: currentRole, content: currentContent.trim() });
}
// Parse tools from frontmatter if present
const tools: Tool[] = [];
// TODO: Add tool parsing if needed
// Strip version suffix from prompt name for display
const promptId = apiResponse?.prompt_spec?.prompt_id || "Unnamed Prompt";
const baseName = stripVersionFromPromptId(promptId) || promptId;
return {
name: baseName,
model: metadata.model || "gpt-4o",
config: {
temperature: metadata.temperature,
max_tokens: metadata.max_tokens,
top_p: metadata.top_p,
},
tools: tools,
developerMessage: developerMessage,
messages: messages.length > 0 ? messages : [{ role: "user", content: "Enter task specifics. Use {{template_variables}} for dynamic inputs" }],
};
};
export const getVersionNumber = (promptId?: string): string => {
if (!promptId) return "1";
// Match version with dot (.v), underscore (_v), or hyphen (-v) separator
const match = promptId.match(/[._-]v(\d+)$/);
return match ? match[1] : "1";
};
export const stripVersionFromPromptId = (promptId?: string): string => {
if (!promptId) return "";
// Remove version suffix with dot (.v), underscore (_v), or hyphen (-v) separator
return promptId.replace(/[._-]v\d+$/, "");
};

View file

@ -13,11 +13,18 @@ import {
TabPanels,
} from "@tremor/react";
import { Button, Modal } from "antd";
import { ArrowLeftIcon, TrashIcon } from "@heroicons/react/outline";
import { ArrowLeftIcon, TrashIcon, PencilIcon } from "@heroicons/react/outline";
import { getPromptInfo, PromptSpec, PromptTemplateBase, deletePromptCall } from "@/components/networking";
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils";
import { CheckIcon, CopyIcon } from "lucide-react";
import NotificationsManager from "../molecules/notifications_manager";
import PromptCodeSnippets from "./prompt_editor_view/PromptCodeSnippets";
import {
extractModel,
extractTemplateVariables,
getBasePromptId,
getCurrentVersion
} from "./prompt_utils";
export interface PromptInfoProps {
promptId: string;
@ -25,9 +32,10 @@ export interface PromptInfoProps {
accessToken: string | null;
isAdmin: boolean;
onDelete?: () => void;
onEdit?: (promptData: any) => void;
}
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin, onDelete }) => {
const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessToken, isAdmin, onDelete, onEdit }) => {
const [promptData, setPromptData] = useState<PromptSpec | null>(null);
const [promptTemplate, setPromptTemplate] = useState<PromptTemplateBase | null>(null);
const [rawApiResponse, setRawApiResponse] = useState<any>(null);
@ -90,8 +98,8 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
setIsDeleting(true);
try {
await deletePromptCall(accessToken, promptData.prompt_id);
NotificationsManager.success(`Prompt "${promptData.prompt_id}" deleted successfully`);
await deletePromptCall(accessToken, basePromptId);
NotificationsManager.success(`Prompt "${basePromptId}" deleted successfully`);
onDelete?.(); // Call the callback to refresh the parent component
onClose(); // Close the info view
} catch (error) {
@ -107,6 +115,11 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
setShowDeleteConfirm(false);
};
// Use utility functions to extract prompt data
const promptModel = promptData ? extractModel(promptData) || "gpt-4o" : "gpt-4o";
const basePromptId = getBasePromptId(promptData);
const currentVersion = getCurrentVersion(promptData);
return (
<div className="p-4">
<div>
@ -117,12 +130,12 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
<div>
<Title>Prompt Details</Title>
<div className="flex items-center cursor-pointer">
<Text className="text-gray-500 font-mono">{promptData.prompt_id}</Text>
<Text className="text-gray-500 font-mono">{basePromptId}</Text>
<Button
type="text"
size="small"
icon={copiedStates["prompt-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
onClick={() => copyToClipboard(promptData.prompt_id, "prompt-id")}
onClick={() => copyToClipboard(basePromptId, "prompt-id")}
className={`left-2 z-10 transition-all duration-200 ${
copiedStates["prompt-id"]
? "text-green-600 bg-green-50 border-green-200"
@ -131,6 +144,22 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
/>
</div>
</div>
<div className="flex gap-2">
<PromptCodeSnippets
promptId={basePromptId}
model={promptModel}
promptVariables={extractTemplateVariables(promptTemplate?.content)}
accessToken={accessToken}
version={currentVersion}
/>
<TremorButton
icon={PencilIcon}
variant="primary"
onClick={() => onEdit?.(rawApiResponse)}
className="flex items-center"
>
Prompt Studio
</TremorButton>
{isAdmin && (
<TremorButton
icon={TrashIcon}
@ -141,6 +170,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
Delete Prompt
</TremorButton>
)}
</div>
</div>
</div>
@ -159,7 +189,17 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
<Card>
<Text>Prompt ID</Text>
<div className="mt-2">
<Title className="font-mono text-sm">{promptData.prompt_id}</Title>
<Title className="font-mono text-sm">{basePromptId}</Title>
</div>
</Card>
<Card>
<Text>Version</Text>
<div className="mt-2">
<Title>{currentVersion}</Title>
<Badge color="blue" className="mt-1">
v{currentVersion}
</Badge>
</div>
</Card>
@ -251,7 +291,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
<div className="space-y-4">
<div>
<Text className="font-medium">Prompt ID</Text>
<div className="font-mono text-sm bg-gray-50 p-2 rounded">{promptData.prompt_id}</div>
<div className="font-mono text-sm bg-gray-50 p-2 rounded">{basePromptId}</div>
</div>
<div>
@ -332,7 +372,7 @@ const PromptInfoView: React.FC<PromptInfoProps> = ({ promptId, onClose, accessTo
okButtonProps={{ danger: true }}
>
<p>
Are you sure you want to delete prompt: <strong>{promptData?.prompt_id}</strong>?
Are you sure you want to delete prompt: <strong>{basePromptId}</strong>?
</p>
<p>This action cannot be undone.</p>
</Modal>

View file

@ -1,8 +1,9 @@
import React, { useState } from "react";
import React, { useState, useEffect } from "react";
import { Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Button } from "@tremor/react";
import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon, TrashIcon } from "@heroicons/react/outline";
import { Tooltip } from "antd";
import { PromptSpec } from "@/components/networking";
import { CopyOutlined } from "@ant-design/icons";
import { PromptSpec, modelHubCall } from "@/components/networking";
import {
ColumnDef,
flexRender,
@ -11,6 +12,8 @@ import {
SortingState,
useReactTable,
} from "@tanstack/react-table";
import { getProviderLogoAndName } from "@/components/provider_info_helpers";
import { extractModel, getProviderFromModelHub } from "./prompt_utils";
interface PromptTableProps {
promptsList: PromptSpec[];
@ -21,6 +24,12 @@ interface PromptTableProps {
isAdmin: boolean;
}
interface ModelGroupInfo {
model_group: string;
providers: string[];
[key: string]: any;
}
const PromptTable: React.FC<PromptTableProps> = ({
promptsList,
isLoading,
@ -30,6 +39,28 @@ const PromptTable: React.FC<PromptTableProps> = ({
isAdmin,
}) => {
const [sorting, setSorting] = useState<SortingState>([{ id: "created_at", desc: true }]);
const [modelHubData, setModelHubData] = useState<Map<string, ModelGroupInfo>>(new Map());
useEffect(() => {
const fetchModelHubData = async () => {
if (!accessToken) return;
try {
const response = await modelHubCall(accessToken);
if (response?.data) {
const modelMap = new Map<string, ModelGroupInfo>();
response.data.forEach((model: ModelGroupInfo) => {
modelMap.set(model.model_group, model);
});
setModelHubData(modelMap);
}
} catch (error) {
console.error("Error fetching model hub data:", error);
}
};
fetchModelHubData();
}, [accessToken]);
// Format date helper function
const formatDate = (dateString?: string) => {
@ -38,22 +69,96 @@ const PromptTable: React.FC<PromptTableProps> = ({
return date.toLocaleString();
};
const copyToClipboard = (text: string) => {
navigator.clipboard.writeText(text);
};
const columns: ColumnDef<PromptSpec>[] = [
{
header: "Prompt ID",
accessorKey: "prompt_id",
cell: (info: any) => (
<Tooltip title={String(info.getValue() || "")}>
<Button
size="xs"
variant="light"
className="font-mono text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs font-normal px-2 py-0.5 text-left overflow-hidden truncate max-w-[200px]"
onClick={() => info.getValue() && onPromptClick?.(info.getValue())}
>
{info.getValue() ? `${String(info.getValue()).slice(0, 7)}...` : ""}
</Button>
</Tooltip>
),
cell: (info: any) => {
const fullId = String(info.getValue() || "");
const displayId = fullId.length > 25 ? `${fullId.slice(0, 25)}...` : fullId;
return (
<div className="flex items-center gap-2">
<Tooltip title={fullId}>
<Button
size="xs"
variant="light"
className="font-mono text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs font-normal px-2 py-0.5 text-left overflow-hidden truncate min-w-[220px] justify-start"
onClick={() => info.getValue() && onPromptClick?.(info.getValue())}
>
{displayId}
</Button>
</Tooltip>
<Tooltip title="Copy prompt ID">
<CopyOutlined
onClick={(e) => {
e.stopPropagation();
copyToClipboard(fullId);
}}
className="cursor-pointer text-gray-500 hover:text-blue-500 text-xs"
/>
</Tooltip>
</div>
);
},
},
{
header: "Model",
accessorKey: "model",
cell: ({ row }) => {
const prompt = row.original;
const model = extractModel(prompt);
if (!model) {
return <span className="text-xs text-gray-400">-</span>;
}
const provider = getProviderFromModelHub(model, modelHubData);
const { logo } = getProviderLogoAndName(provider || "");
return (
<Tooltip title={model}>
<div className="flex items-center space-x-2">
{/* Provider Icon */}
<div className="flex-shrink-0">
{provider && logo ? (
<img
src={logo}
alt={`${provider} logo`}
className="w-4 h-4"
onError={(e) => {
const target = e.currentTarget as HTMLImageElement;
const parent = target.parentElement;
if (!parent || !parent.contains(target)) {
return;
}
try {
const fallbackDiv = document.createElement('div');
fallbackDiv.className = 'w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs';
fallbackDiv.textContent = provider?.charAt(0) || '-';
parent.replaceChild(fallbackDiv, target);
} catch (error) {
console.error('Failed to replace provider logo fallback:', error);
}
}}
/>
) : (
<div className="w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs">
-
</div>
)}
</div>
{/* Model Name */}
<span className="max-w-[15ch] truncate block">{model}</span>
</div>
</Tooltip>
);
},
},
{
header: "Created At",

View file

@ -0,0 +1,105 @@
import { PromptSpec } from "@/components/networking";
import { getVersionNumber } from "./prompt_editor_view/utils";
interface ModelGroupInfo {
model_group: string;
providers: string[];
[key: string]: any;
}
/**
* Extract template variables from prompt content
*/
export const extractTemplateVariables = (content?: string): Record<string, string> => {
if (!content) return {};
const variables: Record<string, string> = {};
const variableRegex = /\{\{(\w+)\}\}/g;
let match;
while ((match = variableRegex.exec(content)) !== null) {
const varName = match[1];
if (!variables[varName]) {
variables[varName] = `example_${varName}`;
}
}
return variables;
};
/**
* Get base prompt ID (stripped of version) from PromptSpec
*/
export const getBasePromptId = (promptData?: PromptSpec): string => {
return promptData?.prompt_id || "";
};
/**
* Get versioned prompt ID from litellm_params (preserves version)
*/
export const getVersionedPromptId = (promptData?: PromptSpec): string => {
const baseId = getBasePromptId(promptData);
const versionedId = (promptData?.litellm_params as any)?.prompt_id || baseId;
return versionedId;
};
/**
* Get current version number from prompt data
*/
export const getCurrentVersion = (promptData?: PromptSpec): string => {
// Use explicit version field if available (from API response)
if (promptData?.version) {
return String(promptData.version);
}
// Fallback: extract from versioned ID in litellm_params
const versionedId = getVersionedPromptId(promptData);
return getVersionNumber(versionedId);
};
/**
* Extract model from prompt litellm_params
*/
export const extractModel = (prompt: PromptSpec): string | null => {
try {
const params = prompt.litellm_params as any;
// Try to extract from dotprompt_content
if (params?.dotprompt_content) {
const match = params.dotprompt_content.match(/model:\s*([^\n]+)/);
if (match) return match[1].trim();
}
// Try to extract from prompt_data
if (params?.prompt_data?.model) {
return params.prompt_data.model;
}
// Try to extract model from litellm_params directly
if (params?.model) {
return params.model;
}
return null;
} catch (error) {
console.error("Error extracting model:", error);
return null;
}
};
/**
* Get provider from model hub data
*/
export const getProviderFromModelHub = (
modelName: string | null,
modelHubData: Map<string, ModelGroupInfo>
): string | null => {
if (!modelName) return null;
const modelInfo = modelHubData.get(modelName);
if (modelInfo && modelInfo.providers && modelInfo.providers.length > 0) {
// Return the first provider from the list
return modelInfo.providers[0];
}
return null;
};

View file

@ -0,0 +1,68 @@
import { describe, it, expect, vi, beforeAll, beforeEach } from "vitest";
import { render } from "@testing-library/react";
import PublicModelHub from "./public_model_hub";
import { FeatureFlagsProvider } from "@/hooks/useFeatureFlags";
vi.mock("next/navigation", () => ({
useRouter: vi.fn(() => ({
replace: vi.fn(),
push: vi.fn(),
refresh: vi.fn(),
})),
}));
vi.mock("./networking", async (importOriginal) => {
const actual = await importOriginal<typeof import("./networking")>();
return {
...actual,
modelHubPublicModelsCall: vi.fn().mockResolvedValue([]),
getPublicModelHubInfo: vi.fn().mockResolvedValue({
docs_title: "LiteLLM Gateway",
custom_docs_description: null,
litellm_version: "1.0.0",
useful_links: {},
}),
agentHubPublicModelsCall: vi.fn().mockResolvedValue([]),
mcpHubPublicServersCall: vi.fn().mockResolvedValue([]),
getUiConfig: vi.fn().mockResolvedValue({}),
};
});
beforeAll(() => {
Object.defineProperty(window, "matchMedia", {
writable: true,
value: (query: string) => ({
matches: false,
media: query,
onchange: null,
addListener: () => {},
removeListener: () => {},
addEventListener: () => {},
removeEventListener: () => {},
dispatchEvent: () => false,
}),
});
});
beforeEach(() => {
Storage.prototype.getItem = vi.fn(() => "false");
Storage.prototype.setItem = vi.fn();
Object.defineProperty(window, "location", {
writable: true,
value: {
pathname: "/",
origin: "http://localhost:3000",
},
});
});
describe("PublicModelHub", () => {
it("renders", () => {
const { container } = render(
<FeatureFlagsProvider>
<PublicModelHub />
</FeatureFlagsProvider>,
);
expect(container).toBeInTheDocument();
});
});

View file

@ -1,5 +1,11 @@
import React, { useEffect, useState, useRef, useMemo } from "react";
import { modelHubPublicModelsCall, getPublicModelHubInfo, agentHubPublicModelsCall, mcpHubPublicServersCall } from "./networking";
import {
modelHubPublicModelsCall,
getPublicModelHubInfo,
agentHubPublicModelsCall,
mcpHubPublicServersCall,
getUiConfig,
} from "./networking";
import { ModelDataTable } from "./model_dashboard/table";
import { ColumnDef } from "@tanstack/react-table";
import { Card, Text, Title, Button } from "@tremor/react";
@ -117,60 +123,72 @@ const PublicModelHub: React.FC<PublicModelHubProps> = ({ accessToken, isEmbedded
const mcpTableRef = useRef<TableInstance<any>>(null);
useEffect(() => {
const fetchPublicData = async () => {
const initializeAndFetch = async () => {
// Initialize proxyBaseUrl first to ensure it includes the server root path
try {
setLoading(true);
const _modelHubData = await modelHubPublicModelsCall();
console.log("ModelHubData:", _modelHubData);
setModelHubData(_modelHubData);
await getUiConfig();
} catch (error) {
console.error("There was an error fetching the public model data", error);
setServiceStatus("Service unavailable");
} finally {
setLoading(false);
console.error("Failed to get UI config:", error);
// Continue anyway - might work with default proxyBaseUrl
}
const fetchPublicData = async () => {
try {
setLoading(true);
const _modelHubData = await modelHubPublicModelsCall();
console.log("ModelHubData:", _modelHubData);
setModelHubData(_modelHubData);
} catch (error) {
console.error("There was an error fetching the public model data", error);
setServiceStatus("Service unavailable");
} finally {
setLoading(false);
}
};
const fetchAgentData = async () => {
try {
setAgentLoading(true);
const _agentHubData = await agentHubPublicModelsCall();
console.log("AgentHubData:", _agentHubData);
setAgentHubData(_agentHubData);
} catch (error) {
console.error("There was an error fetching the public agent data", error);
} finally {
setAgentLoading(false);
}
};
const fetchMcpData = async () => {
try {
setMcpLoading(true);
const _mcpHubData = await mcpHubPublicServersCall();
console.log("MCPHubData:", _mcpHubData);
setMcpHubData(_mcpHubData);
} catch (error) {
console.error("There was an error fetching the public MCP server data", error);
} finally {
setMcpLoading(false);
}
};
const fetchPublicModelHubInfo = async () => {
const publicModelHubInfo = await getPublicModelHubInfo();
console.log("Public Model Hub Info:", publicModelHubInfo);
setPageTitle(publicModelHubInfo.docs_title);
setCustomDocsDescription(publicModelHubInfo.custom_docs_description);
setLitellmVersion(publicModelHubInfo.litellm_version);
setUsefulLinks(publicModelHubInfo.useful_links || {});
};
fetchPublicModelHubInfo();
fetchPublicData();
fetchAgentData();
fetchMcpData();
};
const fetchAgentData = async () => {
try {
setAgentLoading(true);
const _agentHubData = await agentHubPublicModelsCall();
console.log("AgentHubData:", _agentHubData);
setAgentHubData(_agentHubData);
} catch (error) {
console.error("There was an error fetching the public agent data", error);
} finally {
setAgentLoading(false);
}
};
const fetchMcpData = async () => {
try {
setMcpLoading(true);
const _mcpHubData = await mcpHubPublicServersCall();
console.log("MCPHubData:", _mcpHubData);
setMcpHubData(_mcpHubData);
} catch (error) {
console.error("There was an error fetching the public MCP server data", error);
} finally {
setMcpLoading(false);
}
};
const fetchPublicModelHubInfo = async () => {
const publicModelHubInfo = await getPublicModelHubInfo();
console.log("Public Model Hub Info:", publicModelHubInfo);
setPageTitle(publicModelHubInfo.docs_title);
setCustomDocsDescription(publicModelHubInfo.custom_docs_description);
setLitellmVersion(publicModelHubInfo.litellm_version);
setUsefulLinks(publicModelHubInfo.useful_links || {});
};
fetchPublicModelHubInfo();
fetchPublicData();
fetchAgentData();
fetchMcpData();
initializeAndFetch();
}, []);
// Clear filters when filter values change to avoid confusion
@ -400,8 +418,7 @@ const PublicModelHub: React.FC<PublicModelHubProps> = ({ accessToken, isEmbedded
// Apply transport filters
return searchResults.filter((server) => {
const matchesTransport =
selectedMcpTransports.length === 0 || selectedMcpTransports.includes(server.transport);
const matchesTransport = selectedMcpTransports.length === 0 || selectedMcpTransports.includes(server.transport);
return matchesTransport;
});
@ -1183,10 +1200,7 @@ const PublicModelHub: React.FC<PublicModelHubProps> = ({ accessToken, isEmbedded
<div>
<div className="flex items-center space-x-2 mb-3">
<Text className="text-sm font-medium text-gray-700">Search MCP Servers:</Text>
<Tooltip
title="Search MCP servers by name or description"
placement="top"
>
<Tooltip title="Search MCP servers by name or description" placement="top">
<Info className="w-4 h-4 text-gray-400 cursor-help" />
</Tooltip>
</div>
@ -1842,9 +1856,7 @@ print(response.model_dump(mode='json', exclude_none=True))`;
<div>
<Text className="text-lg font-semibold mb-4">Additional Information</Text>
<div className="bg-gray-50 p-4 rounded-lg">
<pre className="text-xs overflow-x-auto">
{JSON.stringify(selectedMcpServer.mcp_info, null, 2)}
</pre>
<pre className="text-xs overflow-x-auto">{JSON.stringify(selectedMcpServer.mcp_info, null, 2)}</pre>
</div>
</div>
)}
@ -1854,7 +1866,7 @@ print(response.model_dump(mode='json', exclude_none=True))`;
<Text className="text-lg font-semibold mb-4">Usage Example</Text>
<div className="bg-gray-900 text-gray-100 p-4 rounded-lg overflow-x-auto">
<pre className="text-sm">
{`# Using MCP Server with Python FastMCP
{`# Using MCP Server with Python FastMCP
from fastmcp import Client
import asyncio