mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge remote-tracking branch 'origin' into litellm_ui_callback_fix
This commit is contained in:
commit
5dad3c9708
82 changed files with 5501 additions and 550 deletions
|
|
@ -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 && \
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
|
|
|||
277
docs/my-website/docs/providers/docker_model_runner.md
Normal file
277
docs/my-website/docs/providers/docker_model_runner.md
Normal 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/).
|
||||
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
:::
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
},
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
144
litellm/llms/docker_model_runner/chat/transformation.py
Normal file
144
litellm/llms/docker_model_runner/chat/transformation.py
Normal 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
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
10
litellm/types/proxy/prompt_endpoints.py
Normal file
10
litellm/types/proxy/prompt_endpoints.py
Normal 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
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ test("admin login test", async ({ page }) => {
|
|||
await loginButton.click();
|
||||
const tabs = [
|
||||
"Virtual Keys",
|
||||
"Test Key",
|
||||
"Playground",
|
||||
"Models",
|
||||
"Usage",
|
||||
"Teams",
|
||||
|
|
|
|||
134
tests/proxy_unit_tests/test_prompt_test_endpoint.py
Normal file
134
tests/proxy_unit_tests/test_prompt_test_endpoint.py
Normal 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
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
50
tests/search_tests/test_search_tool_name_filtering.py
Normal file
50
tests/search_tests/test_search_tool_name_filtering.py
Normal 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
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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?"
|
||||
|
||||
|
|
@ -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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
307
tests/test_litellm/proxy/prompts/test_prompt_endpoints.py
Normal file
307
tests/test_litellm/proxy/prompts/test_prompt_endpoints.py
Normal 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
|
||||
|
||||
|
|
@ -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__])
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
/>
|
||||
) : (
|
||||
<>
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
@ -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,
|
||||
};
|
||||
};
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -26,5 +26,6 @@ export interface PromptEditorViewProps {
|
|||
onClose: () => void;
|
||||
onSuccess: () => void;
|
||||
accessToken: string | null;
|
||||
initialPromptData?: any;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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+$/, "");
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
105
ui/litellm-dashboard/src/components/prompts/prompt_utils.tsx
Normal file
105
ui/litellm-dashboard/src/components/prompts/prompt_utils.tsx
Normal 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;
|
||||
};
|
||||
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue