mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-23 12:24:08 +00:00
@@ -69,9 +69,11 @@ jobs:
|
||||
- run:
|
||||
name: Install Python
|
||||
command: |
|
||||
choco install python --version=3.11.0 -y
|
||||
choco install python --version=3.11.0 -y --no-progress --force
|
||||
refreshenv
|
||||
python --version
|
||||
environment:
|
||||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
command: |
|
||||
|
||||
@@ -19,6 +19,7 @@ jobs:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
defaults:
|
||||
run:
|
||||
working-directory: enterprise
|
||||
@@ -56,14 +57,33 @@ jobs:
|
||||
- name: Build
|
||||
run: poetry build
|
||||
|
||||
- name: Commit version bump
|
||||
- name: Commit version bump and create PR
|
||||
id: create-pr
|
||||
run: |
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
cd ..
|
||||
BRANCH="bump/enterprise-${{ steps.bump.outputs.new }}"
|
||||
git checkout -b "$BRANCH"
|
||||
git add enterprise/pyproject.toml pyproject.toml requirements.txt poetry.lock
|
||||
git commit -m "bump: litellm-enterprise ${{ steps.bump.outputs.old }} → ${{ steps.bump.outputs.new }}"
|
||||
git push
|
||||
git push origin "$BRANCH" --force
|
||||
gh pr create \
|
||||
--title "bump: litellm-enterprise ${{ steps.bump.outputs.old }} → ${{ steps.bump.outputs.new }}" \
|
||||
--body "Version bump for litellm-enterprise. Merge to update main." \
|
||||
--head "$BRANCH" \
|
||||
--base main \
|
||||
|| true
|
||||
PR_URL=$(gh pr list --head "$BRANCH" --json url -q '.[0].url')
|
||||
echo "pr_url=$PR_URL" >> $GITHUB_OUTPUT
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
|
||||
- name: Enable auto-merge
|
||||
run: |
|
||||
gh pr merge "${{ steps.create-pr.outputs.pr_url }}" --auto --squash
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
|
||||
- name: Publish to PyPI
|
||||
env:
|
||||
|
||||
@@ -355,7 +355,7 @@ router_settings:
|
||||
| set_verbose | boolean | [DEPRECATED PARAM - see debug docs](./debugging) If true, sets the logging level to verbose. |
|
||||
| retry_after | int | Time to wait before retrying a request in seconds. Defaults to 0. If `x-retry-after` is received from LLM API, this value is overridden. |
|
||||
| provider_budget_config | ProviderBudgetConfig | Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None. [Further Docs](./provider_budget_routing.md) |
|
||||
| enable_pre_call_checks | boolean | If true, checks if a call is within the model's context window before making the call. [More information here](reliability) |
|
||||
| enable_pre_call_checks | boolean | If true, checks if a call is within the model's context window before making the call. **Required** for `model_info.max_input_tokens` enforcement. Default: false. [More information here](reliability) |
|
||||
| model_group_retry_policy | Dict[str, RetryPolicy] | [SDK-only arg] Set retry policy for model groups. |
|
||||
| context_window_fallbacks | List[Dict[str, List[str]]] | Fallback models for context window violations. |
|
||||
| redis_url | str | URL for Redis server. **Known performance issue with Redis URL.** |
|
||||
@@ -804,6 +804,7 @@ router_settings:
|
||||
| PYROSCOPE_SERVER_ADDRESS | Pyroscope server URL to send profiles to. Required when LITELLM_ENABLE_PYROSCOPE is true. No default.
|
||||
| PYROSCOPE_SAMPLE_RATE | Optional. Sample rate for Pyroscope profiling (integer). No default; when unset, the pyroscope-io library default is used.
|
||||
| LITELLM_MASTER_KEY | Master key for proxy authentication
|
||||
| LITELLM_MAX_BUDGET_PER_SESSION_TTL | TTL in seconds for session budget counters used by the max-budget-per-session limiter. Default is 3600 (1 hour)
|
||||
| LITELLM_MAX_ITERATIONS_TTL | TTL in seconds for session iteration counters used by the max-iterations limiter. Default is 3600 (1 hour)
|
||||
| LITELLM_MAX_STREAMING_DURATION_SECONDS | Maximum duration in seconds allowed for a streaming response. Streams exceeding this duration are terminated with a Timeout error. Default is None (no limit)
|
||||
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
|
||||
|
||||
@@ -713,6 +713,34 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
|
||||
|
||||
[**See Code**](https://github.com/BerriAI/litellm/blob/c9e6b05cfb20dfb17272218e2555d6b496c47f6f/litellm/router.py#L2163)
|
||||
|
||||
:::important
|
||||
**`enable_pre_call_checks` is required** for context-window enforcement. Without it, requests are sent to the provider regardless of input token count. Set `enable_pre_call_checks: true` in `router_settings` in your config.
|
||||
:::
|
||||
|
||||
#### Custom max_input_tokens per deployment
|
||||
|
||||
You can override the default context limit for a deployment by setting `max_input_tokens` in `model_info`. This is useful for testing, rate-limiting long prompts, or enforcing stricter limits than the provider's default.
|
||||
|
||||
**Both** of the following are required:
|
||||
|
||||
1. **`router_settings.enable_pre_call_checks: true`** — enables pre-call checks
|
||||
2. **`model_info.max_input_tokens`** on the deployment — overrides the limit for that model
|
||||
|
||||
```yaml
|
||||
router_settings:
|
||||
enable_pre_call_checks: true # Required for enforcement
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
model_info:
|
||||
max_input_tokens: 10 # Override: reject prompts > 10 tokens
|
||||
```
|
||||
|
||||
If a request exceeds the limit, LiteLLM raises `ContextWindowExceededError` with details like `Model=gpt-4o, Max Input Tokens=10, Got=306`.
|
||||
|
||||
**1. Setup config**
|
||||
|
||||
For azure deployments, set the base model. Pick the base model from [this list](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json), all the azure models start with azure/.
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
| Feature | Supported |
|
||||
|---------|-----------|
|
||||
| Supported Providers | `perplexity`, `tavily`, `parallel_ai`, `exa_ai`, `brave`, `google_pse`, `dataforseo`, `firecrawl`, `searxng`, `linkup`, `duckduckgo`, `searchapi` |
|
||||
| Supported Providers | `perplexity`, `tavily`, `parallel_ai`, `exa_ai`, `brave`, `google_pse`, `dataforseo`, `firecrawl`, `searxng`, `linkup`, `duckduckgo`, `searchapi`, `serper` |
|
||||
| Cost Tracking | ✅ |
|
||||
| Logging | ✅ |
|
||||
| Load Balancing | ❌ |
|
||||
@@ -210,7 +210,7 @@ See the [official Perplexity Search documentation](https://docs.perplexity.ai/ap
|
||||
| Parameter | Type | Required | Description |
|
||||
|-----------|------|----------|-------------|
|
||||
| `query` | string or array | Yes | Search query. Can be a single string or array of strings |
|
||||
| `search_provider` | string | Yes (SDK) | The search provider to use: `"perplexity"`, `"tavily"`, `"parallel_ai"`, `"exa_ai"`, `"brave"`, `"google_pse"`, `"dataforseo"`, `"firecrawl"`, `"searxng"`, `"linkup"`, `"duckduckgo"`, or `"searchapi"` |
|
||||
| `search_provider` | string | Yes (SDK) | The search provider to use: `"perplexity"`, `"tavily"`, `"parallel_ai"`, `"exa_ai"`, `"brave"`, `"google_pse"`, `"dataforseo"`, `"firecrawl"`, `"searxng"`, `"linkup"`, `"duckduckgo"`, `"searchapi"`, or `"serper"` |
|
||||
| `search_tool_name` | string | Yes (Proxy) | Name of the search tool configured in `config.yaml` |
|
||||
| `max_results` | integer | No | Maximum number of results to return (1-20). Default: 10 |
|
||||
| `search_domain_filter` | array | No | List of domains to filter results (max 20 domains) |
|
||||
@@ -276,6 +276,7 @@ The response follows Perplexity's search format with the following structure:
|
||||
| Firecrawl | `FIRECRAWL_API_KEY` | `firecrawl` |
|
||||
| SearXNG | `SEARXNG_API_BASE` (required) | `searxng` |
|
||||
| Linkup | `LINKUP_API_KEY` | `linkup` |
|
||||
| Serper | `SERPER_API_KEY` | `serper` |
|
||||
| DuckDuckGo | `DUCKDUCKGO_API_BASE` | `duckduckgo` |
|
||||
| SearchAPI.io | `SEARCHAPI_API_KEY` | `searchapi` |
|
||||
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
# Serper Search
|
||||
|
||||
**Get API Key:** [https://serper.dev](https://serper.dev)
|
||||
|
||||
## LiteLLM Python SDK
|
||||
|
||||
```python showLineNumbers title="Serper Search"
|
||||
import os
|
||||
from litellm import search
|
||||
|
||||
os.environ["SERPER_API_KEY"] = "your-api-key"
|
||||
|
||||
response = search(
|
||||
query="latest AI developments",
|
||||
search_provider="serper",
|
||||
max_results=5
|
||||
)
|
||||
```
|
||||
|
||||
## LiteLLM AI Gateway
|
||||
|
||||
### 1. Setup config.yaml
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-5
|
||||
litellm_params:
|
||||
model: gpt-5
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
search_tools:
|
||||
- search_tool_name: serper-search
|
||||
litellm_params:
|
||||
search_provider: serper
|
||||
api_key: os.environ/SERPER_API_KEY
|
||||
```
|
||||
|
||||
### 2. Start the proxy
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
### 3. Test the search endpoint
|
||||
|
||||
```bash showLineNumbers title="Test Request"
|
||||
curl http://0.0.0.0:4000/v1/search/serper-search \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"query": "latest AI developments",
|
||||
"max_results": 5
|
||||
}'
|
||||
```
|
||||
|
||||
## Provider-specific Parameters
|
||||
|
||||
```python showLineNumbers title="Serper Search with Provider-specific Parameters"
|
||||
import os
|
||||
from litellm import search
|
||||
|
||||
os.environ["SERPER_API_KEY"] = "your-api-key"
|
||||
|
||||
response = search(
|
||||
query="latest tech news",
|
||||
search_provider="serper",
|
||||
max_results=10,
|
||||
# Serper-specific parameters
|
||||
gl="us", # Country/geolocation code
|
||||
hl="en", # Language code
|
||||
autocorrect=False, # Disable autocorrect
|
||||
tbs="qdr:d", # Time filter: past day ('qdr:h' hour, 'qdr:w' week, 'qdr:m' month)
|
||||
page=2 # Page number
|
||||
)
|
||||
```
|
||||
@@ -684,6 +684,7 @@ const sidebars = {
|
||||
"search/firecrawl",
|
||||
"search/searxng",
|
||||
"search/linkup",
|
||||
"search/serper",
|
||||
]
|
||||
},
|
||||
"skills",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.33"
|
||||
version = "0.1.34"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.52"
|
||||
version = "0.4.53"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
@@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.52"
|
||||
version = "0.4.53"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
||||
@@ -2,7 +2,7 @@ import asyncio
|
||||
import json
|
||||
import time
|
||||
import traceback
|
||||
from typing import Dict, Iterable, List, Literal, Optional, Tuple, Union
|
||||
from typing import Dict, Iterable, List, Literal, Optional, Tuple, Union, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
@@ -13,6 +13,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
from litellm.types.llms.databricks import DatabricksTool
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionThinkingBlock,
|
||||
ImageURLListItem,
|
||||
OpenAIModerationResponse,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
@@ -26,13 +27,13 @@ from litellm.types.utils import (
|
||||
Function,
|
||||
HiddenParams,
|
||||
ImageResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
)
|
||||
from litellm.types.utils import Logprobs as TextCompletionLogprobs
|
||||
from litellm.types.utils import (
|
||||
Message,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
PromptTokensDetailsWrapper,
|
||||
RerankResponse,
|
||||
StreamingChoices,
|
||||
TextChoices,
|
||||
@@ -52,6 +53,24 @@ _MODEL_RESPONSE_FIELDS: frozenset = frozenset(ModelResponse.model_fields.keys())
|
||||
}
|
||||
|
||||
|
||||
def _normalize_images_for_message(
|
||||
images: Optional[List[dict]],
|
||||
) -> Optional[List[ImageURLListItem]]:
|
||||
"""
|
||||
Ensure each image has an 'index' field, as required by ImageURLListItem.
|
||||
Some providers (e.g. OpenRouter) return images without index.
|
||||
"""
|
||||
if not images:
|
||||
return cast(Optional[List[ImageURLListItem]], images)
|
||||
normalized: List[ImageURLListItem] = []
|
||||
for i, img in enumerate(images):
|
||||
if isinstance(img, dict) and "index" not in img:
|
||||
normalized.append(cast(ImageURLListItem, {**img, "index": i}))
|
||||
else:
|
||||
normalized.append(cast(ImageURLListItem, img))
|
||||
return normalized
|
||||
|
||||
|
||||
def _safe_convert_created_field(created_value) -> int:
|
||||
"""
|
||||
Safely convert a 'created' field value to an integer.
|
||||
@@ -591,7 +610,9 @@ def convert_to_model_response_object( # noqa: PLR0915
|
||||
reasoning_content=reasoning_content,
|
||||
thinking_blocks=thinking_blocks,
|
||||
annotations=choice["message"].get("annotations", None),
|
||||
images=choice["message"].get("images", None),
|
||||
images=_normalize_images_for_message(
|
||||
choice["message"].get("images", None)
|
||||
),
|
||||
)
|
||||
finish_reason = choice.get("finish_reason", None)
|
||||
if finish_reason is None:
|
||||
|
||||
@@ -73,6 +73,53 @@ def _redact_responses_api_output(output_items):
|
||||
summary_item.text = "redacted-by-litellm"
|
||||
|
||||
|
||||
def _redact_standard_logging_object(model_call_details: dict):
|
||||
"""Redact messages and response inside standard_logging_object if present."""
|
||||
standard_logging_object = model_call_details.get("standard_logging_object")
|
||||
if standard_logging_object is None:
|
||||
return
|
||||
|
||||
redacted_str = "redacted-by-litellm"
|
||||
|
||||
if standard_logging_object.get("messages") is not None:
|
||||
standard_logging_object["messages"] = [
|
||||
{"role": "user", "content": redacted_str}
|
||||
]
|
||||
|
||||
response = standard_logging_object.get("response")
|
||||
if response is not None:
|
||||
if isinstance(response, dict) and "output" in response:
|
||||
# ResponsesAPIResponse format - redact content in output items
|
||||
if isinstance(response.get("output"), list):
|
||||
for output_item in response["output"]:
|
||||
if isinstance(output_item, dict) and "content" in output_item:
|
||||
if isinstance(output_item["content"], list):
|
||||
for content_item in output_item["content"]:
|
||||
if (
|
||||
isinstance(content_item, dict)
|
||||
and "text" in content_item
|
||||
):
|
||||
content_item["text"] = redacted_str
|
||||
elif isinstance(response, dict) and "choices" in response:
|
||||
# ModelResponse dict format - redact content in choices
|
||||
if isinstance(response.get("choices"), list):
|
||||
for choice in response["choices"]:
|
||||
if isinstance(choice, dict):
|
||||
if "message" in choice and isinstance(choice["message"], dict):
|
||||
choice["message"]["content"] = redacted_str
|
||||
if "audio" in choice["message"]:
|
||||
choice["message"]["audio"] = None
|
||||
elif "delta" in choice and isinstance(choice["delta"], dict):
|
||||
choice["delta"]["content"] = redacted_str
|
||||
if "audio" in choice["delta"]:
|
||||
choice["delta"]["audio"] = None
|
||||
elif isinstance(response, str):
|
||||
standard_logging_object["response"] = redacted_str
|
||||
else:
|
||||
# For other formats (empty dict, None, etc.), use simple text format
|
||||
standard_logging_object["response"] = {"text": redacted_str}
|
||||
|
||||
|
||||
def perform_redaction(model_call_details: dict, result):
|
||||
"""
|
||||
Performs the actual redaction on the logging object and result.
|
||||
@@ -114,6 +161,29 @@ def perform_redaction(model_call_details: dict, result):
|
||||
if hasattr(_result, "choices") and _result.choices is not None:
|
||||
for choice in _result.choices:
|
||||
_redact_choice_content(choice)
|
||||
elif isinstance(_result, dict) and "choices" in _result:
|
||||
# Handle dict representation of ModelResponse (e.g., from model_dump())
|
||||
if _result.get("choices") is not None:
|
||||
for choice in _result["choices"]:
|
||||
if isinstance(choice, dict):
|
||||
if "message" in choice and isinstance(choice["message"], dict):
|
||||
choice["message"]["content"] = "redacted-by-litellm"
|
||||
if "reasoning_content" in choice["message"]:
|
||||
choice["message"]["reasoning_content"] = "redacted-by-litellm"
|
||||
if "thinking_blocks" in choice["message"]:
|
||||
choice["message"]["thinking_blocks"] = None
|
||||
if "audio" in choice["message"]:
|
||||
choice["message"]["audio"] = None
|
||||
elif "delta" in choice and isinstance(choice["delta"], dict):
|
||||
choice["delta"]["content"] = "redacted-by-litellm"
|
||||
if "reasoning_content" in choice["delta"]:
|
||||
choice["delta"]["reasoning_content"] = "redacted-by-litellm"
|
||||
if "thinking_blocks" in choice["delta"]:
|
||||
choice["delta"]["thinking_blocks"] = None
|
||||
if "audio" in choice["delta"]:
|
||||
choice["delta"]["audio"] = None
|
||||
else:
|
||||
_redact_choice_content(choice)
|
||||
elif isinstance(_result, litellm.ResponsesAPIResponse):
|
||||
if hasattr(_result, "output"):
|
||||
_redact_responses_api_output(_result.output)
|
||||
|
||||
@@ -334,24 +334,67 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
||||
"""
|
||||
Parse direct JSON response (non-streaming).
|
||||
|
||||
JSON response structure:
|
||||
{
|
||||
"result": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "..."}]
|
||||
}
|
||||
}
|
||||
Supports multiple agent response schemas:
|
||||
1. {"result": {"role": "assistant", "content": [{"text": "..."}]}} - standard AgentCore
|
||||
2. {"response": [{"text": "..."}]} - Strands agent format
|
||||
3. {"result": "plain text"} or {"response": "plain text"} - simple string
|
||||
4. Fallback: raw JSON as content string
|
||||
"""
|
||||
result = response_json.get("result", {})
|
||||
# Guard: if json.loads() returned a non-dict (e.g. array or primitive),
|
||||
# skip strategy matching and fall back to raw JSON string
|
||||
if not isinstance(response_json, dict):
|
||||
verbose_logger.warning(
|
||||
"AgentCore: JSON response is not a dict. "
|
||||
"Returning raw JSON as content."
|
||||
)
|
||||
return AgentCoreParsedResponse(
|
||||
content=json.dumps(response_json),
|
||||
usage=None,
|
||||
final_message=None,
|
||||
)
|
||||
|
||||
# Extract content using the same helper as SSE parsing
|
||||
content = self._extract_content_from_message(result) # type: ignore
|
||||
# Strategy 1: {"result": {"content": [{"text": "..."}]}} - standard AgentCore format
|
||||
if "result" in response_json and isinstance(response_json["result"], dict):
|
||||
result = response_json["result"]
|
||||
content = self._extract_content_from_message(result) # type: ignore
|
||||
return AgentCoreParsedResponse(
|
||||
content=content,
|
||||
usage=None,
|
||||
final_message=result, # type: ignore
|
||||
)
|
||||
|
||||
# JSON responses don't include usage data
|
||||
# Strategy 2: {"response": [{"text": "..."}]} - Strands agent content blocks
|
||||
if "response" in response_json and isinstance(
|
||||
response_json["response"], list
|
||||
):
|
||||
content = self._extract_content_from_message(
|
||||
{"content": response_json["response"]} # type: ignore
|
||||
)
|
||||
return AgentCoreParsedResponse(
|
||||
content=content,
|
||||
usage=None,
|
||||
final_message=None,
|
||||
)
|
||||
|
||||
# Strategy 3: string values - {"result": "text"} or {"response": "text"}
|
||||
for key in ("result", "response"):
|
||||
val = response_json.get(key)
|
||||
if isinstance(val, str):
|
||||
return AgentCoreParsedResponse(
|
||||
content=val,
|
||||
usage=None,
|
||||
final_message=None,
|
||||
)
|
||||
|
||||
# Strategy 4: fallback - return raw JSON as content
|
||||
verbose_logger.warning(
|
||||
f"AgentCore: Could not extract content from JSON response keys "
|
||||
f"{list(response_json.keys())}. Returning raw JSON as content."
|
||||
)
|
||||
return AgentCoreParsedResponse(
|
||||
content=content,
|
||||
content=json.dumps(response_json),
|
||||
usage=None,
|
||||
final_message=result, # type: ignore
|
||||
final_message=None,
|
||||
)
|
||||
|
||||
def _get_parsed_response(
|
||||
@@ -589,7 +632,64 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
# Wrap the generator in CustomStreamWrapper
|
||||
# Check if response is JSON (agent used sync return) instead of SSE
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
if "application/json" in content_type:
|
||||
verbose_logger.debug(
|
||||
"AgentCore streaming: received JSON response instead of SSE, "
|
||||
"converting to single-chunk stream"
|
||||
)
|
||||
try:
|
||||
body = response.read()
|
||||
response_json = json.loads(body)
|
||||
except (json.JSONDecodeError, Exception) as e:
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=f"AgentCore: Failed to read/parse JSON response body: {e}",
|
||||
)
|
||||
parsed = self._parse_json_response(response_json)
|
||||
|
||||
def _json_as_sync_stream():
|
||||
# Content chunk
|
||||
content_chunk = ModelResponseStream(
|
||||
id=f"chatcmpl-{uuid.uuid4()}",
|
||||
created=0,
|
||||
model=model,
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
content_chunk.choices = [
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content=parsed["content"], role="assistant"),
|
||||
)
|
||||
]
|
||||
yield content_chunk
|
||||
|
||||
# Stop sentinel chunk (matches SSE path convention)
|
||||
stop_chunk = ModelResponseStream(
|
||||
id=f"chatcmpl-{uuid.uuid4()}",
|
||||
created=0,
|
||||
model=model,
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
stop_chunk.choices = [
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(),
|
||||
)
|
||||
]
|
||||
yield stop_chunk
|
||||
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=_json_as_sync_stream(),
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
# SSE stream (text/event-stream or default) - use existing SSE parser
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=self._stream_agentcore_response_sync(response, model),
|
||||
model=model,
|
||||
@@ -746,7 +846,64 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
# Wrap the async generator in CustomStreamWrapper
|
||||
# Check if response is JSON (agent used sync return) instead of SSE
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
if "application/json" in content_type:
|
||||
verbose_logger.debug(
|
||||
"AgentCore streaming: received JSON response instead of SSE, "
|
||||
"converting to single-chunk stream"
|
||||
)
|
||||
try:
|
||||
body = await response.aread()
|
||||
response_json = json.loads(body)
|
||||
except (json.JSONDecodeError, Exception) as e:
|
||||
raise BedrockError(
|
||||
status_code=response.status_code,
|
||||
message=f"AgentCore: Failed to read/parse JSON response body: {e}",
|
||||
)
|
||||
parsed = self._parse_json_response(response_json)
|
||||
|
||||
async def _json_as_async_stream() -> AsyncGenerator[ModelResponseStream, None]:
|
||||
# Content chunk
|
||||
content_chunk = ModelResponseStream(
|
||||
id=f"chatcmpl-{uuid.uuid4()}",
|
||||
created=0,
|
||||
model=model,
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
content_chunk.choices = [
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content=parsed["content"], role="assistant"),
|
||||
)
|
||||
]
|
||||
yield content_chunk
|
||||
|
||||
# Stop sentinel chunk (matches SSE path convention)
|
||||
stop_chunk = ModelResponseStream(
|
||||
id=f"chatcmpl-{uuid.uuid4()}",
|
||||
created=0,
|
||||
model=model,
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
stop_chunk.choices = [
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(),
|
||||
)
|
||||
]
|
||||
yield stop_chunk
|
||||
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=_json_as_async_stream(),
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
# SSE stream (text/event-stream or default) - use existing SSE parser
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=self._stream_agentcore_response(response, model),
|
||||
model=model,
|
||||
|
||||
@@ -1206,6 +1206,7 @@ class AmazonConverseConfig(BaseConfig):
|
||||
self._validate_request_metadata(request_metadata)
|
||||
|
||||
output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None)
|
||||
inference_params.pop("output_config", None) # Bedrock Converse doesn't support it
|
||||
|
||||
# keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params'
|
||||
additional_request_params = {
|
||||
|
||||
@@ -94,5 +94,12 @@
|
||||
"assemblyai": {
|
||||
"base_url": "https://llm-gateway.assemblyai.com/v1",
|
||||
"api_key_env": "ASSEMBLYAI_API_KEY"
|
||||
},
|
||||
"charity_engine": {
|
||||
"base_url": "https://api.charityengine.services/remotejobs/v2/inference",
|
||||
"api_key_env": "CHARITY_ENGINE_API_KEY",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
Serper Search API module.
|
||||
"""
|
||||
from litellm.llms.serper.search.transformation import SerperSearchConfig
|
||||
|
||||
__all__ = ["SerperSearchConfig"]
|
||||
@@ -0,0 +1,167 @@
|
||||
"""
|
||||
Calls Serper's /search endpoint to search Google.
|
||||
|
||||
Serper API Reference: https://serper.dev
|
||||
"""
|
||||
from typing import Dict, List, Optional, TypedDict, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.search.transformation import (
|
||||
BaseSearchConfig,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class _SerperSearchRequestRequired(TypedDict):
|
||||
"""Required fields for Serper Search API request."""
|
||||
q: str # Required - search query
|
||||
|
||||
|
||||
class SerperSearchRequest(_SerperSearchRequestRequired, total=False):
|
||||
"""
|
||||
Serper Search API request format.
|
||||
Based on: https://serper.dev
|
||||
"""
|
||||
num: int # Optional - number of results to return, default 10
|
||||
page: int # Optional - page number (default 1)
|
||||
gl: str # Optional - country/geolocation code (e.g., "us", "gb")
|
||||
hl: str # Optional - language code (e.g., "en", "de")
|
||||
location: str # Optional - specific location for search targeting
|
||||
autocorrect: bool # Optional - enable autocorrect (default True)
|
||||
tbs: str # Optional - time-based search filter (e.g., "qdr:h", "qdr:d", "qdr:w")
|
||||
|
||||
|
||||
class SerperSearchConfig(BaseSearchConfig):
|
||||
SERPER_API_BASE = "https://google.serper.dev"
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Serper"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Validate environment and return headers.
|
||||
"""
|
||||
api_key = api_key or get_secret_str("SERPER_API_KEY")
|
||||
if not api_key:
|
||||
raise ValueError("SERPER_API_KEY is not set. Set `SERPER_API_KEY` environment variable.")
|
||||
headers["X-API-KEY"] = api_key
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
optional_params: dict,
|
||||
data: Optional[Union[Dict, List[Dict]]] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
Get complete URL for Search endpoint.
|
||||
"""
|
||||
api_base = api_base or get_secret_str("SERPER_API_BASE") or self.SERPER_API_BASE
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
if not api_base.endswith("/search"):
|
||||
api_base = f"{api_base}/search"
|
||||
|
||||
return api_base
|
||||
|
||||
def transform_search_request(
|
||||
self,
|
||||
query: Union[str, List[str]],
|
||||
optional_params: dict,
|
||||
**kwargs,
|
||||
) -> Dict:
|
||||
"""
|
||||
Transform Search request to Serper API format.
|
||||
|
||||
Args:
|
||||
query: Search query (string or list of strings). Serper only supports single string queries.
|
||||
optional_params: Optional parameters for the request
|
||||
- max_results: Maximum number of search results -> maps to `num`
|
||||
- search_domain_filter: List of domains -> appended as site: clauses to `q`
|
||||
- country: Country code filter (e.g., 'US', 'GB') -> maps to `gl` (lowercased)
|
||||
|
||||
Returns:
|
||||
Dict with typed request data following SerperSearchRequest spec
|
||||
"""
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
|
||||
request_data: SerperSearchRequest = {
|
||||
"q": query,
|
||||
}
|
||||
|
||||
if "max_results" in optional_params:
|
||||
request_data["num"] = optional_params["max_results"]
|
||||
|
||||
if "country" in optional_params:
|
||||
request_data["gl"] = optional_params["country"].lower()
|
||||
|
||||
if "search_domain_filter" in optional_params:
|
||||
domains = optional_params["search_domain_filter"]
|
||||
if isinstance(domains, list) and len(domains) > 0:
|
||||
domain_clauses = " OR ".join(f"site:{d}" for d in domains)
|
||||
request_data["q"] = f"({request_data['q']}) ({domain_clauses})"
|
||||
|
||||
# Convert to dict before dynamic key assignments
|
||||
result_data = dict(request_data)
|
||||
|
||||
# pass through all other parameters as-is
|
||||
for param, value in optional_params.items():
|
||||
if param not in self.get_supported_perplexity_optional_params() and param not in result_data:
|
||||
result_data[param] = value
|
||||
|
||||
return result_data
|
||||
|
||||
def transform_search_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs,
|
||||
) -> SearchResponse:
|
||||
"""
|
||||
Transform Serper API response to LiteLLM unified SearchResponse format.
|
||||
|
||||
Serper -> LiteLLM mappings:
|
||||
- organic[].title -> SearchResult.title
|
||||
- organic[].link -> SearchResult.url
|
||||
- organic[].snippet -> SearchResult.snippet
|
||||
- organic[].date -> SearchResult.date (optional, not always present)
|
||||
|
||||
Args:
|
||||
raw_response: Raw httpx response from Serper API
|
||||
logging_obj: Logging object for tracking
|
||||
|
||||
Returns:
|
||||
SearchResponse with standardized format
|
||||
"""
|
||||
response_json = raw_response.json()
|
||||
|
||||
results = []
|
||||
for result in response_json.get("organic", []):
|
||||
search_result = SearchResult(
|
||||
title=result.get("title", ""),
|
||||
url=result.get("link", ""),
|
||||
snippet=result.get("snippet", ""),
|
||||
date=result.get("date"),
|
||||
last_updated=None,
|
||||
)
|
||||
results.append(search_result)
|
||||
|
||||
return SearchResponse(
|
||||
results=results,
|
||||
object="search",
|
||||
)
|
||||
|
||||
@@ -4207,6 +4207,41 @@
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.3-chat": {
|
||||
"cache_read_input_token_cost": 1.75e-07,
|
||||
"cache_read_input_token_cost_priority": 3.5e-07,
|
||||
"input_cost_per_token": 1.75e-06,
|
||||
"input_cost_per_token_priority": 3.5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 16384,
|
||||
"max_tokens": 16384,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.4e-05,
|
||||
"output_cost_per_token_priority": 2.8e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_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_service_tier": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.3-codex": {
|
||||
"cache_read_input_token_cost": 1.75e-07,
|
||||
"input_cost_per_token": 1.75e-06,
|
||||
@@ -4299,6 +4334,160 @@
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.4": {
|
||||
"cache_read_input_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
|
||||
"cache_read_input_token_cost_priority": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": 1e-06,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 5e-06,
|
||||
"input_cost_per_token_priority": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens_priority": 1e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 2.25e-05,
|
||||
"output_cost_per_token_priority": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 4.5e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_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_service_tier": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.4-2026-03-05": {
|
||||
"cache_read_input_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
|
||||
"cache_read_input_token_cost_priority": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": 1e-06,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 5e-06,
|
||||
"input_cost_per_token_priority": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens_priority": 1e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 2.25e-05,
|
||||
"output_cost_per_token_priority": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 4.5e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_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_service_tier": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.4-pro": {
|
||||
"cache_read_input_token_cost": 3e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 6e-06,
|
||||
"input_cost_per_token": 3e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 6e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.4-pro-2026-03-05": {
|
||||
"cache_read_input_token_cost": 3e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 6e-06,
|
||||
"input_cost_per_token": 3e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 6e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-image-1": {
|
||||
"cache_read_input_image_token_cost": 2.5e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
@@ -12090,6 +12279,14 @@
|
||||
"notes": "SearXNG is an open-source metasearch engine. Free to use when self-hosted or using public instances."
|
||||
}
|
||||
},
|
||||
"serper/search": {
|
||||
"input_cost_per_query": 0.001,
|
||||
"litellm_provider": "serper",
|
||||
"mode": "search",
|
||||
"metadata": {
|
||||
"notes": "Serper Google Search API. Pricing: $1.00/1k queries (Starter), $0.75/1k (Standard), $0.50/1k (Scale), $0.30/1k (Ultimate)."
|
||||
}
|
||||
},
|
||||
"elevenlabs/scribe_v1": {
|
||||
"input_cost_per_second": 6.11e-05,
|
||||
"litellm_provider": "elevenlabs",
|
||||
@@ -16799,6 +16996,42 @@
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gemini/gemini-3.1-flash-image-preview": {
|
||||
"input_cost_per_token": 2.5e-07,
|
||||
"input_cost_per_token_batches": 1.25e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 65536,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.045,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_image_token_batches": 3e-05,
|
||||
"output_cost_per_token": 1.5e-06,
|
||||
"output_cost_per_token_batches": 7.5e-07,
|
||||
"rpm": 1000,
|
||||
"tpm": 4000000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image-preview",
|
||||
"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/deep-research-pro-preview-12-2025": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"input_cost_per_token": 2e-06,
|
||||
@@ -21083,7 +21316,7 @@
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"output_cost_per_token_flex": 9e-05,
|
||||
@@ -21091,9 +21324,8 @@
|
||||
"output_cost_per_token_priority": 0.00027,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 0.000405,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
"/v1/responses",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
@@ -21132,7 +21364,7 @@
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"output_cost_per_token_flex": 9e-05,
|
||||
@@ -21140,9 +21372,8 @@
|
||||
"output_cost_per_token_priority": 0.00027,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 0.000405,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
"/v1/responses",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
||||
@@ -2061,6 +2061,13 @@
|
||||
"search": true
|
||||
}
|
||||
},
|
||||
"serper": {
|
||||
"display_name": "Serper (`serper`)",
|
||||
"url": "https://docs.litellm.ai/docs/search/serper",
|
||||
"endpoints": {
|
||||
"search": true
|
||||
}
|
||||
},
|
||||
"triton": {
|
||||
"display_name": "Triton (`triton`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/triton-inference-server",
|
||||
|
||||
@@ -92,7 +92,24 @@ def get_base_url(spec: Dict[str, Any], spec_path: Optional[str] = None) -> str:
|
||||
"""Extract base URL from OpenAPI spec."""
|
||||
# OpenAPI 3.x
|
||||
if "servers" in spec and spec["servers"]:
|
||||
return spec["servers"][0]["url"]
|
||||
server_url = spec["servers"][0]["url"]
|
||||
|
||||
# If the server URL is relative (starts with /), derive base from spec_path
|
||||
if server_url.startswith("/") and spec_path:
|
||||
if spec_path.startswith("http://") or spec_path.startswith("https://"):
|
||||
# Extract base URL from spec_path (e.g., https://petstore3.swagger.io/api/v3/openapi.json)
|
||||
# Combine domain with the relative server URL
|
||||
from urllib.parse import urlparse
|
||||
parsed = urlparse(spec_path)
|
||||
base_domain = f"{parsed.scheme}://{parsed.netloc}"
|
||||
full_base_url = base_domain + server_url
|
||||
verbose_logger.info(
|
||||
f"OpenAPI spec has relative server URL '{server_url}'. "
|
||||
f"Deriving base from spec_path: {full_base_url}"
|
||||
)
|
||||
return full_base_url
|
||||
|
||||
return server_url
|
||||
# OpenAPI 2.x (Swagger)
|
||||
elif "host" in spec:
|
||||
scheme = spec.get("schemes", ["https"])[0]
|
||||
|
||||
@@ -711,6 +711,7 @@ if MCP_AVAILABLE:
|
||||
|
||||
Checks both the full tool name and unprefixed version (without server prefix).
|
||||
This allows users to configure simple tool names regardless of prefixing.
|
||||
Comparison is case-insensitive to handle OpenAPI operationIds that may be in camelCase.
|
||||
|
||||
Args:
|
||||
tool_name: The tool name to check (may be prefixed like "server-tool_name")
|
||||
@@ -723,13 +724,15 @@ if MCP_AVAILABLE:
|
||||
split_server_prefix_from_name,
|
||||
)
|
||||
|
||||
# Check if the full name is in the list
|
||||
if tool_name in filter_list:
|
||||
# Normalize filter list to lowercase for case-insensitive comparison
|
||||
filter_list_lower = [f.lower() for f in filter_list]
|
||||
|
||||
if tool_name.lower() in filter_list_lower:
|
||||
return True
|
||||
|
||||
# Check if the unprefixed name is in the list
|
||||
# Check if the unprefixed name is in the list (case-insensitive)
|
||||
unprefixed_name, _ = split_server_prefix_from_name(tool_name)
|
||||
return unprefixed_name in filter_list
|
||||
return unprefixed_name.lower() in filter_list_lower
|
||||
|
||||
def filter_tools_by_allowed_tools(
|
||||
tools: List[MCPTool],
|
||||
|
||||
@@ -108,16 +108,23 @@ def get_key_models(
|
||||
"""
|
||||
all_models: List[str] = []
|
||||
if len(user_api_key_dict.models) > 0:
|
||||
all_models = user_api_key_dict.models
|
||||
all_models = list(user_api_key_dict.models) # copy to avoid mutating cached objects
|
||||
if SpecialModelNames.all_team_models.value in all_models:
|
||||
all_models = user_api_key_dict.team_models
|
||||
all_models = list(user_api_key_dict.team_models) # copy to avoid mutating cached objects
|
||||
if SpecialModelNames.all_proxy_models.value in all_models:
|
||||
all_models = proxy_model_list
|
||||
all_models = list(proxy_model_list) # copy to avoid mutating caller's list
|
||||
if include_model_access_groups:
|
||||
all_models.extend(model_access_groups.keys())
|
||||
|
||||
all_models = _get_models_from_access_groups(
|
||||
model_access_groups=model_access_groups, all_models=all_models
|
||||
model_access_groups=model_access_groups,
|
||||
all_models=all_models,
|
||||
include_model_access_groups=include_model_access_groups,
|
||||
)
|
||||
|
||||
# deduplicate while preserving order
|
||||
all_models = list(dict.fromkeys(all_models))
|
||||
|
||||
verbose_proxy_logger.debug("ALL KEY MODELS - {}".format(len(all_models)))
|
||||
return all_models
|
||||
|
||||
@@ -141,8 +148,8 @@ def get_team_models(
|
||||
all_models_set.update(team_models)
|
||||
if SpecialModelNames.all_proxy_models.value in all_models_set:
|
||||
all_models_set.update(proxy_model_list)
|
||||
|
||||
all_models = list(all_models_set)
|
||||
if include_model_access_groups:
|
||||
all_models_set.update(model_access_groups.keys())
|
||||
|
||||
all_models = _get_models_from_access_groups(
|
||||
model_access_groups=model_access_groups,
|
||||
@@ -150,6 +157,9 @@ def get_team_models(
|
||||
include_model_access_groups=include_model_access_groups,
|
||||
)
|
||||
|
||||
# deduplicate while preserving order
|
||||
all_models = list(dict.fromkeys(all_models))
|
||||
|
||||
verbose_proxy_logger.debug("ALL TEAM MODELS - {}".format(len(all_models)))
|
||||
return all_models
|
||||
|
||||
|
||||
@@ -615,17 +615,23 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
# This allows UI SSO to work separately from API M2M authentication
|
||||
# Note: Info routes are already scoped to the user
|
||||
if RouteChecks.is_llm_api_route(route=route) or RouteChecks.is_info_route(route=route):
|
||||
# return UserAPIKeyAuth object
|
||||
# helper to check if the api_key is a valid oauth2 token
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
# When both OAuth2 and JWT auth are enabled, use token format to decide:
|
||||
# - JWT tokens (3 dot-separated parts) -> skip OAuth2, fall through to JWT handler
|
||||
# - Opaque tokens -> use OAuth2 handler
|
||||
# This allows JWT for users and OAuth2 for M2M on the same instance
|
||||
is_jwt_token = jwt_handler.is_jwt(token=api_key) if general_settings.get("enable_jwt_auth", False) is True else False
|
||||
if not is_jwt_token:
|
||||
# return UserAPIKeyAuth object
|
||||
# helper to check if the api_key is a valid oauth2 token
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
||||
if premium_user is not True:
|
||||
raise ValueError(
|
||||
"Oauth2 token validation is only available for premium users"
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
)
|
||||
if premium_user is not True:
|
||||
raise ValueError(
|
||||
"Oauth2 token validation is only available for premium users"
|
||||
+ CommonProxyErrors.not_premium_user.value
|
||||
)
|
||||
|
||||
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
||||
return await Oauth2Handler.check_oauth2_token(token=api_key)
|
||||
|
||||
if general_settings.get("enable_oauth2_proxy_auth", False) is True:
|
||||
return await handle_oauth2_proxy_request(request=request)
|
||||
|
||||
@@ -27,6 +27,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
encode_file_id_with_model,
|
||||
get_batch_from_database,
|
||||
get_credentials_for_model,
|
||||
get_model_id_from_unified_batch_id,
|
||||
get_models_from_unified_file_id,
|
||||
get_original_file_id,
|
||||
prepare_data_with_credentials,
|
||||
@@ -487,6 +488,10 @@ async def retrieve_batch( # noqa: PLR0915
|
||||
|
||||
response = await llm_router.aretrieve_batch(**data) # type: ignore
|
||||
response._hidden_params["unified_batch_id"] = unified_batch_id
|
||||
if unified_batch_id:
|
||||
model_id_from_batch = get_model_id_from_unified_batch_id(unified_batch_id)
|
||||
if model_id_from_batch:
|
||||
response._hidden_params["model_id"] = model_id_from_batch
|
||||
|
||||
# SCENARIO 3: Fallback to custom_llm_provider (uses env variables)
|
||||
else:
|
||||
|
||||
@@ -1844,7 +1844,13 @@ async def delete_user(
|
||||
|
||||
## DELETE ASSOCIATED INVITATION LINKS
|
||||
await prisma_client.db.litellm_invitationlink.delete_many(
|
||||
where={"user_id": {"in": data.user_ids}}
|
||||
where={
|
||||
"OR": [
|
||||
{"user_id": {"in": data.user_ids}},
|
||||
{"created_by": {"in": data.user_ids}},
|
||||
{"updated_by": {"in": data.user_ids}},
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
## DELETE ASSOCIATED ORGANIZATION MEMBERSHIPS
|
||||
|
||||
@@ -2827,21 +2827,6 @@ async def validate_membership(
|
||||
)
|
||||
|
||||
|
||||
def _unfurl_all_proxy_models(
|
||||
team_info: LiteLLM_TeamTable, llm_router: Router
|
||||
) -> LiteLLM_TeamTable:
|
||||
if (
|
||||
SpecialModelNames.all_proxy_models.value in team_info.models
|
||||
and llm_router is not None
|
||||
):
|
||||
team_models: set[str] = set() # make set to avoid duplicates
|
||||
for model in team_info.models:
|
||||
if model != SpecialModelNames.all_proxy_models.value:
|
||||
team_models.add(model)
|
||||
for model in llm_router.get_model_names():
|
||||
team_models.add(model)
|
||||
team_info.models = list(team_models)
|
||||
return team_info
|
||||
|
||||
|
||||
async def _add_team_member_budget_table(
|
||||
@@ -2972,9 +2957,6 @@ async def team_info(
|
||||
team_info_response_object=_team_info,
|
||||
)
|
||||
|
||||
# ## UNFURL 'all-proxy-models' into the team_info.models list ##
|
||||
# if llm_router is not None:
|
||||
# _team_info = _unfurl_all_proxy_models(_team_info, llm_router)
|
||||
response_object = TeamInfoResponseObject(
|
||||
team_id=team_id,
|
||||
team_info=_team_info,
|
||||
|
||||
@@ -377,9 +377,6 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import user_upda
|
||||
from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import (
|
||||
router as jwt_key_mapping_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import (
|
||||
router as jwt_key_mapping_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
delete_verification_tokens,
|
||||
duration_in_seconds,
|
||||
|
||||
+19
-112
@@ -8,7 +8,7 @@ from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Literal, Optional, Tuple, Union, get_type_hints
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm._uuid import uuid
|
||||
@@ -16,7 +16,6 @@ from litellm._uuid import uuid
|
||||
from .completion import CompletionRequest
|
||||
from .embedding import EmbeddingRequest
|
||||
from .llms.openai import OpenAIFileObject
|
||||
from .llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
|
||||
from .search import SearchProvider
|
||||
from .utils import CustomPricingLiteLLMParams, ModelResponse
|
||||
|
||||
@@ -162,6 +161,9 @@ class CredentialLiteLLMParams(BaseModel):
|
||||
watsonx_region_name: Optional[str] = None
|
||||
|
||||
|
||||
_RESERVED_INIT_KEYS = frozenset({"self", "params", "__class__"})
|
||||
|
||||
|
||||
class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
||||
"""
|
||||
LiteLLM Params without 'model' arg (used across completion / assistants api)
|
||||
@@ -215,76 +217,21 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
||||
vector_store_id: Optional[str] = None
|
||||
milvus_text_field: Optional[str] = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
max_retries: Optional[Union[int, str]] = None,
|
||||
tpm: Optional[int] = None,
|
||||
rpm: Optional[int] = None,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
timeout: Optional[Union[float, str]] = None, # if str, pass in as os.environ/
|
||||
stream_timeout: Optional[Union[float, str]] = (
|
||||
None # timeout when making stream=True calls, if str, pass in as os.environ/
|
||||
),
|
||||
organization: Optional[str] = None, # for openai orgs
|
||||
## LOGGING PARAMS ##
|
||||
litellm_trace_id: Optional[str] = None,
|
||||
## UNIFIED PROJECT/REGION ##
|
||||
region_name: Optional[str] = None,
|
||||
## VERTEX AI ##
|
||||
vertex_project: Optional[str] = None,
|
||||
vertex_location: Optional[str] = None,
|
||||
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
|
||||
## AWS BEDROCK / SAGEMAKER ##
|
||||
aws_access_key_id: Optional[str] = None,
|
||||
aws_secret_access_key: Optional[str] = None,
|
||||
aws_region_name: Optional[str] = None,
|
||||
## IBM WATSONX ##
|
||||
watsonx_region_name: Optional[str] = None,
|
||||
input_cost_per_token: Optional[float] = None,
|
||||
output_cost_per_token: Optional[float] = None,
|
||||
input_cost_per_second: Optional[float] = None,
|
||||
output_cost_per_second: Optional[float] = None,
|
||||
max_file_size_mb: Optional[float] = None,
|
||||
# Deployment budgets
|
||||
max_budget: Optional[float] = None,
|
||||
budget_duration: Optional[str] = None,
|
||||
# Pass through params
|
||||
use_in_pass_through: Optional[bool] = False,
|
||||
# Dynamic param to force using litellm proxy
|
||||
use_litellm_proxy: Optional[bool] = False,
|
||||
# This will merge the reasoning content in the choices
|
||||
merge_reasoning_content_in_choices: Optional[bool] = False,
|
||||
model_info: Optional[Dict] = None,
|
||||
mock_response: Optional[Union[str, ModelResponse, Exception, Any]] = None,
|
||||
# auto-router params
|
||||
auto_router_config_path: Optional[str] = None,
|
||||
auto_router_config: Optional[str] = None,
|
||||
auto_router_default_model: Optional[str] = None,
|
||||
auto_router_embedding_model: Optional[str] = None,
|
||||
# complexity-router params
|
||||
complexity_router_config: Optional[Dict] = None,
|
||||
complexity_router_default_model: Optional[str] = None,
|
||||
# Batch/File API Params
|
||||
s3_bucket_name: Optional[str] = None,
|
||||
s3_encryption_key_id: Optional[str] = None,
|
||||
gcs_bucket_name: Optional[str] = None,
|
||||
**params,
|
||||
):
|
||||
args = locals()
|
||||
args.pop("max_retries", None)
|
||||
args.pop("self", None)
|
||||
args.pop("params", None)
|
||||
args.pop("__class__", None)
|
||||
if max_retries is not None and isinstance(max_retries, str):
|
||||
max_retries = int(max_retries) # cast to int
|
||||
# We need to keep max_retries in args since it's a parameter of GenericLiteLLMParams
|
||||
args[
|
||||
"max_retries"
|
||||
] = max_retries # Put max_retries back in args after popping it
|
||||
super().__init__(**args, **params)
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def preprocess_input_data(cls, data: Any) -> Any:
|
||||
"""
|
||||
Pre-process input data before validation:
|
||||
1. Filter out reserved Python keywords ('self', 'params', '__class__') to prevent
|
||||
'got multiple values for argument' errors when user data contains these keys.
|
||||
2. Convert max_retries from string to int if needed.
|
||||
"""
|
||||
if isinstance(data, dict):
|
||||
filtered = {k: v for k, v in data.items() if k not in _RESERVED_INIT_KEYS}
|
||||
if "max_retries" in filtered and isinstance(filtered["max_retries"], str):
|
||||
filtered["max_retries"] = int(filtered["max_retries"])
|
||||
return filtered
|
||||
return data
|
||||
|
||||
def __contains__(self, key):
|
||||
# Define custom behavior for the 'in' operator
|
||||
@@ -311,46 +258,6 @@ class LiteLLM_Params(GenericLiteLLMParams):
|
||||
model: str
|
||||
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
max_retries: Optional[Union[int, str]] = None,
|
||||
tpm: Optional[int] = None,
|
||||
rpm: Optional[int] = None,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
timeout: Optional[Union[float, str]] = None, # if str, pass in as os.environ/
|
||||
stream_timeout: Optional[Union[float, str]] = (
|
||||
None # timeout when making stream=True calls, if str, pass in as os.environ/
|
||||
),
|
||||
organization: Optional[str] = None, # for openai orgs
|
||||
## VERTEX AI ##
|
||||
vertex_project: Optional[str] = None,
|
||||
vertex_location: Optional[str] = None,
|
||||
## AWS BEDROCK / SAGEMAKER ##
|
||||
aws_access_key_id: Optional[str] = None,
|
||||
aws_secret_access_key: Optional[str] = None,
|
||||
aws_region_name: Optional[str] = None,
|
||||
# OpenAI / Azure Whisper
|
||||
# set a max-size of file that can be passed to litellm proxy
|
||||
max_file_size_mb: Optional[float] = None,
|
||||
# will use deployment on pass-through endpoints if True
|
||||
use_in_pass_through: Optional[bool] = False,
|
||||
use_litellm_proxy: Optional[bool] = False,
|
||||
**params,
|
||||
):
|
||||
args = locals()
|
||||
args.pop("max_retries", None)
|
||||
args.pop("self", None)
|
||||
args.pop("params", None)
|
||||
args.pop("__class__", None)
|
||||
if max_retries is not None and isinstance(max_retries, str):
|
||||
max_retries = int(max_retries) # cast to int
|
||||
args["max_retries"] = max_retries
|
||||
super().__init__(**{**args, **params})
|
||||
|
||||
def __contains__(self, key):
|
||||
# Define custom behavior for the 'in' operator
|
||||
return hasattr(self, key)
|
||||
|
||||
@@ -3177,6 +3177,7 @@ class LlmProviders(str, Enum):
|
||||
TOPAZ = "topaz"
|
||||
SAP_GENERATIVE_AI_HUB = "sap"
|
||||
ASSEMBLYAI = "assemblyai"
|
||||
CHARITY_ENGINE = "charity_engine"
|
||||
GITHUB_COPILOT = "github_copilot"
|
||||
SNOWFLAKE = "snowflake"
|
||||
GRADIENT_AI = "gradient_ai"
|
||||
@@ -3249,6 +3250,7 @@ class SearchProviders(str, Enum):
|
||||
LINKUP = "linkup"
|
||||
DUCKDUCKGO = "duckduckgo"
|
||||
SEARCHAPI = "searchapi"
|
||||
SERPER = "serper"
|
||||
|
||||
|
||||
# Create a set of all search provider values for quick lookup
|
||||
|
||||
@@ -8884,6 +8884,7 @@ class ProviderConfigManager:
|
||||
from litellm.llms.perplexity.search.transformation import PerplexitySearchConfig
|
||||
from litellm.llms.searchapi.search.transformation import SearchAPIConfig
|
||||
from litellm.llms.searxng.search.transformation import SearXNGSearchConfig
|
||||
from litellm.llms.serper.search.transformation import SerperSearchConfig
|
||||
from litellm.llms.tavily.search.transformation import TavilySearchConfig
|
||||
|
||||
PROVIDER_TO_CONFIG_MAP = {
|
||||
@@ -8899,6 +8900,7 @@ class ProviderConfigManager:
|
||||
SearchProviders.LINKUP: LinkupSearchConfig,
|
||||
SearchProviders.DUCKDUCKGO: DuckDuckGoSearchConfig,
|
||||
SearchProviders.SEARCHAPI: SearchAPIConfig,
|
||||
SearchProviders.SERPER: SerperSearchConfig,
|
||||
}
|
||||
config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None)
|
||||
if config_class is None:
|
||||
|
||||
@@ -4207,6 +4207,41 @@
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.3-chat": {
|
||||
"cache_read_input_token_cost": 1.75e-07,
|
||||
"cache_read_input_token_cost_priority": 3.5e-07,
|
||||
"input_cost_per_token": 1.75e-06,
|
||||
"input_cost_per_token_priority": 3.5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 16384,
|
||||
"max_tokens": 16384,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.4e-05,
|
||||
"output_cost_per_token_priority": 2.8e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_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_service_tier": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.3-codex": {
|
||||
"cache_read_input_token_cost": 1.75e-07,
|
||||
"input_cost_per_token": 1.75e-06,
|
||||
@@ -4299,6 +4334,160 @@
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.4": {
|
||||
"cache_read_input_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
|
||||
"cache_read_input_token_cost_priority": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": 1e-06,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 5e-06,
|
||||
"input_cost_per_token_priority": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens_priority": 1e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 2.25e-05,
|
||||
"output_cost_per_token_priority": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 4.5e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_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_service_tier": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.4-2026-03-05": {
|
||||
"cache_read_input_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
|
||||
"cache_read_input_token_cost_priority": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": 1e-06,
|
||||
"input_cost_per_token": 2.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 5e-06,
|
||||
"input_cost_per_token_priority": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens_priority": 1e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 2.25e-05,
|
||||
"output_cost_per_token_priority": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 4.5e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_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_service_tier": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.4-pro": {
|
||||
"cache_read_input_token_cost": 3e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 6e-06,
|
||||
"input_cost_per_token": 3e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 6e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.4-pro-2026-03-05": {
|
||||
"cache_read_input_token_cost": 3e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 6e-06,
|
||||
"input_cost_per_token": 3e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 6e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"supported_endpoints": [
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-image-1": {
|
||||
"cache_read_input_image_token_cost": 2.5e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
@@ -12090,6 +12279,14 @@
|
||||
"notes": "SearXNG is an open-source metasearch engine. Free to use when self-hosted or using public instances."
|
||||
}
|
||||
},
|
||||
"serper/search": {
|
||||
"input_cost_per_query": 0.001,
|
||||
"litellm_provider": "serper",
|
||||
"mode": "search",
|
||||
"metadata": {
|
||||
"notes": "Serper Google Search API. Pricing: $1.00/1k queries (Starter), $0.75/1k (Standard), $0.50/1k (Scale), $0.30/1k (Ultimate)."
|
||||
}
|
||||
},
|
||||
"elevenlabs/scribe_v1": {
|
||||
"input_cost_per_second": 6.11e-05,
|
||||
"litellm_provider": "elevenlabs",
|
||||
@@ -16799,6 +16996,42 @@
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"gemini/gemini-3.1-flash-image-preview": {
|
||||
"input_cost_per_token": 2.5e-07,
|
||||
"input_cost_per_token_batches": 1.25e-07,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 65536,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.045,
|
||||
"output_cost_per_image_token": 6e-05,
|
||||
"output_cost_per_image_token_batches": 3e-05,
|
||||
"output_cost_per_token": 1.5e-06,
|
||||
"output_cost_per_token_batches": 7.5e-07,
|
||||
"rpm": 1000,
|
||||
"tpm": 4000000,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image-preview",
|
||||
"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/deep-research-pro-preview-12-2025": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"input_cost_per_token": 2e-06,
|
||||
@@ -21083,7 +21316,7 @@
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"output_cost_per_token_flex": 9e-05,
|
||||
@@ -21091,9 +21324,8 @@
|
||||
"output_cost_per_token_priority": 0.00027,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 0.000405,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
"/v1/responses",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
@@ -21132,7 +21364,7 @@
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 0.00018,
|
||||
"output_cost_per_token_above_272k_tokens": 0.00027,
|
||||
"output_cost_per_token_flex": 9e-05,
|
||||
@@ -21140,9 +21372,8 @@
|
||||
"output_cost_per_token_priority": 0.00027,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 0.000405,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
"/v1/responses",
|
||||
"/v1/batch"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
||||
Generated
+4
-4
@@ -3222,15 +3222,15 @@ files = [
|
||||
|
||||
[[package]]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.52"
|
||||
version = "0.4.53"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
optional = true
|
||||
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"proxy\""
|
||||
files = [
|
||||
{file = "litellm_proxy_extras-0.4.52-py3-none-any.whl", hash = "sha256:5cdfeb5b93f6e4329299b3eabdb1e51beb264b075e1b5179149d8ded084b4aaa"},
|
||||
{file = "litellm_proxy_extras-0.4.52.tar.gz", hash = "sha256:fcac06b212ef12bb0f79fe465680f2f0e85e4aaab9234780fd3dc18e3598e743"},
|
||||
{file = "litellm_proxy_extras-0.4.53-py3-none-any.whl", hash = "sha256:9224c667144774b6119e4de9b4b2d52fafc58442e6db317785c43b2d833665d6"},
|
||||
{file = "litellm_proxy_extras-0.4.53.tar.gz", hash = "sha256:22c53fa8890d93d4a0d24171726e4e2bba8be6fef4838317cb74284fa9d27f70"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -8002,4 +8002,4 @@ utils = ["numpydoc"]
|
||||
[metadata]
|
||||
lock-version = "2.1"
|
||||
python-versions = ">=3.9,<4.0"
|
||||
content-hash = "fa110a048c30d0ad4e66414290ec103dba7707d99474827ea0cf3e4a2058d165"
|
||||
content-hash = "3036cfcdc06fb4293e248a2edd9c32a7afe6846920167527e247b2aefd74cfa6"
|
||||
|
||||
@@ -458,6 +458,24 @@
|
||||
"interactions": true
|
||||
}
|
||||
},
|
||||
"charity_engine": {
|
||||
"display_name": "Charity Engine (`charity_engine`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/charity_engine",
|
||||
"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,
|
||||
"a2a": false,
|
||||
"interactions": false
|
||||
}
|
||||
},
|
||||
"chutes": {
|
||||
"display_name": "Chutes (`chutes`)",
|
||||
"endpoints": {
|
||||
@@ -2061,6 +2079,13 @@
|
||||
"search": true
|
||||
}
|
||||
},
|
||||
"serper": {
|
||||
"display_name": "Serper (`serper`)",
|
||||
"url": "https://docs.litellm.ai/docs/search/serper",
|
||||
"endpoints": {
|
||||
"search": true
|
||||
}
|
||||
},
|
||||
"triton": {
|
||||
"display_name": "Triton (`triton`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/triton-inference-server",
|
||||
|
||||
+1
-1
@@ -61,7 +61,7 @@ boto3 = { version = "^1.40.76", optional = true }
|
||||
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
|
||||
mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
|
||||
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
|
||||
litellm-proxy-extras = {version = "^0.4.52", optional = true}
|
||||
litellm-proxy-extras = {version = "^0.4.53", optional = true}
|
||||
rich = {version = "^13.7.1", optional = true}
|
||||
litellm-enterprise = {version = "^0.1.33", optional = true}
|
||||
diskcache = {version = "^5.6.1", optional = true}
|
||||
|
||||
+2
-2
@@ -57,7 +57,7 @@ grpcio>=1.75.0; python_version >= "3.14"
|
||||
sentry_sdk==2.21.0 # for sentry error handling
|
||||
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
|
||||
tzdata==2025.1 # IANA time zone database
|
||||
litellm-proxy-extras==0.4.52 # for proxy extras - e.g. prisma migrations
|
||||
litellm-proxy-extras==0.4.53 # for proxy extras - e.g. prisma migrations
|
||||
llm-sandbox==0.3.31 # for skill execution in sandbox
|
||||
### LITELLM PACKAGE DEPENDENCIES
|
||||
python-dotenv==1.0.1 # for env
|
||||
@@ -80,4 +80,4 @@ pypdf>=6.7.3 # for PDF text extraction in RAG ingestion (CVE-2026-27888)
|
||||
########################
|
||||
# LITELLM ENTERPRISE DEPENDENCIES
|
||||
########################
|
||||
litellm-enterprise==0.1.33
|
||||
litellm-enterprise==0.1.34
|
||||
|
||||
@@ -18,6 +18,7 @@ SEARCH_PROVIDERS = [
|
||||
"linkup",
|
||||
"duckduckgo",
|
||||
"searchapi",
|
||||
"serper",
|
||||
]
|
||||
|
||||
ALLOWED_FILES_IN_LLMS_FOLDER = [
|
||||
|
||||
@@ -42,44 +42,52 @@ def image_url():
|
||||
image_file = BytesIO()
|
||||
img.save(image_file, format="PNG")
|
||||
image_file.seek(0)
|
||||
# openai>=2.24.0 requires BytesIO to have .name for MIME type detection in multipart uploads
|
||||
image_file.name = "litellm_logo.png"
|
||||
|
||||
return image_file
|
||||
|
||||
|
||||
def test_openai_image_variation_openai_sdk(image_url):
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI()
|
||||
response = client.images.create_variation(image=image_url, n=2, size="1024x1024")
|
||||
print(response)
|
||||
# Commented out: OpenAI /images/variations endpoint deprecated (DALL-E 2 shutdown May 12, 2026)
|
||||
# def test_openai_image_variation_openai_sdk(image_url):
|
||||
# from openai import OpenAI
|
||||
#
|
||||
# client = OpenAI()
|
||||
# response = client.images.create_variation(image=image_url, n=2, size="1024x1024")
|
||||
# print(response)
|
||||
#
|
||||
#
|
||||
# @pytest.mark.parametrize("sync_mode", [True, False])
|
||||
# @pytest.mark.asyncio
|
||||
# async def test_openai_image_variation_litellm_sdk(image_url, sync_mode):
|
||||
# from litellm import image_variation, aimage_variation
|
||||
#
|
||||
# if sync_mode:
|
||||
# image_variation(image=image_url, n=2, size="1024x1024")
|
||||
# else:
|
||||
# await aimage_variation(image=image_url, n=2, size="1024x1024")
|
||||
#
|
||||
#
|
||||
# def test_topaz_image_variation(image_url):
|
||||
# from litellm import image_variation, aimage_variation
|
||||
# from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
# from unittest.mock import patch
|
||||
#
|
||||
# client = HTTPHandler()
|
||||
# with patch.object(client, "post") as mock_post:
|
||||
# try:
|
||||
# image_variation(
|
||||
# model="topaz/Standard V2",
|
||||
# image=image_url,
|
||||
# n=2,
|
||||
# size="1024x1024",
|
||||
# client=client,
|
||||
# )
|
||||
# except Exception as e:
|
||||
# print(e)
|
||||
# mock_post.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_image_variation_litellm_sdk(image_url, sync_mode):
|
||||
from litellm import image_variation, aimage_variation
|
||||
|
||||
if sync_mode:
|
||||
image_variation(image=image_url, n=2, size="1024x1024")
|
||||
else:
|
||||
await aimage_variation(image=image_url, n=2, size="1024x1024")
|
||||
|
||||
|
||||
def test_topaz_image_variation(image_url):
|
||||
from litellm import image_variation, aimage_variation
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from unittest.mock import patch
|
||||
|
||||
client = HTTPHandler()
|
||||
with patch.object(client, "post") as mock_post:
|
||||
try:
|
||||
image_variation(
|
||||
model="topaz/Standard V2",
|
||||
image=image_url,
|
||||
n=2,
|
||||
size="1024x1024",
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
mock_post.assert_called_once()
|
||||
def test_image_variation_placeholder():
|
||||
"""Placeholder: variation tests commented out - OpenAI /images/variations deprecated (DALL-E 2 shutdown May 12, 2026)."""
|
||||
pass
|
||||
|
||||
@@ -868,8 +868,9 @@ class BaseLLMChatTest(ABC):
|
||||
base_completion_call_args = self.get_base_completion_call_args()
|
||||
if not supports_vision(base_completion_call_args["model"], None):
|
||||
pytest.skip("Model does not support image input")
|
||||
elif "http://" in image_url and "fireworks_ai" in base_completion_call_args.get(
|
||||
"model"
|
||||
elif "http://" in image_url and (
|
||||
"fireworks_ai" in base_completion_call_args.get("model", "")
|
||||
or "mistral" in base_completion_call_args.get("model", "")
|
||||
):
|
||||
pytest.skip("Model does not support http:// input")
|
||||
|
||||
|
||||
@@ -23,27 +23,48 @@ from litellm.types.llms.anthropic_skills import (
|
||||
|
||||
|
||||
@contextmanager
|
||||
def create_skill_zip(skill_name: str):
|
||||
def create_skill_zip(skill_name: str, unique_suffix: Optional[str] = None):
|
||||
"""
|
||||
Helper context manager to create a zip file for a skill.
|
||||
|
||||
|
||||
Args:
|
||||
skill_name: Name of the skill directory in test_skills_data/
|
||||
|
||||
unique_suffix: Optional suffix to make the skill name unique in the zip.
|
||||
When provided, the SKILL.md frontmatter name is rewritten
|
||||
to avoid duplicate-name conflicts on the API side.
|
||||
|
||||
Yields:
|
||||
File handle to the zip file
|
||||
|
||||
|
||||
The zip file is automatically cleaned up after use.
|
||||
"""
|
||||
import time
|
||||
|
||||
test_dir = Path(__file__).parent / "test_skills_data"
|
||||
skill_dir = test_dir / skill_name
|
||||
|
||||
|
||||
# Create a zip file containing the skill directory
|
||||
# When unique_suffix is set, folder name must match skill name in SKILL.md (Anthropic requirement)
|
||||
zip_folder_name = f"{skill_name}-{unique_suffix}" if unique_suffix else skill_name
|
||||
zip_path = test_dir / f"{skill_name}.zip"
|
||||
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zip_file:
|
||||
zip_file.write(skill_dir, arcname=skill_name)
|
||||
zip_file.write(skill_dir / "SKILL.md", arcname=f"{skill_name}/SKILL.md")
|
||||
|
||||
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zf:
|
||||
if unique_suffix is not None:
|
||||
# Rewrite SKILL.md with a unique name and use matching folder name
|
||||
skill_md = (skill_dir / "SKILL.md").read_text()
|
||||
skill_md = skill_md.replace(
|
||||
f"name: {skill_name}",
|
||||
f"name: {zip_folder_name}",
|
||||
)
|
||||
zf.writestr(f"{zip_folder_name}/SKILL.md", skill_md)
|
||||
# Add any other files in the skill dir (e.g. subdirs) under the new folder name
|
||||
for f in skill_dir.rglob("*"):
|
||||
if f.is_file() and f.name != "SKILL.md":
|
||||
rel = f.relative_to(skill_dir)
|
||||
zf.write(f, arcname=f"{zip_folder_name}/{rel}")
|
||||
else:
|
||||
zf.write(skill_dir, arcname=skill_name)
|
||||
zf.write(skill_dir / "SKILL.md", arcname=f"{skill_name}/SKILL.md")
|
||||
|
||||
try:
|
||||
with open(zip_path, "rb") as f:
|
||||
yield f
|
||||
@@ -77,13 +98,13 @@ class BaseSkillsAPITest(ABC):
|
||||
def test_create_skill(self):
|
||||
"""
|
||||
Test creating a skill.
|
||||
|
||||
|
||||
Note: This test creates a skill but does not clean it up,
|
||||
as we want to verify it was created successfully.
|
||||
The test_delete_skill test will handle cleanup.
|
||||
"""
|
||||
import time
|
||||
|
||||
|
||||
custom_llm_provider = self.get_custom_llm_provider()
|
||||
api_key = self.get_api_key()
|
||||
api_base = self.get_api_base()
|
||||
@@ -96,12 +117,14 @@ class BaseSkillsAPITest(ABC):
|
||||
|
||||
# Use helper to create skill zip
|
||||
skill_name = "test-skill-litellm"
|
||||
|
||||
# Use unique title to avoid conflicts with previous test runs
|
||||
unique_title = f"Test Skill {int(time.time())}"
|
||||
|
||||
|
||||
# Use unique title and unique skill name to avoid conflicts
|
||||
# with previous test runs (skills are never cleaned up in CI)
|
||||
ts = str(int(time.time()))
|
||||
unique_title = f"Test Skill {ts}"
|
||||
|
||||
# Upload the skill with the zip file
|
||||
with create_skill_zip(skill_name) as zip_file:
|
||||
with create_skill_zip(skill_name, unique_suffix=ts) as zip_file:
|
||||
response = litellm.create_skill(
|
||||
display_title=unique_title,
|
||||
files=[zip_file],
|
||||
@@ -217,12 +240,13 @@ class BaseSkillsAPITest(ABC):
|
||||
|
||||
# Use helper to create skill zip
|
||||
skill_name = "test-delete-skill"
|
||||
|
||||
# Use unique title to avoid conflicts
|
||||
unique_title = f"Test Delete Skill {int(time.time())}"
|
||||
|
||||
|
||||
# Use unique title and skill name to avoid conflicts
|
||||
ts = str(int(time.time()))
|
||||
unique_title = f"Test Delete Skill {ts}"
|
||||
|
||||
# Create a skill specifically to delete
|
||||
with create_skill_zip(skill_name) as zip_file:
|
||||
with create_skill_zip(skill_name, unique_suffix=ts) as zip_file:
|
||||
created_skill = litellm.create_skill(
|
||||
display_title=unique_title,
|
||||
files=[zip_file],
|
||||
|
||||
@@ -1300,9 +1300,11 @@ def test_logging_async_cache_hit_sync_call(turn_off_message_logging):
|
||||
"redacted-by-litellm"
|
||||
== standard_logging_object["messages"][0]["content"]
|
||||
)
|
||||
assert {"text": "redacted-by-litellm"} == standard_logging_object[
|
||||
"response"
|
||||
]
|
||||
# response is a full ModelResponse dict (choices format) since d84e5e381acf
|
||||
assert (
|
||||
standard_logging_object["response"]["choices"][0]["message"]["content"]
|
||||
== "redacted-by-litellm"
|
||||
)
|
||||
|
||||
|
||||
def test_logging_standard_payload_failure_call():
|
||||
|
||||
@@ -636,7 +636,7 @@ def test_stream_chunk_builder_openai_prompt_caching():
|
||||
assert response_usage_value == v
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=2)
|
||||
@pytest.mark.flaky(retries=5, delay=2)
|
||||
def test_stream_chunk_builder_openai_audio_output_usage():
|
||||
from pydantic import BaseModel
|
||||
from openai import OpenAI
|
||||
@@ -667,13 +667,15 @@ def test_stream_chunk_builder_openai_audio_output_usage():
|
||||
usage_obj: Optional[litellm.Usage] = None
|
||||
|
||||
for index, chunk in enumerate(chunks):
|
||||
if hasattr(chunk, "usage"):
|
||||
if hasattr(chunk, "usage") and chunk.usage is not None:
|
||||
usage_obj = chunk.usage
|
||||
print(f"chunk usage: {chunk.usage}")
|
||||
print(f"index: {index}")
|
||||
print(f"len chunks: {len(chunks)}")
|
||||
|
||||
print(f"usage_obj: {usage_obj}")
|
||||
if usage_obj is None:
|
||||
pytest.skip("OpenAI did not return usage data in streaming response")
|
||||
response = stream_chunk_builder(chunks=chunks)
|
||||
print(f"response usage: {response.usage}")
|
||||
check_non_streaming_response(response)
|
||||
|
||||
@@ -45,7 +45,8 @@ async def test_global_redaction_on():
|
||||
await asyncio.sleep(1)
|
||||
standard_logging_payload = test_custom_logger.logged_standard_logging_payload
|
||||
assert standard_logging_payload is not None
|
||||
assert standard_logging_payload["response"] == {"text": "redacted-by-litellm"}
|
||||
response = standard_logging_payload["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
|
||||
print(
|
||||
"logged standard logging payload",
|
||||
@@ -75,7 +76,8 @@ async def test_global_redaction_with_dynamic_params(turn_off_message_logging):
|
||||
)
|
||||
|
||||
if turn_off_message_logging is True:
|
||||
assert standard_logging_payload["response"] == {"text": "redacted-by-litellm"}
|
||||
response = standard_logging_payload["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert (
|
||||
standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
|
||||
)
|
||||
@@ -108,7 +110,8 @@ async def test_global_redaction_off_with_dynamic_params(turn_off_message_logging
|
||||
json.dumps(standard_logging_payload, indent=2),
|
||||
)
|
||||
if turn_off_message_logging is True:
|
||||
assert standard_logging_payload["response"] == {"text": "redacted-by-litellm"}
|
||||
response = standard_logging_payload["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert (
|
||||
standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
|
||||
)
|
||||
@@ -390,7 +393,8 @@ async def test_redaction_with_streaming_response():
|
||||
assert standard_logging_payload is not None
|
||||
|
||||
# Verify that redaction worked without pickle errors
|
||||
assert standard_logging_payload["response"] == {"text": "redacted-by-litellm"}
|
||||
response = standard_logging_payload["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
|
||||
print(
|
||||
"logged standard logging payload for streaming with coroutine handling",
|
||||
@@ -477,5 +481,6 @@ async def test_redaction_with_metadata_completion_api():
|
||||
|
||||
# Verify the helper function works correctly - with get_metadata_variable_name_from_kwargs,
|
||||
# the system checks the appropriate field for headers
|
||||
assert standard_logging_payload["response"] == {"text": "redacted-by-litellm"}
|
||||
response = standard_logging_payload["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
|
||||
|
||||
@@ -395,7 +395,7 @@ async def test_anthropic_messages_openai_model_streaming_cost_injection():
|
||||
|
||||
payload = {
|
||||
"model": "openai/gpt-4o",
|
||||
"max_tokens": 10,
|
||||
"max_tokens": 20,
|
||||
"stream": True,
|
||||
"messages": [{"role": "user", "content": "Say 'Hi'"}],
|
||||
}
|
||||
|
||||
@@ -53,4 +53,5 @@ general_settings:
|
||||
forward_client_headers_to_llm_api: true
|
||||
|
||||
litellm_settings:
|
||||
drop_params: true
|
||||
drop_params: true
|
||||
modify_params: true
|
||||
|
||||
@@ -1,110 +1,327 @@
|
||||
import pytest
|
||||
import litellm
|
||||
"""
|
||||
Unit tests for SearXNG Search request/response transformation.
|
||||
|
||||
These tests validate the request payload and response parsing without
|
||||
requiring a live SearXNG instance.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import List, Union
|
||||
from unittest.mock import MagicMock, patch
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
from tests.search_tests.base_search_unit_tests import BaseSearchTest
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.searxng.search.transformation import SearXNGSearchConfig
|
||||
|
||||
|
||||
class TestSearXNGSearch(BaseSearchTest):
|
||||
class TestSearXNGSearchRequestTransformation:
|
||||
"""
|
||||
Tests for SearXNG Search functionality.
|
||||
Tests that SearXNG search requests are transformed into the expected payload.
|
||||
"""
|
||||
|
||||
def get_search_provider(self) -> str:
|
||||
"""
|
||||
Return search_provider for SearXNG Search.
|
||||
"""
|
||||
return "searxng"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_basic_search(self):
|
||||
"""
|
||||
Test basic search functionality with a simple query.
|
||||
Override to handle free (0.0 cost) provider.
|
||||
"""
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm._turn_on_debug()
|
||||
search_provider = self.get_search_provider()
|
||||
print("Search Provider=", search_provider)
|
||||
|
||||
try:
|
||||
response = await litellm.asearch(
|
||||
query="latest developments in AI",
|
||||
search_provider=search_provider,
|
||||
)
|
||||
print("Search response=", response.model_dump_json(indent=4))
|
||||
def setup_method(self):
|
||||
self.config = SearXNGSearchConfig()
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"Response type: {type(response)}")
|
||||
print(f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}")
|
||||
|
||||
# Check if response has expected Search format
|
||||
assert hasattr(response, "results"), "Response should have 'results' attribute"
|
||||
assert hasattr(response, "object"), "Response should have 'object' attribute"
|
||||
assert response.object == "search", f"Expected object='search', got '{response.object}'"
|
||||
|
||||
# Validate results structure
|
||||
assert isinstance(response.results, list), "results should be a list"
|
||||
assert len(response.results) > 0, "Should have at least one result"
|
||||
|
||||
# Check first result structure
|
||||
first_result = response.results[0]
|
||||
assert hasattr(first_result, "title"), "Result should have 'title' attribute"
|
||||
assert hasattr(first_result, "url"), "Result should have 'url' attribute"
|
||||
assert hasattr(first_result, "snippet"), "Result should have 'snippet' attribute"
|
||||
|
||||
print(f"Total results: {len(response.results)}")
|
||||
print(f"First result title: {first_result.title}")
|
||||
print(f"First result URL: {first_result.url}")
|
||||
print(f"First result snippet: {first_result.snippet[:100]}...")
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
assert len(first_result.title) > 0, "Title should not be empty"
|
||||
assert len(first_result.url) > 0, "URL should not be empty"
|
||||
assert len(first_result.snippet) > 0, "Snippet should not be empty"
|
||||
|
||||
# Validate cost tracking in _hidden_params
|
||||
# For SearXNG (free provider), cost can be None or 0.0
|
||||
assert hasattr(response, "_hidden_params"), "Response should have '_hidden_params' attribute"
|
||||
hidden_params = response._hidden_params
|
||||
assert "response_cost" in hidden_params, "_hidden_params should contain 'response_cost'"
|
||||
|
||||
response_cost = hidden_params["response_cost"]
|
||||
# SearXNG is free, so cost can be None or 0.0
|
||||
if response_cost is not None:
|
||||
assert isinstance(response_cost, (int, float)), "response_cost should be a number"
|
||||
assert response_cost >= 0, "response_cost should be non-negative"
|
||||
print(f"Cost tracking: ${response_cost:.6f}")
|
||||
else:
|
||||
print(f"Cost tracking: Free (None)")
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Search call failed: {str(e)}")
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=5)
|
||||
def test_search_with_optional_params(self):
|
||||
"""
|
||||
Test search with optional parameters.
|
||||
Override for SearXNG since it doesn't natively limit results.
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
search_provider = self.get_search_provider()
|
||||
|
||||
response = litellm.search(
|
||||
query="machine learning",
|
||||
search_provider=search_provider,
|
||||
max_results=5,
|
||||
def test_basic_query_request(self):
|
||||
"""Test that a basic query produces the expected SearXNG request params."""
|
||||
result = self.config.transform_search_request(
|
||||
query="artificial intelligence recent news",
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
# Validate response
|
||||
assert hasattr(response, "results"), "Response should have 'results' attribute"
|
||||
assert isinstance(response.results, list), "results should be a list"
|
||||
assert len(response.results) > 0, "Should have at least one result"
|
||||
# Note: SearXNG doesn't natively limit results, so we don't check <= 5
|
||||
|
||||
print(f"\nSearch with optional params validated:")
|
||||
print(f" - Requested max_results: 5")
|
||||
print(f" - Received results: {len(response.results)}")
|
||||
assert "_searxng_params" in result
|
||||
params = result["_searxng_params"]
|
||||
assert params["q"] == "artificial intelligence recent news"
|
||||
assert params["format"] == "json"
|
||||
|
||||
def test_list_query_joined(self):
|
||||
"""Test that a list query is joined into a single string."""
|
||||
result = self.config.transform_search_request(
|
||||
query=["artificial intelligence", "recent news"],
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
params = result["_searxng_params"]
|
||||
assert params["q"] == "artificial intelligence recent news"
|
||||
assert params["format"] == "json"
|
||||
|
||||
def test_country_to_language_mapping(self):
|
||||
"""Test that country codes are mapped to SearXNG language params."""
|
||||
test_cases = {
|
||||
"us": "en",
|
||||
"uk": "en",
|
||||
"de": "de",
|
||||
"fr": "fr",
|
||||
"es": "es",
|
||||
"jp": "ja",
|
||||
"br": "br", # unmapped country passed through as-is
|
||||
}
|
||||
for country, expected_language in test_cases.items():
|
||||
result = self.config.transform_search_request(
|
||||
query="test",
|
||||
optional_params={"country": country},
|
||||
)
|
||||
params = result["_searxng_params"]
|
||||
assert params["language"] == expected_language, (
|
||||
f"country={country} should map to language={expected_language}"
|
||||
)
|
||||
|
||||
def test_max_results_ignored(self):
|
||||
"""Test that max_results is accepted but doesn't add extra params."""
|
||||
result = self.config.transform_search_request(
|
||||
query="test",
|
||||
optional_params={"max_results": 5},
|
||||
)
|
||||
|
||||
params = result["_searxng_params"]
|
||||
assert params["q"] == "test"
|
||||
assert params["format"] == "json"
|
||||
# max_results should not appear in the SearXNG params
|
||||
assert "max_results" not in params
|
||||
|
||||
def test_searxng_specific_params_passthrough(self):
|
||||
"""Test that SearXNG-specific params are passed through as-is."""
|
||||
result = self.config.transform_search_request(
|
||||
query="test",
|
||||
optional_params={"categories": "general,news", "engines": "google,bing", "time_range": "month"},
|
||||
)
|
||||
|
||||
params = result["_searxng_params"]
|
||||
assert params["q"] == "test"
|
||||
assert params["format"] == "json"
|
||||
assert params["categories"] == "general,news"
|
||||
assert params["engines"] == "google,bing"
|
||||
assert params["time_range"] == "month"
|
||||
|
||||
|
||||
class TestSearXNGSearchURLConstruction:
|
||||
"""
|
||||
Tests that the complete URL is built correctly from api_base and request params.
|
||||
"""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = SearXNGSearchConfig()
|
||||
|
||||
def test_url_with_search_suffix(self):
|
||||
"""Test URL construction appends /search."""
|
||||
data = {"_searxng_params": {"q": "test query", "format": "json"}}
|
||||
url = self.config.get_complete_url(
|
||||
api_base="https://searxng.example.com",
|
||||
optional_params={},
|
||||
data=data,
|
||||
)
|
||||
|
||||
parsed = urlparse(url)
|
||||
assert parsed.scheme == "https"
|
||||
assert parsed.netloc == "searxng.example.com"
|
||||
assert parsed.path == "/search"
|
||||
query_params = parse_qs(parsed.query)
|
||||
assert query_params["q"] == ["test query"]
|
||||
assert query_params["format"] == ["json"]
|
||||
|
||||
def test_url_already_has_search_suffix(self):
|
||||
"""Test URL construction doesn't double-append /search."""
|
||||
data = {"_searxng_params": {"q": "test", "format": "json"}}
|
||||
url = self.config.get_complete_url(
|
||||
api_base="https://searxng.example.com/search",
|
||||
optional_params={},
|
||||
data=data,
|
||||
)
|
||||
|
||||
parsed = urlparse(url)
|
||||
assert parsed.path == "/search"
|
||||
assert "/search/search" not in url
|
||||
|
||||
def test_url_with_trailing_slash(self):
|
||||
"""Test URL construction with trailing slash on api_base."""
|
||||
data = {"_searxng_params": {"q": "test", "format": "json"}}
|
||||
url = self.config.get_complete_url(
|
||||
api_base="https://searxng.example.com/",
|
||||
optional_params={},
|
||||
data=data,
|
||||
)
|
||||
|
||||
parsed = urlparse(url)
|
||||
assert parsed.path == "/search"
|
||||
|
||||
def test_url_from_env_variable(self):
|
||||
"""Test URL construction falls back to SEARXNG_API_BASE env var."""
|
||||
data = {"_searxng_params": {"q": "test", "format": "json"}}
|
||||
with patch(
|
||||
"litellm.llms.searxng.search.transformation.get_secret_str",
|
||||
return_value="https://env-searxng.example.com",
|
||||
):
|
||||
url = self.config.get_complete_url(
|
||||
api_base=None,
|
||||
optional_params={},
|
||||
data=data,
|
||||
)
|
||||
|
||||
assert url.startswith("https://env-searxng.example.com/search?")
|
||||
|
||||
def test_url_missing_api_base_raises(self):
|
||||
"""Test that missing api_base and env var raises ValueError."""
|
||||
with patch(
|
||||
"litellm.llms.searxng.search.transformation.get_secret_str",
|
||||
return_value=None,
|
||||
):
|
||||
with pytest.raises(ValueError, match="SEARXNG_API_BASE is not set"):
|
||||
self.config.get_complete_url(
|
||||
api_base=None,
|
||||
optional_params={},
|
||||
data={"_searxng_params": {"q": "test"}},
|
||||
)
|
||||
|
||||
def test_url_without_data_returns_base(self):
|
||||
"""Test URL construction without data returns just the api_base/search."""
|
||||
url = self.config.get_complete_url(
|
||||
api_base="https://searxng.example.com",
|
||||
optional_params={},
|
||||
data=None,
|
||||
)
|
||||
|
||||
assert url == "https://searxng.example.com/search"
|
||||
|
||||
|
||||
class TestSearXNGSearchResponseTransformation:
|
||||
"""
|
||||
Tests that SearXNG API responses are correctly transformed to SearchResponse.
|
||||
"""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = SearXNGSearchConfig()
|
||||
self.logging_obj = MagicMock()
|
||||
|
||||
def _make_mock_response(self, json_data: dict) -> httpx.Response:
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json=json_data,
|
||||
request=httpx.Request("GET", "https://searxng.example.com/search"),
|
||||
)
|
||||
return response
|
||||
|
||||
def test_response_with_results(self):
|
||||
"""Test transforming a typical SearXNG response with results."""
|
||||
raw = self._make_mock_response({
|
||||
"results": [
|
||||
{
|
||||
"title": "AI News Article",
|
||||
"url": "https://example.com/ai-news",
|
||||
"content": "Latest developments in artificial intelligence.",
|
||||
"publishedDate": "2025-01-15",
|
||||
},
|
||||
{
|
||||
"title": "ML Research Paper",
|
||||
"url": "https://example.com/ml-paper",
|
||||
"content": "New machine learning research findings.",
|
||||
"pubdate": "2025-01-10",
|
||||
},
|
||||
]
|
||||
})
|
||||
|
||||
response = self.config.transform_search_response(
|
||||
raw_response=raw, logging_obj=self.logging_obj
|
||||
)
|
||||
|
||||
assert response.object == "search"
|
||||
assert len(response.results) == 2
|
||||
|
||||
first = response.results[0]
|
||||
assert first.title == "AI News Article"
|
||||
assert first.url == "https://example.com/ai-news"
|
||||
assert first.snippet == "Latest developments in artificial intelligence."
|
||||
assert first.date == "2025-01-15"
|
||||
assert first.last_updated is None
|
||||
|
||||
second = response.results[1]
|
||||
assert second.title == "ML Research Paper"
|
||||
assert second.date == "2025-01-10" # from pubdate field
|
||||
|
||||
def test_response_empty_results(self):
|
||||
"""Test transforming a response with no results."""
|
||||
raw = self._make_mock_response({"results": []})
|
||||
|
||||
response = self.config.transform_search_response(
|
||||
raw_response=raw, logging_obj=self.logging_obj
|
||||
)
|
||||
|
||||
assert response.object == "search"
|
||||
assert response.results == []
|
||||
|
||||
def test_response_missing_results_key(self):
|
||||
"""Test transforming a response that has no 'results' key."""
|
||||
raw = self._make_mock_response({"query": "test"})
|
||||
|
||||
response = self.config.transform_search_response(
|
||||
raw_response=raw, logging_obj=self.logging_obj
|
||||
)
|
||||
|
||||
assert response.object == "search"
|
||||
assert response.results == []
|
||||
|
||||
def test_response_missing_optional_fields(self):
|
||||
"""Test transforming results with missing optional fields."""
|
||||
raw = self._make_mock_response({
|
||||
"results": [
|
||||
{
|
||||
"title": "Minimal Result",
|
||||
"url": "https://example.com",
|
||||
}
|
||||
]
|
||||
})
|
||||
|
||||
response = self.config.transform_search_response(
|
||||
raw_response=raw, logging_obj=self.logging_obj
|
||||
)
|
||||
|
||||
result = response.results[0]
|
||||
assert result.title == "Minimal Result"
|
||||
assert result.url == "https://example.com"
|
||||
assert result.snippet == "" # defaults to empty string
|
||||
assert result.date is None
|
||||
assert result.last_updated is None
|
||||
|
||||
|
||||
class TestSearXNGSearchHeaders:
|
||||
"""
|
||||
Tests for header/environment validation.
|
||||
"""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = SearXNGSearchConfig()
|
||||
|
||||
def test_headers_without_api_key(self):
|
||||
"""Test that headers are set correctly without an API key."""
|
||||
with patch(
|
||||
"litellm.llms.searxng.search.transformation.get_secret_str",
|
||||
return_value=None,
|
||||
):
|
||||
headers = self.config.validate_environment(headers={})
|
||||
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
assert "Authorization" not in headers
|
||||
|
||||
def test_headers_with_api_key(self):
|
||||
"""Test that headers include Authorization when API key is provided."""
|
||||
headers = self.config.validate_environment(
|
||||
headers={}, api_key="test-key-123"
|
||||
)
|
||||
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
assert headers["Authorization"] == "Bearer test-key-123"
|
||||
|
||||
def test_headers_with_env_api_key(self):
|
||||
"""Test that headers use SEARXNG_API_KEY from env."""
|
||||
with patch(
|
||||
"litellm.llms.searxng.search.transformation.get_secret_str",
|
||||
return_value="env-key-456",
|
||||
):
|
||||
headers = self.config.validate_environment(headers={})
|
||||
|
||||
assert headers["Authorization"] == "Bearer env-key-456"
|
||||
|
||||
def test_http_method_is_get(self):
|
||||
"""Test that the HTTP method is GET."""
|
||||
assert self.config.get_http_method() == "GET"
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
"""
|
||||
Tests for Serper Search API integration.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
)
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestSerperSearch:
|
||||
"""
|
||||
Tests for Serper Search functionality with mocked network responses.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serper_search_request_payload(self):
|
||||
"""
|
||||
Test that validates the Serper search request payload structure without making real API calls.
|
||||
"""
|
||||
# Set environment variable for API key
|
||||
os.environ["SERPER_API_KEY"] = "test-api-key"
|
||||
|
||||
# Create a mock response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"organic": [
|
||||
{
|
||||
"title": "Test Result 1",
|
||||
"link": "https://example.com/1",
|
||||
"snippet": "This is a test snippet for result 1",
|
||||
"position": 1,
|
||||
},
|
||||
{
|
||||
"title": "Test Result 2",
|
||||
"link": "https://example.com/2",
|
||||
"snippet": "This is a test snippet for result 2",
|
||||
"position": 2,
|
||||
"date": "Jan 15, 2025",
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
# Mock the httpx AsyncClient post method
|
||||
with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Make the search call
|
||||
response = await litellm.asearch(
|
||||
query="latest developments in AI",
|
||||
search_provider="serper",
|
||||
max_results=5
|
||||
)
|
||||
|
||||
# Verify the post method was called once
|
||||
assert mock_post.call_count == 1
|
||||
|
||||
# Get the actual call arguments
|
||||
call_args = mock_post.call_args
|
||||
|
||||
# Verify URL
|
||||
assert call_args.kwargs["url"] == "https://google.serper.dev/search"
|
||||
|
||||
# Verify headers contain X-API-KEY
|
||||
headers = call_args.kwargs.get("headers", {})
|
||||
assert "X-API-KEY" in headers
|
||||
assert headers["X-API-KEY"] == "test-api-key"
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
|
||||
# Verify request payload
|
||||
json_data = call_args.kwargs.get("json")
|
||||
assert json_data is not None
|
||||
assert json_data["q"] == "latest developments in AI"
|
||||
assert json_data["num"] == 5
|
||||
|
||||
# Verify response structure
|
||||
assert hasattr(response, "results")
|
||||
assert hasattr(response, "object")
|
||||
assert response.object == "search"
|
||||
assert len(response.results) == 2
|
||||
|
||||
# Verify first result
|
||||
first_result = response.results[0]
|
||||
assert first_result.title == "Test Result 1"
|
||||
assert first_result.url == "https://example.com/1"
|
||||
assert first_result.snippet == "This is a test snippet for result 1"
|
||||
|
||||
# Verify date on second result
|
||||
second_result = response.results[1]
|
||||
assert second_result.date == "Jan 15, 2025"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serper_search_with_country(self):
|
||||
"""
|
||||
Test that country parameter is mapped to 'gl' in Serper request.
|
||||
"""
|
||||
os.environ["SERPER_API_KEY"] = "test-api-key"
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"organic": [
|
||||
{
|
||||
"title": "Result",
|
||||
"link": "https://example.com",
|
||||
"snippet": "Snippet",
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
await litellm.asearch(
|
||||
query="test query",
|
||||
search_provider="serper",
|
||||
country="US",
|
||||
)
|
||||
|
||||
json_data = mock_post.call_args.kwargs.get("json")
|
||||
assert json_data["gl"] == "us"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serper_search_with_domain_filter(self):
|
||||
"""
|
||||
Test that search_domain_filter is appended as site: clauses to the query.
|
||||
"""
|
||||
os.environ["SERPER_API_KEY"] = "test-api-key"
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"organic": [
|
||||
{
|
||||
"title": "Result",
|
||||
"link": "https://arxiv.org/paper/1",
|
||||
"snippet": "Snippet",
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
await litellm.asearch(
|
||||
query="machine learning",
|
||||
search_provider="serper",
|
||||
search_domain_filter=["arxiv.org", "nature.com"],
|
||||
)
|
||||
|
||||
json_data = mock_post.call_args.kwargs.get("json")
|
||||
assert "site:arxiv.org" in json_data["q"]
|
||||
assert "site:nature.com" in json_data["q"]
|
||||
assert "machine learning" in json_data["q"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serper_search_empty_organic(self):
|
||||
"""
|
||||
Test handling of response with no organic results.
|
||||
"""
|
||||
os.environ["SERPER_API_KEY"] = "test-api-key"
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"searchParameters": {"q": "xyznonexistent"},
|
||||
}
|
||||
|
||||
with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
response = await litellm.asearch(
|
||||
query="xyznonexistent",
|
||||
search_provider="serper",
|
||||
)
|
||||
|
||||
assert response.object == "search"
|
||||
assert len(response.results) == 0
|
||||
@@ -1,20 +1,22 @@
|
||||
"""
|
||||
Unit tests for Bedrock AgentCore transformation — Accept header fix.
|
||||
Unit tests for Bedrock AgentCore transformation.
|
||||
|
||||
Verifies that AmazonAgentCoreConfig.sign_request() sets the
|
||||
Accept: application/json, text/event-stream header required by
|
||||
MCP servers on Bedrock AgentCore.
|
||||
Tests:
|
||||
- Accept header fix (sign_request sets Accept: application/json, text/event-stream)
|
||||
- JSON response parsing fallback chain (_parse_json_response supports multiple schemas)
|
||||
- Streaming Content-Type fallback (JSON responses converted to single-chunk streams)
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../../.."))
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
|
||||
@@ -81,3 +83,237 @@ class TestAgentCoreAcceptHeader:
|
||||
headers = mock_post.call_args.kwargs["headers"]
|
||||
assert "Accept" in headers
|
||||
assert headers["Accept"] == "application/json, text/event-stream"
|
||||
|
||||
|
||||
class TestAgentCoreJsonResponseParsing:
|
||||
"""Tests for _parse_json_response fallback chain."""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
return AmazonAgentCoreConfig()
|
||||
|
||||
def test_parse_json_standard_agentcore_format(self, config):
|
||||
"""Strategy 1: standard {"result": {"content": [{"text": "..."}]}} format."""
|
||||
response_json = {
|
||||
"result": {
|
||||
"role": "assistant",
|
||||
"content": [{"text": "Hello from standard format"}],
|
||||
}
|
||||
}
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == "Hello from standard format"
|
||||
assert parsed["usage"] is None
|
||||
assert parsed["final_message"] == response_json["result"]
|
||||
|
||||
def test_parse_json_strands_format(self, config):
|
||||
"""Strategy 2: Strands {"response": [{"text": "..."}]} format."""
|
||||
response_json = {
|
||||
"response": [
|
||||
{"text": "Based on my research, "},
|
||||
{"text": "iOS 18.2 was released."},
|
||||
]
|
||||
}
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == "Based on my research, iOS 18.2 was released."
|
||||
assert parsed["usage"] is None
|
||||
assert parsed["final_message"] is None
|
||||
|
||||
def test_parse_json_string_result(self, config):
|
||||
"""Strategy 3: plain string {"result": "text"} format."""
|
||||
response_json = {"result": "Simple text response"}
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == "Simple text response"
|
||||
assert parsed["usage"] is None
|
||||
|
||||
def test_parse_json_string_response(self, config):
|
||||
"""Strategy 3: plain string {"response": "text"} format."""
|
||||
response_json = {"response": "Another text response"}
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == "Another text response"
|
||||
assert parsed["usage"] is None
|
||||
|
||||
def test_parse_json_unknown_format_fallback(self, config):
|
||||
"""Strategy 4: unknown keys fall back to raw JSON."""
|
||||
response_json = {"custom_key": "custom_value", "data": [1, 2, 3]}
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == json.dumps(response_json)
|
||||
assert parsed["usage"] is None
|
||||
assert parsed["final_message"] is None
|
||||
|
||||
def test_parse_json_non_dict_response(self, config):
|
||||
"""Guard: non-dict JSON (e.g. array) falls back to raw JSON string."""
|
||||
response_json = [{"text": "array response"}]
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == json.dumps(response_json)
|
||||
assert parsed["usage"] is None
|
||||
assert parsed["final_message"] is None
|
||||
|
||||
def test_parse_json_empty_content_in_result(self, config):
|
||||
"""Standard format with empty content list - preserves existing behavior."""
|
||||
response_json = {
|
||||
"result": {
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
}
|
||||
}
|
||||
parsed = config._parse_json_response(response_json)
|
||||
assert parsed["content"] == ""
|
||||
assert parsed["final_message"] == response_json["result"]
|
||||
|
||||
|
||||
class TestAgentCoreNonStreamingJsonFormats:
|
||||
"""Tests for _get_parsed_response with different JSON formats (non-streaming path)."""
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
return AmazonAgentCoreConfig()
|
||||
|
||||
def test_get_parsed_response_strands_json(self, config):
|
||||
"""
|
||||
Non-streaming path: _get_parsed_response routes application/json
|
||||
to _parse_json_response which handles the Strands format.
|
||||
"""
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.json.return_value = {
|
||||
"response": [{"text": "Strands agent response via non-streaming"}]
|
||||
}
|
||||
parsed = config._get_parsed_response(mock_response)
|
||||
assert parsed["content"] == "Strands agent response via non-streaming"
|
||||
assert parsed["usage"] is None
|
||||
|
||||
def test_get_parsed_response_raw_json_fallback(self, config):
|
||||
"""
|
||||
Non-streaming path: unknown JSON schema falls back to raw JSON string.
|
||||
"""
|
||||
response_json = {"output": "some value"}
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.json.return_value = response_json
|
||||
parsed = config._get_parsed_response(mock_response)
|
||||
assert parsed["content"] == json.dumps(response_json)
|
||||
|
||||
|
||||
class TestAgentCoreStreamingJsonFallback:
|
||||
"""Tests for streaming Content-Type check (JSON -> single-chunk stream)."""
|
||||
|
||||
def test_sync_streaming_with_json_response(self):
|
||||
"""
|
||||
When stream=True but the agent returns Content-Type: application/json,
|
||||
content is extracted and returned instead of silently returning empty.
|
||||
Exercises the full path through litellm.completion().
|
||||
"""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
client = HTTPHandler()
|
||||
json_body = {"response": [{"text": "Strands sync response"}]}
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.read.return_value = json.dumps(json_body).encode()
|
||||
|
||||
with patch.object(client, "post", return_value=mock_response):
|
||||
response = litellm.completion(
|
||||
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
|
||||
# Collect content across all chunks
|
||||
# CustomStreamWrapper yields content chunk(s) + a synthetic stop chunk
|
||||
content = ""
|
||||
for chunk in response:
|
||||
if chunk.choices[0].delta.content:
|
||||
content += chunk.choices[0].delta.content
|
||||
|
||||
assert content == "Strands sync response"
|
||||
|
||||
async def test_async_streaming_with_json_response(self):
|
||||
"""
|
||||
Async streaming: same Content-Type: application/json fallback via
|
||||
litellm.acompletion(stream=True).
|
||||
"""
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
json_body = {"response": [{"text": "Strands async response"}]}
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.aread = AsyncMock(
|
||||
return_value=json.dumps(json_body).encode()
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
client, "post", new_callable=AsyncMock, return_value=mock_response
|
||||
):
|
||||
response = await litellm.acompletion(
|
||||
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
|
||||
# Collect content across all chunks
|
||||
content = ""
|
||||
async for chunk in response:
|
||||
if chunk.choices[0].delta.content:
|
||||
content += chunk.choices[0].delta.content
|
||||
|
||||
assert content == "Strands async response"
|
||||
|
||||
def test_sync_streaming_malformed_json_raises_error(self):
|
||||
"""
|
||||
When stream=True and Content-Type is application/json but the body
|
||||
is malformed JSON, an error is raised with a descriptive message
|
||||
(not a raw JSONDecodeError).
|
||||
"""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
client = HTTPHandler()
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.read.return_value = b"not valid json {{"
|
||||
|
||||
with patch.object(client, "post", return_value=mock_response):
|
||||
with pytest.raises(Exception, match="Failed to read/parse JSON response body"):
|
||||
litellm.completion(
|
||||
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
|
||||
async def test_async_streaming_malformed_json_raises_error(self):
|
||||
"""
|
||||
Async mirror: malformed JSON body raises a structured error, not a
|
||||
raw JSONDecodeError.
|
||||
"""
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.aread = AsyncMock(return_value=b"not valid json {{")
|
||||
|
||||
with patch.object(
|
||||
client, "post", new_callable=AsyncMock, return_value=mock_response
|
||||
):
|
||||
with pytest.raises(Exception, match="Failed to read/parse JSON response body"):
|
||||
await litellm.acompletion(
|
||||
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=True,
|
||||
client=client,
|
||||
)
|
||||
|
||||
@@ -3170,6 +3170,33 @@ def test_transform_request_with_output_config():
|
||||
assert result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["name"] == "TestSchema"
|
||||
|
||||
|
||||
def test_output_config_snake_case_stripped_from_bedrock_converse_request():
|
||||
"""Test that output_config (snake_case) is stripped from Bedrock Converse requests.
|
||||
|
||||
Bedrock Converse API doesn't support the output_config parameter (Anthropic-only).
|
||||
Nova and other Converse models reject requests with extraneous output_config.
|
||||
"""
|
||||
config = AmazonConverseConfig()
|
||||
messages = [{"role": "user", "content": "test"}]
|
||||
optional_params = {
|
||||
"output_config": {"effort": "high"},
|
||||
}
|
||||
|
||||
result = config._transform_request(
|
||||
model="us.amazon.nova-pro-v1:0",
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
# output_config must not appear in additionalModelRequestFields
|
||||
additional = result.get("additionalModelRequestFields", {})
|
||||
assert "output_config" not in additional, (
|
||||
f"output_config should be stripped for Bedrock Converse, got: {list(additional.keys())}"
|
||||
)
|
||||
|
||||
|
||||
def test_transform_response_native_structured_output():
|
||||
"""Test response handling when model returns JSON as text content (native structured output)."""
|
||||
response_json = {
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
"""
|
||||
Tests for Charity Engine provider configuration and integration.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
try:
|
||||
import pytest
|
||||
except ImportError:
|
||||
pytest = None
|
||||
|
||||
# Add workspace to path
|
||||
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
|
||||
sys.path.insert(0, workspace_path)
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestCharityEngineProviderConfig:
|
||||
"""Test Charity Engine provider configuration"""
|
||||
|
||||
def test_charity_engine_in_provider_list(self):
|
||||
"""Test that charity_engine is in the provider list"""
|
||||
from litellm import LlmProviders
|
||||
|
||||
assert hasattr(LlmProviders, "CHARITY_ENGINE")
|
||||
assert LlmProviders.CHARITY_ENGINE.value == "charity_engine"
|
||||
assert "charity_engine" in litellm.provider_list
|
||||
|
||||
def test_charity_engine_json_config_exists(self):
|
||||
"""Test that charity_engine is configured in providers.json"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
assert JSONProviderRegistry.exists("charity_engine")
|
||||
|
||||
charity_engine = JSONProviderRegistry.get("charity_engine")
|
||||
assert charity_engine is not None
|
||||
assert charity_engine.base_url == "https://api.charityengine.services/remotejobs/v2/inference"
|
||||
assert charity_engine.api_key_env == "CHARITY_ENGINE_API_KEY"
|
||||
assert charity_engine.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_charity_engine_provider_resolution(self):
|
||||
"""Test that provider resolution finds charity_engine"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="charity_engine/gemma3:270m",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "gemma3:270m"
|
||||
assert provider == "charity_engine"
|
||||
assert api_base == "https://api.charityengine.services/remotejobs/v2/inference"
|
||||
|
||||
def test_charity_engine_router_config(self):
|
||||
"""Test that charity_engine can be used in Router configuration"""
|
||||
from litellm import Router
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gemma3-270m",
|
||||
"litellm_params": {
|
||||
"model": "charity_engine/gemma3:270m",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert len(router.model_list) == 1
|
||||
assert router.model_list[0]["model_name"] == "gemma3-270m"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Testing Charity Engine Provider...")
|
||||
|
||||
test_config = TestCharityEngineProviderConfig()
|
||||
|
||||
print("\n1. Testing provider in list...")
|
||||
test_config.test_charity_engine_in_provider_list()
|
||||
print(" ✓ charity_engine in provider list")
|
||||
|
||||
print("\n2. Testing JSON config...")
|
||||
test_config.test_charity_engine_json_config_exists()
|
||||
print(" ✓ charity_engine JSON config loaded")
|
||||
|
||||
print("\n3. Testing provider resolution...")
|
||||
test_config.test_charity_engine_provider_resolution()
|
||||
print(" ✓ Provider resolution works")
|
||||
|
||||
print("\n4. Testing router configuration...")
|
||||
test_config.test_charity_engine_router_config()
|
||||
print(" ✓ Router configuration works")
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
print("✓ All configuration tests passed!")
|
||||
print("=" * 50)
|
||||
@@ -207,13 +207,11 @@ def test_watsonx_completion_regular_model_includes_model_id(
|
||||
assert "project_id" in json_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.xdist_group("watsonx_heavy")
|
||||
async def test_watsonx_gpt_oss_prompt_transformation(monkeypatch): # noqa: PLR0915
|
||||
def test_watsonx_gpt_oss_prompt_transformation(monkeypatch):
|
||||
"""
|
||||
Test that gpt-oss-120b model transforms messages to proper format instead of simple concatenation.
|
||||
|
||||
This test starts from litellm.acompletion and verifies what gets sent in the final POST request body.
|
||||
This test calls litellm.completion (sync) and verifies what gets sent in the final POST request body.
|
||||
Input messages should be transformed using the HuggingFace chat template from openai/gpt-oss-120b,
|
||||
not just concatenated as "You are chatgpt Hi there".
|
||||
"""
|
||||
@@ -229,39 +227,12 @@ async def test_watsonx_gpt_oss_prompt_transformation(monkeypatch): # noqa: PLR0
|
||||
{"role": "user", "content": "Hi there"},
|
||||
]
|
||||
|
||||
# Mock the HTTP client
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
# Mock the token call
|
||||
mock_token_response = Mock()
|
||||
mock_token_response.json.return_value = {
|
||||
"access_token": "mock_access_token",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
mock_token_response.raise_for_status = Mock()
|
||||
|
||||
# Mock the completion call
|
||||
mock_completion_response = Mock()
|
||||
mock_completion_response.status_code = 200
|
||||
mock_completion_response.json.return_value = {
|
||||
"results": [
|
||||
{
|
||||
"generated_text": "Hello! How can I help you?",
|
||||
"generated_token_count": 10,
|
||||
"input_token_count": 5,
|
||||
"stop_reason": "stop", # Required field for response transformation
|
||||
}
|
||||
],
|
||||
"model_id": "openai/gpt-oss-120b",
|
||||
}
|
||||
client = HTTPHandler()
|
||||
|
||||
# Mock HuggingFace template fetch to make test deterministic and avoid network flakiness.
|
||||
# The test verifies that prompt transformation occurs (not simple concatenation), not the exact
|
||||
# HuggingFace template format. Using a mock template that produces the correct format is sufficient.
|
||||
from unittest.mock import patch
|
||||
|
||||
#
|
||||
# Mock template that produces gpt-oss-120b-like format.
|
||||
# Note: This is a simplified version of the actual template. The real template is more complex
|
||||
# (adds metadata, handles tools, thinking messages, etc.), but this captures the key aspects:
|
||||
@@ -277,105 +248,46 @@ async def test_watsonx_gpt_oss_prompt_transformation(monkeypatch): # noqa: PLR0
|
||||
},
|
||||
}
|
||||
|
||||
async def mock_aget_tokenizer_config(hf_model_name: str):
|
||||
return mock_tokenizer_config
|
||||
|
||||
async def mock_aget_chat_template_file(hf_model_name: str):
|
||||
# Return failure to use tokenizer_config instead
|
||||
return {"status": "failure"}
|
||||
|
||||
# Set cached tokenizer config directly to avoid race conditions with parallel tests.
|
||||
# When running with pytest-xdist (-n 16), another test might populate the cache between
|
||||
# clearing it and the actual usage. By setting the cache directly, we ensure the correct
|
||||
# template is always used regardless of test execution order.
|
||||
# Isolate known_tokenizer_config so parallel tests don't interfere.
|
||||
# monkeypatch.setitem restores the original value on teardown.
|
||||
hf_model = "openai/gpt-oss-120b"
|
||||
litellm.known_tokenizer_config[hf_model] = mock_tokenizer_config
|
||||
monkeypatch.setitem(litellm.known_tokenizer_config, hf_model, mock_tokenizer_config)
|
||||
|
||||
# Also create sync mock functions in case the fallback sync path is used
|
||||
def mock_get_tokenizer_config(hf_model_name: str):
|
||||
return mock_tokenizer_config
|
||||
|
||||
def mock_get_chat_template_file(hf_model_name: str):
|
||||
return {"status": "failure"}
|
||||
|
||||
# Async mock function for client.post to properly handle async method mocking
|
||||
async def mock_post_func(*args, **kwargs):
|
||||
return mock_completion_response
|
||||
|
||||
# Mock the token generation response to avoid actual API call
|
||||
mock_token_get_response = Mock()
|
||||
mock_token_get_response.json.return_value = {
|
||||
# Mock IAM token generation to avoid real HTTP calls.
|
||||
mock_token_response = Mock()
|
||||
mock_token_response.json.return_value = {
|
||||
"access_token": "mock_access_token",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
mock_token_get_response.raise_for_status = Mock()
|
||||
mock_token_response.raise_for_status = Mock()
|
||||
|
||||
# Pre-populate the WatsonX IAM token cache to avoid any HTTP calls for token generation.
|
||||
# This prevents parallel test interference with litellm.module_level_client.
|
||||
from litellm.llms.watsonx.common_utils import iam_token_cache
|
||||
iam_token_cache.set_cache(key="test_api_key", value="mock_access_token", ttl=3600)
|
||||
|
||||
with patch.object(client, "post", side_effect=mock_post_func) as mock_post, patch.object(
|
||||
litellm.module_level_client, "post", return_value=mock_token_get_response
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._aget_tokenizer_config",
|
||||
side_effect=mock_aget_tokenizer_config,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._aget_chat_template_file",
|
||||
side_effect=mock_aget_chat_template_file,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._get_tokenizer_config",
|
||||
side_effect=mock_get_tokenizer_config,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._get_chat_template_file",
|
||||
side_effect=mock_get_chat_template_file,
|
||||
with patch.object(client, "post") as mock_post, patch.object(
|
||||
litellm.module_level_client, "post", return_value=mock_token_response
|
||||
):
|
||||
try:
|
||||
# Call acompletion with messages
|
||||
await litellm.acompletion(
|
||||
completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_key="test_api_key",
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
# May fail due to incomplete mocking, but we should have captured the request
|
||||
print(f"Exception (may be expected): {e}")
|
||||
print(f"Caught expected exception: {e}")
|
||||
|
||||
# Verify the POST was called
|
||||
assert (
|
||||
mock_post.call_count >= 1
|
||||
), f"POST should have been called at least once, got {mock_post.call_count}"
|
||||
mock_post.call_count == 1
|
||||
), f"POST should have been called exactly once, got {mock_post.call_count}"
|
||||
|
||||
# Get the request body from the first call
|
||||
# Use call_args_list to be more robust - get the first call's arguments
|
||||
assert len(mock_post.call_args_list) > 0, "mock_post should have at least one call"
|
||||
call_args = mock_post.call_args_list[0]
|
||||
assert call_args is not None, "call_args should not be None"
|
||||
# Get the request body
|
||||
call_args = mock_post.call_args
|
||||
assert "data" in call_args.kwargs, "call_args.kwargs should contain 'data'"
|
||||
json_data = json.loads(call_args.kwargs["data"])
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"Input messages to litellm.acompletion:")
|
||||
print(json.dumps(messages, indent=2))
|
||||
print(f"\n{'='*80}")
|
||||
print(f"Final POST request body:")
|
||||
print(json.dumps(json_data, indent=2))
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# Verify the transformed input is in the request
|
||||
assert "input" in json_data, "Request should have 'input' field"
|
||||
transformed_prompt = json_data["input"]
|
||||
|
||||
# Verify transformation occurred
|
||||
assert transformed_prompt is not None, (
|
||||
"Prompt transformation failed - the template should have been applied to transform "
|
||||
"messages into the correct format for gpt-oss-120b."
|
||||
)
|
||||
|
||||
print(f"Transformed prompt: {repr(transformed_prompt)}")
|
||||
print(f"Prompt length: {len(transformed_prompt)}")
|
||||
|
||||
# Verify it's NOT simple concatenation
|
||||
simple_concat = "You are chatgpt Hi there"
|
||||
assert transformed_prompt != simple_concat, (
|
||||
|
||||
@@ -2093,3 +2093,150 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
|
||||
assert spend_meta["tool_count_total"] == 1
|
||||
assert spend_meta["allowed_server_count"] == 1
|
||||
assert spend_meta["per_server_tool_counts"]["server_a"] == 1
|
||||
|
||||
|
||||
def test_tool_name_matches_case_insensitive():
|
||||
"""Test that _tool_name_matches performs case-insensitive comparison.
|
||||
|
||||
This is critical for OpenAPI-based MCP servers where:
|
||||
1. operationIds are often in camelCase (e.g., 'addPet', 'updatePet')
|
||||
2. Tool names are lowercased during registration (e.g., 'addpet', 'updatepet')
|
||||
3. allowed_tools configuration may use the original camelCase names
|
||||
|
||||
Without case-insensitive matching, all tools would be filtered out.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import _tool_name_matches
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Test case 1: Unprefixed tool name with camelCase in filter list
|
||||
assert _tool_name_matches("addpet", ["addPet", "updatePet"]) is True
|
||||
assert _tool_name_matches("updatepet", ["addPet", "updatePet"]) is True
|
||||
assert _tool_name_matches("deletepet", ["addPet", "updatePet"]) is False
|
||||
|
||||
# Test case 2: Prefixed tool name with camelCase in filter list
|
||||
assert _tool_name_matches("per_store-addpet", ["addPet", "updatePet"]) is True
|
||||
assert _tool_name_matches("per_store-updatepet", ["addPet", "updatePet"]) is True
|
||||
assert _tool_name_matches("per_store-deletepet", ["addPet", "updatePet"]) is False
|
||||
|
||||
# Test case 3: Mixed case variations
|
||||
assert _tool_name_matches("findPetsByStatus", ["findpetsbystatus"]) is True
|
||||
assert _tool_name_matches("findpetsbystatus", ["findPetsByStatus"]) is True
|
||||
assert _tool_name_matches("FINDPETSBYSTATUS", ["findPetsByStatus"]) is True
|
||||
|
||||
# Test case 4: Full prefixed name in filter list (case-insensitive)
|
||||
assert _tool_name_matches("server-addPet", ["server-addpet"]) is True
|
||||
assert _tool_name_matches("server-addpet", ["server-addPet"]) is True
|
||||
|
||||
# Test case 5: Ensure non-matching names still don't match
|
||||
assert _tool_name_matches("addpet", ["deletePet", "updatePet"]) is False
|
||||
assert _tool_name_matches("server-addpet", ["deletePet", "updatePet"]) is False
|
||||
|
||||
|
||||
def test_filter_tools_by_allowed_tools_case_insensitive():
|
||||
"""Test that filter_tools_by_allowed_tools handles case-insensitive matching.
|
||||
|
||||
Ensures that OpenAPI tools with lowercase names can be filtered using
|
||||
camelCase allowed_tools configuration from the OpenAPI spec.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
filter_tools_by_allowed_tools,
|
||||
)
|
||||
from litellm.types.mcp_server.tool_registry import MCPTool
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Mock handler function
|
||||
def mock_handler(**kwargs):
|
||||
return kwargs
|
||||
|
||||
# Create mock tools with lowercase names (as registered from OpenAPI)
|
||||
tools = [
|
||||
MCPTool(
|
||||
name="per_store-addpet",
|
||||
description="Add a pet",
|
||||
input_schema={"type": "object"},
|
||||
handler=mock_handler,
|
||||
),
|
||||
MCPTool(
|
||||
name="per_store-updatepet",
|
||||
description="Update a pet",
|
||||
input_schema={"type": "object"},
|
||||
handler=mock_handler,
|
||||
),
|
||||
MCPTool(
|
||||
name="per_store-deletepet",
|
||||
description="Delete a pet",
|
||||
input_schema={"type": "object"},
|
||||
handler=mock_handler,
|
||||
),
|
||||
MCPTool(
|
||||
name="per_store-findpetsbystatus",
|
||||
description="Find pets by status",
|
||||
input_schema={"type": "object"},
|
||||
handler=mock_handler,
|
||||
),
|
||||
]
|
||||
|
||||
# Create mock server with camelCase allowed_tools (as from OpenAPI spec)
|
||||
server = MCPServer(
|
||||
server_id="test-server",
|
||||
name="per_store",
|
||||
transport=MCPTransport.http,
|
||||
allowed_tools=["addPet", "updatePet", "findPetsByStatus"],
|
||||
)
|
||||
|
||||
# Filter tools
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Should return 3 tools (case-insensitive match)
|
||||
assert len(filtered_tools) == 3
|
||||
assert any(t.name == "per_store-addpet" for t in filtered_tools)
|
||||
assert any(t.name == "per_store-updatepet" for t in filtered_tools)
|
||||
assert any(t.name == "per_store-findpetsbystatus" for t in filtered_tools)
|
||||
assert not any(t.name == "per_store-deletepet" for t in filtered_tools)
|
||||
|
||||
|
||||
def test_filter_tools_by_allowed_tools_no_filter():
|
||||
"""Test that filter_tools_by_allowed_tools returns all tools when no filter is set."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
filter_tools_by_allowed_tools,
|
||||
)
|
||||
from litellm.types.mcp_server.tool_registry import MCPTool
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
# Mock handler function
|
||||
def mock_handler(**kwargs):
|
||||
return kwargs
|
||||
|
||||
tools = [
|
||||
MCPTool(
|
||||
name="fusion_litellm_mcp-model_list",
|
||||
description="List models",
|
||||
input_schema={"type": "object"},
|
||||
handler=mock_handler,
|
||||
),
|
||||
MCPTool(
|
||||
name="fusion_litellm_mcp-chat_completion",
|
||||
description="Chat completion",
|
||||
input_schema={"type": "object"},
|
||||
handler=mock_handler,
|
||||
),
|
||||
]
|
||||
|
||||
# Server with no allowed_tools filter
|
||||
server = MCPServer(
|
||||
server_id="test-server",
|
||||
name="fusion_litellm_mcp",
|
||||
transport=MCPTransport.http,
|
||||
allowed_tools=None,
|
||||
)
|
||||
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Should return all tools when no filter is configured
|
||||
assert len(filtered_tools) == 2
|
||||
|
||||
@@ -46,6 +46,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
||||
"url": "http://backend-agent:10001",
|
||||
"name": "Test Agent",
|
||||
}
|
||||
mock_agent.litellm_params = None
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock()
|
||||
|
||||
@@ -295,6 +295,9 @@ class TestAgentRBACInternalUser:
|
||||
return_value=_sample_agent_response()
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
resp = self.internal_client.get(
|
||||
"/v1/agents/agent-123", headers={"Authorization": "Bearer k"}
|
||||
)
|
||||
|
||||
@@ -21,6 +21,140 @@ def test_get_team_models_for_all_models_and_team_only_models():
|
||||
assert set(result) == set(combined_models)
|
||||
|
||||
|
||||
def test_get_team_models_all_proxy_models_includes_access_groups():
|
||||
"""
|
||||
When a team has 'all-proxy-models' and include_model_access_groups=True,
|
||||
the result should include model access group names (e.g. 'claude-model-group')
|
||||
in addition to individual model names.
|
||||
"""
|
||||
from litellm.proxy.auth.model_checks import get_team_models
|
||||
|
||||
team_models = ["all-proxy-models"]
|
||||
proxy_model_list = ["model1", "model2"]
|
||||
model_access_groups = {
|
||||
"group-a": ["model1"],
|
||||
"group-b": ["model2"],
|
||||
}
|
||||
|
||||
result = get_team_models(
|
||||
team_models, proxy_model_list, model_access_groups, include_model_access_groups=True
|
||||
)
|
||||
assert "group-a" in result
|
||||
assert "group-b" in result
|
||||
assert "model1" in result
|
||||
assert "model2" in result
|
||||
assert len(result) == len(set(result)), "result should have no duplicates"
|
||||
|
||||
|
||||
def test_get_team_models_all_proxy_models_without_include_flag():
|
||||
"""
|
||||
When include_model_access_groups=False, access group names should NOT
|
||||
appear in the result even with 'all-proxy-models'.
|
||||
"""
|
||||
from litellm.proxy.auth.model_checks import get_team_models
|
||||
|
||||
team_models = ["all-proxy-models"]
|
||||
proxy_model_list = ["model1", "model2"]
|
||||
model_access_groups = {
|
||||
"group-a": ["model1"],
|
||||
"group-b": ["model2"],
|
||||
}
|
||||
|
||||
result = get_team_models(
|
||||
team_models, proxy_model_list, model_access_groups, include_model_access_groups=False
|
||||
)
|
||||
assert "group-a" not in result
|
||||
assert "group-b" not in result
|
||||
assert "model1" in result
|
||||
assert "model2" in result
|
||||
|
||||
|
||||
def test_get_key_models_all_proxy_models_includes_access_groups():
|
||||
"""
|
||||
When a key has 'all-proxy-models' and include_model_access_groups=True,
|
||||
the result should include model access group names.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.model_checks import get_key_models
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
models=["all-proxy-models"],
|
||||
api_key="test-key",
|
||||
)
|
||||
proxy_model_list = ["model1", "model2"]
|
||||
model_access_groups = {
|
||||
"group-a": ["model1"],
|
||||
}
|
||||
|
||||
result = get_key_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
include_model_access_groups=True,
|
||||
)
|
||||
assert "group-a" in result
|
||||
assert "model1" in result
|
||||
assert "model2" in result
|
||||
assert len(result) == len(set(result)), "result should have no duplicates"
|
||||
|
||||
|
||||
def test_get_key_models_passes_include_model_access_groups():
|
||||
"""
|
||||
When a key explicitly has an access group name in its models list and
|
||||
include_model_access_groups=True, the group name should be retained
|
||||
(not stripped by _get_models_from_access_groups).
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.model_checks import get_key_models
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
models=["group-a"],
|
||||
api_key="test-key",
|
||||
)
|
||||
proxy_model_list = ["model1", "model2"]
|
||||
model_access_groups = {
|
||||
"group-a": ["model1", "model2"],
|
||||
}
|
||||
|
||||
result = get_key_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
include_model_access_groups=True,
|
||||
)
|
||||
assert "group-a" in result
|
||||
assert "model1" in result
|
||||
assert "model2" in result
|
||||
|
||||
|
||||
def test_get_key_models_does_not_mutate_input():
|
||||
"""
|
||||
get_key_models must not mutate user_api_key_dict.models in-place.
|
||||
_get_models_from_access_groups uses .pop()/.extend() which would corrupt
|
||||
cached UserAPIKeyAuth objects if all_models were an alias instead of a copy.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.model_checks import get_key_models
|
||||
|
||||
original_models = ["group-a", "extra-model"]
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
models=list(original_models), # give it a list
|
||||
api_key="test-key",
|
||||
)
|
||||
model_access_groups = {
|
||||
"group-a": ["model1", "model2"],
|
||||
}
|
||||
|
||||
_ = get_key_models(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_model_list=["model1", "model2"],
|
||||
model_access_groups=model_access_groups,
|
||||
include_model_access_groups=False,
|
||||
)
|
||||
# The original models list on the auth object must be unchanged
|
||||
assert user_api_key_dict.models == original_models
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_models,team_models,proxy_model_list,model_list,expected",
|
||||
[
|
||||
|
||||
@@ -13,8 +13,12 @@ from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.proxy.proxy_server
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.user_api_key_auth import get_api_key
|
||||
from litellm.proxy.auth.user_api_key_auth import get_api_key, user_api_key_auth
|
||||
|
||||
|
||||
def test_get_api_key():
|
||||
@@ -515,3 +519,169 @@ def test_proxy_admin_jwt_auth_handles_no_team_object():
|
||||
assert result.team_metadata is None
|
||||
assert result.org_id is None
|
||||
assert result.end_user_id is None
|
||||
|
||||
|
||||
class TestJWTOAuth2Coexistence:
|
||||
"""
|
||||
Test that JWT and OAuth2 auth can coexist on the same instance.
|
||||
|
||||
When both enable_jwt_auth and enable_oauth2_auth are True, the proxy should
|
||||
route tokens based on their format:
|
||||
- JWT tokens (3 dot-separated parts) -> JWT auth handler
|
||||
- Opaque tokens -> OAuth2 auth handler
|
||||
"""
|
||||
|
||||
def test_is_jwt_detects_jwt_tokens(self):
|
||||
"""JWT tokens have 3 dot-separated parts."""
|
||||
assert JWTHandler.is_jwt("header.payload.signature") is True
|
||||
assert JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") is True
|
||||
|
||||
def test_is_jwt_rejects_opaque_tokens(self):
|
||||
"""Opaque OAuth2 tokens do not have 3 dot-separated parts."""
|
||||
assert JWTHandler.is_jwt("some-opaque-oauth2-token") is False
|
||||
assert JWTHandler.is_jwt("sk-12345678") is False
|
||||
assert JWTHandler.is_jwt("Bearer token") is False
|
||||
assert JWTHandler.is_jwt("two.parts") is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_both_enabled_opaque_token_uses_oauth2(self):
|
||||
"""
|
||||
When both enable_jwt_auth and enable_oauth2_auth are True,
|
||||
an opaque token should be handled by OAuth2 auth (not JWT).
|
||||
"""
|
||||
opaque_token = "some-opaque-m2m-oauth2-token"
|
||||
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": True,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
|
||||
mock_oauth2_response = UserAPIKeyAuth(
|
||||
api_key=opaque_token,
|
||||
user_id="machine-client-1",
|
||||
team_id="m2m-team",
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.headers = {"authorization": f"Bearer {opaque_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", general_settings), \
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), \
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"), \
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None), \
|
||||
patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock, return_value=mock_oauth2_response) as mock_oauth2, \
|
||||
patch("litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", new_callable=AsyncMock) as mock_jwt_auth:
|
||||
|
||||
litellm.proxy.proxy_server.jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(),
|
||||
)
|
||||
|
||||
result = await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {opaque_token}",
|
||||
)
|
||||
|
||||
# OAuth2 SHOULD be called for opaque tokens
|
||||
mock_oauth2.assert_called_once_with(token=opaque_token)
|
||||
# JWT auth should NOT be called
|
||||
mock_jwt_auth.assert_not_called()
|
||||
assert result.user_id == "machine-client-1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_both_enabled_jwt_token_skips_oauth2(self):
|
||||
"""
|
||||
When both enable_jwt_auth and enable_oauth2_auth are True,
|
||||
a JWT-formatted token should skip OAuth2 and reach the JWT handler.
|
||||
"""
|
||||
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
|
||||
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": True,
|
||||
"enable_jwt_auth": True,
|
||||
}
|
||||
|
||||
mock_jwt_result = {
|
||||
"is_proxy_admin": True,
|
||||
"team_object": None,
|
||||
"user_object": None,
|
||||
"end_user_object": None,
|
||||
"org_object": None,
|
||||
"token": jwt_token,
|
||||
"team_id": "jwt-team",
|
||||
"user_id": "jwt-human-user",
|
||||
"end_user_id": None,
|
||||
"org_id": None,
|
||||
"team_membership": None,
|
||||
"jwt_claims": {"sub": "user1"},
|
||||
}
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", general_settings), \
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), \
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"), \
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None), \
|
||||
patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock) as mock_oauth2, \
|
||||
patch("litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", new_callable=AsyncMock, return_value=mock_jwt_result) as mock_jwt_auth:
|
||||
|
||||
litellm.proxy.proxy_server.jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(),
|
||||
)
|
||||
|
||||
result = await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {jwt_token}",
|
||||
)
|
||||
|
||||
# OAuth2 should NOT be called for JWT tokens
|
||||
mock_oauth2.assert_not_called()
|
||||
# JWT auth SHOULD be called
|
||||
mock_jwt_auth.assert_called_once()
|
||||
assert result.user_id == "jwt-human-user"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_oauth2_enabled_handles_all_tokens(self):
|
||||
"""
|
||||
When only enable_oauth2_auth is True (no JWT), all LLM API tokens
|
||||
should go through OAuth2 - backward compatible behavior.
|
||||
"""
|
||||
jwt_like_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
|
||||
|
||||
general_settings = {
|
||||
"enable_oauth2_auth": True,
|
||||
"enable_jwt_auth": False,
|
||||
}
|
||||
|
||||
mock_oauth2_response = UserAPIKeyAuth(
|
||||
api_key=jwt_like_token,
|
||||
user_id="oauth2-user",
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.headers = {"authorization": f"Bearer {jwt_like_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", general_settings), \
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), \
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"), \
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None), \
|
||||
patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock, return_value=mock_oauth2_response) as mock_oauth2:
|
||||
|
||||
result = await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {jwt_like_token}",
|
||||
)
|
||||
|
||||
# OAuth2 should handle it since JWT auth is disabled
|
||||
mock_oauth2.assert_called_once_with(token=jwt_like_token)
|
||||
assert result.user_id == "oauth2-user"
|
||||
|
||||
@@ -1643,4 +1643,89 @@ async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch)
|
||||
model="gpt-4",
|
||||
api_key=None,
|
||||
timezone_offset_minutes=480,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
|
||||
"""
|
||||
Test that delete_user removes invitation links where the deleted user is the
|
||||
creator (created_by) or updater (updated_by), not just the invited person (user_id).
|
||||
|
||||
This prevents FK constraint violations when deleting a user who created pending invites.
|
||||
"""
|
||||
from litellm.proxy._types import DeleteUserRequest, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user
|
||||
|
||||
mock_prisma_client = mocker.MagicMock()
|
||||
|
||||
# Mock user lookup
|
||||
mock_user_row = mocker.MagicMock()
|
||||
mock_user_row.user_id = "admin-creator"
|
||||
mock_user_row.user_email = "admin@example.com"
|
||||
mock_user_row.teams = []
|
||||
mock_user_row.json.return_value = "{}"
|
||||
mock_user_row.model_dump.return_value = {
|
||||
"user_id": "admin-creator",
|
||||
"user_email": "admin@example.com",
|
||||
"teams": [],
|
||||
}
|
||||
|
||||
async def mock_find_unique(*args, **kwargs):
|
||||
return mock_user_row
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(
|
||||
side_effect=mock_find_unique
|
||||
)
|
||||
|
||||
# Mock find_many for teams (no teams)
|
||||
mock_prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(
|
||||
return_value=[]
|
||||
)
|
||||
|
||||
# Mock all delete_many calls
|
||||
mock_prisma_client.db.litellm_verificationtoken.delete_many = mocker.AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
mock_prisma_client.db.litellm_invitationlink.delete_many = mocker.AsyncMock(
|
||||
return_value=1
|
||||
)
|
||||
mock_prisma_client.db.litellm_organizationmembership.delete_many = mocker.AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
mock_prisma_client.db.litellm_teammembership.delete_many = mocker.AsyncMock(
|
||||
return_value=0
|
||||
)
|
||||
mock_prisma_client.db.litellm_usertable.delete_many = mocker.AsyncMock(
|
||||
return_value=1
|
||||
)
|
||||
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
# Call delete_user
|
||||
data = DeleteUserRequest(user_ids=["admin-creator"])
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_id="proxy-admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
await delete_user(data=data, user_api_key_dict=user_api_key_dict)
|
||||
|
||||
# Verify invitation link deletion uses OR with user_id, created_by, updated_by
|
||||
mock_prisma_client.db.litellm_invitationlink.delete_many.assert_called_once()
|
||||
call_kwargs = mock_prisma_client.db.litellm_invitationlink.delete_many.call_args
|
||||
where_clause = call_kwargs.kwargs.get("where") or call_kwargs[1].get("where")
|
||||
|
||||
assert "OR" in where_clause, "Should use OR to match user_id, created_by, and updated_by"
|
||||
or_conditions = where_clause["OR"]
|
||||
assert len(or_conditions) == 3, "Should have 3 OR conditions"
|
||||
|
||||
# Verify all three FK fields are covered
|
||||
condition_keys = [list(c.keys())[0] for c in or_conditions]
|
||||
assert "user_id" in condition_keys
|
||||
assert "created_by" in condition_keys
|
||||
assert "updated_by" in condition_keys
|
||||
|
||||
# Verify each condition uses {"in": ["admin-creator"]}
|
||||
for condition in or_conditions:
|
||||
field = list(condition.keys())[0]
|
||||
assert condition[field] == {"in": ["admin-creator"]}
|
||||
@@ -1071,9 +1071,10 @@ def test_spend_logs_redacts_request_and_response_when_turn_off_message_logging_e
|
||||
response_result = _get_response_for_spend_logs_payload(payload=payload, kwargs=kwargs)
|
||||
|
||||
# When redaction is enabled and response is a dict (not ModelResponse),
|
||||
# perform_redaction returns {"text": "redacted-by-litellm"}
|
||||
# perform_redaction redacts content in-place within the choices structure
|
||||
parsed_response = json.loads(response_result)
|
||||
assert parsed_response == {"text": "redacted-by-litellm"}
|
||||
assert parsed_response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert parsed_response["choices"][0]["message"]["role"] == "assistant"
|
||||
|
||||
|
||||
@patch("litellm.secret_managers.main.get_secret_bool")
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
"""
|
||||
Test that LiteLLM_Params and GenericLiteLLMParams handle reserved keys gracefully.
|
||||
|
||||
This test verifies the fix for the bug where passing a dict containing 'self',
|
||||
'params', or '__class__' keys to LiteLLM_Params() would cause:
|
||||
TypeError: LiteLLM_Params.__init__() got multiple values for argument 'self'
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.types.router import GenericLiteLLMParams, LiteLLM_Params
|
||||
|
||||
|
||||
class TestLiteLLMParamsReservedKeys:
|
||||
"""Test that reserved keys in input data are filtered out gracefully."""
|
||||
|
||||
def test_litellm_params_with_self_key(self):
|
||||
"""Test LiteLLM_Params handles 'self' key in input dict."""
|
||||
params_dict = {"model": "gpt-4", "self": "some_value", "api_key": "test-key"}
|
||||
params = LiteLLM_Params(**params_dict)
|
||||
assert params.model == "gpt-4"
|
||||
assert params.api_key == "test-key"
|
||||
assert not hasattr(params, "self") or params.get("self") is None
|
||||
|
||||
def test_litellm_params_with_params_key(self):
|
||||
"""Test LiteLLM_Params handles 'params' key in input dict."""
|
||||
params_dict = {"model": "gpt-4", "params": "bad_value"}
|
||||
params = LiteLLM_Params(**params_dict)
|
||||
assert params.model == "gpt-4"
|
||||
|
||||
def test_litellm_params_with_class_key(self):
|
||||
"""Test LiteLLM_Params handles '__class__' key in input dict."""
|
||||
params_dict = {"model": "gpt-4", "__class__": "bad_value"}
|
||||
params = LiteLLM_Params(**params_dict)
|
||||
assert params.model == "gpt-4"
|
||||
|
||||
def test_generic_litellm_params_with_self_key(self):
|
||||
"""Test GenericLiteLLMParams handles 'self' key in input dict."""
|
||||
params_dict = {"self": "some_value", "api_key": "test-key"}
|
||||
params = GenericLiteLLMParams(**params_dict)
|
||||
assert params.api_key == "test-key"
|
||||
|
||||
def test_generic_litellm_params_with_params_key(self):
|
||||
"""Test GenericLiteLLMParams handles 'params' key in input dict."""
|
||||
params_dict = {"params": "bad_value", "api_key": "test-key"}
|
||||
params = GenericLiteLLMParams(**params_dict)
|
||||
assert params.api_key == "test-key"
|
||||
|
||||
def test_generic_litellm_params_with_class_key(self):
|
||||
"""Test GenericLiteLLMParams handles '__class__' key in input dict."""
|
||||
params_dict = {"__class__": "bad_value", "api_key": "test-key"}
|
||||
params = GenericLiteLLMParams(**params_dict)
|
||||
assert params.api_key == "test-key"
|
||||
|
||||
def test_max_retries_string_conversion(self):
|
||||
"""Test that max_retries is converted from string to int."""
|
||||
params = LiteLLM_Params(model="gpt-4", max_retries="5")
|
||||
assert params.max_retries == 5
|
||||
assert isinstance(params.max_retries, int)
|
||||
|
||||
def test_extra_fields_preserved(self):
|
||||
"""Test that extra fields are preserved when reserved keys are filtered."""
|
||||
params_dict = {
|
||||
"model": "gpt-4",
|
||||
"self": "ignored",
|
||||
"custom_field": "custom_value",
|
||||
}
|
||||
params = LiteLLM_Params(**params_dict)
|
||||
assert params.model == "gpt-4"
|
||||
assert params.custom_field == "custom_value"
|
||||
|
||||
def test_normal_instantiation_still_works(self):
|
||||
"""Test that normal instantiation without reserved keys works."""
|
||||
params = LiteLLM_Params(
|
||||
model="gpt-4", api_key="test-key", custom_llm_provider="openai"
|
||||
)
|
||||
assert params.model == "gpt-4"
|
||||
assert params.api_key == "test-key"
|
||||
assert params.custom_llm_provider == "openai"
|
||||
|
||||
def test_multiple_reserved_keys(self):
|
||||
"""Test filtering multiple reserved keys at once."""
|
||||
params_dict = {
|
||||
"model": "gpt-4",
|
||||
"self": "value1",
|
||||
"params": "value2",
|
||||
"__class__": "value3",
|
||||
"api_key": "test-key",
|
||||
}
|
||||
params = LiteLLM_Params(**params_dict)
|
||||
assert params.model == "gpt-4"
|
||||
assert params.api_key == "test-key"
|
||||
@@ -609,6 +609,24 @@ def test_responses_api_bridge_check_strips_responses_prefix():
|
||||
assert model_info["mode"] == "responses"
|
||||
|
||||
|
||||
def test_responses_api_bridge_check_gpt_5_4_pro():
|
||||
"""Test that gpt-5.4-pro routes through responses API bridge, not chat completions.
|
||||
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/23014
|
||||
gpt-5.4-pro is a responses-only model and must not be sent to /v1/chat/completions.
|
||||
"""
|
||||
from litellm.main import responses_api_bridge_check
|
||||
|
||||
for model_name in ["gpt-5.4-pro", "gpt-5.4-pro-2026-03-05"]:
|
||||
model_info, model = responses_api_bridge_check(
|
||||
model=model_name,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
assert model_info.get("mode") == "responses", (
|
||||
f"{model_name} should have mode='responses', got '{model_info.get('mode')}'"
|
||||
)
|
||||
|
||||
|
||||
def test_responses_api_bridge_check_handles_exception():
|
||||
"""Test that responses_api_bridge_check handles exceptions and still processes responses/ models."""
|
||||
from litellm.main import responses_api_bridge_check
|
||||
|
||||
@@ -507,6 +507,7 @@ def validate_model_cost_values(model_data, exceptions=None):
|
||||
"input_cost_per_audio_token",
|
||||
"output_cost_per_audio_token",
|
||||
"output_cost_per_image_token",
|
||||
"output_cost_per_image_token_batches",
|
||||
"input_cost_per_audio_per_second",
|
||||
"input_cost_per_video_per_second",
|
||||
"input_cost_per_token_above_128k_tokens",
|
||||
@@ -696,6 +697,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
||||
"output_cost_per_character_above_128k_tokens": {"type": "number"},
|
||||
"output_cost_per_image": {"type": "number"},
|
||||
"output_cost_per_image_token": {"type": "number"},
|
||||
"output_cost_per_image_token_batches": {"type": "number"},
|
||||
"output_cost_per_pixel": {"type": "number"},
|
||||
"output_cost_per_second": {"type": "number"},
|
||||
"output_cost_per_token": {"type": "number"},
|
||||
|
||||
@@ -262,8 +262,8 @@ it("should display user email correctly", async () => {
|
||||
});
|
||||
});
|
||||
|
||||
it("should show skeleton loaders when isLoading is true", () => {
|
||||
// Mock loading state
|
||||
it("should show loading message only on initial load (isPending)", () => {
|
||||
// Mock initial loading state
|
||||
mockUseKeys.mockReturnValue({
|
||||
data: null,
|
||||
isPending: true,
|
||||
@@ -283,7 +283,7 @@ it("should show skeleton loaders when isLoading is true", () => {
|
||||
|
||||
renderWithProviders(<VirtualKeysTable {...mockProps} />);
|
||||
|
||||
// Check that loading message is shown
|
||||
// Check that loading message is shown on initial load
|
||||
expect(screen.getByText("🚅 Loading keys...")).toBeInTheDocument();
|
||||
|
||||
// Check that actual key data is not shown
|
||||
@@ -795,3 +795,79 @@ describe("pagination display – total count and page count", () => {
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("refetch button", () => {
|
||||
it("should show Fetch button in normal state", () => {
|
||||
renderWithProviders(<VirtualKeysTable {...defaultMockProps} />);
|
||||
|
||||
const fetchButton = screen.getByTitle("Fetch data");
|
||||
expect(fetchButton).toBeInTheDocument();
|
||||
expect(fetchButton).not.toBeDisabled();
|
||||
expect(screen.getByText("Fetch")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show Fetching state and keep table data visible during refetch", () => {
|
||||
mockUseKeys.mockReturnValue({
|
||||
data: {
|
||||
keys: [mockKey],
|
||||
total_count: 1,
|
||||
current_page: 1,
|
||||
total_pages: 1,
|
||||
} as KeysResponse,
|
||||
isPending: false,
|
||||
isFetching: true,
|
||||
refetch: vi.fn(),
|
||||
} as any);
|
||||
|
||||
renderWithProviders(<VirtualKeysTable {...defaultMockProps} />);
|
||||
|
||||
// Button should show "Fetching" and be disabled
|
||||
expect(screen.getByText("Fetching")).toBeInTheDocument();
|
||||
const fetchButton = screen.getByTitle("Fetch data");
|
||||
expect(fetchButton).toBeDisabled();
|
||||
|
||||
// Table data should still be visible (stale data)
|
||||
expect(screen.getByText("Test Key Alias")).toBeInTheDocument();
|
||||
|
||||
// "Loading keys..." should NOT appear during refetch
|
||||
expect(screen.queryByText("🚅 Loading keys...")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call refetch when Fetch button is clicked", () => {
|
||||
const mockRefetch = vi.fn();
|
||||
mockUseKeys.mockReturnValue({
|
||||
data: {
|
||||
keys: [mockKey],
|
||||
total_count: 1,
|
||||
current_page: 1,
|
||||
total_pages: 1,
|
||||
} as KeysResponse,
|
||||
isPending: false,
|
||||
isFetching: false,
|
||||
refetch: mockRefetch,
|
||||
} as any);
|
||||
|
||||
renderWithProviders(<VirtualKeysTable {...defaultMockProps} />);
|
||||
|
||||
const fetchButton = screen.getByTitle("Fetch data");
|
||||
fireEvent.click(fetchButton);
|
||||
|
||||
expect(mockRefetch).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should show Fetch button enabled on error so user can retry", () => {
|
||||
mockUseKeys.mockReturnValue({
|
||||
data: null,
|
||||
isPending: false,
|
||||
isFetching: false,
|
||||
isError: true,
|
||||
refetch: vi.fn(),
|
||||
} as any);
|
||||
|
||||
renderWithProviders(<VirtualKeysTable {...defaultMockProps} />);
|
||||
|
||||
const fetchButton = screen.getByTitle("Fetch data");
|
||||
expect(fetchButton).not.toBeDisabled();
|
||||
expect(screen.getByText("Fetch")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -24,9 +24,9 @@ import {
|
||||
TableRow,
|
||||
Text,
|
||||
} from "@tremor/react";
|
||||
import { InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { Popover, Skeleton, Tooltip } from "antd";
|
||||
import React, { useEffect, useMemo, useState } from "react";
|
||||
import { InfoCircleOutlined, SyncOutlined } from "@ant-design/icons";
|
||||
import { Button as AntButton, Popover, Skeleton, Tooltip } from "antd";
|
||||
import React, { useEffect, useDeferredValue, useMemo, useState } from "react";
|
||||
import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key";
|
||||
import { useFilterLogic } from "../key_team_helpers/filter_logic";
|
||||
import { PaginatedKeyAliasSelect } from "../KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect";
|
||||
@@ -81,6 +81,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
|
||||
data: keys,
|
||||
isPending: isLoading,
|
||||
isFetching,
|
||||
isError,
|
||||
refetch,
|
||||
} = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize, {
|
||||
sortBy: sortBy || undefined,
|
||||
@@ -97,6 +98,15 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
|
||||
organizations,
|
||||
});
|
||||
|
||||
// Defer the transition so the button stays in loading state until the table
|
||||
// has rendered with the new data (mirrors the spend-logs pattern)
|
||||
const isFetchingDeferred = useDeferredValue(isFetching);
|
||||
const isButtonLoading = (isFetching || isFetchingDeferred) && !isError;
|
||||
|
||||
const handleRefresh = () => {
|
||||
refetch();
|
||||
};
|
||||
|
||||
const totalCount = filteredTotalCount ?? keys?.total_count ?? 0;
|
||||
|
||||
// Add a useEffect to call refresh when a key is created
|
||||
@@ -606,16 +616,28 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
|
||||
</div>
|
||||
|
||||
<div className="flex items-center justify-between w-full mb-4">
|
||||
{isLoading || isFetching ? (
|
||||
<Skeleton.Node active style={{ width: 200, height: 20 }} />
|
||||
) : (
|
||||
<span className="inline-flex text-sm text-gray-700">
|
||||
Showing {rangeLabel} of {totalCount} results
|
||||
</span>
|
||||
)}
|
||||
<div className="inline-flex items-center gap-2">
|
||||
{isLoading ? (
|
||||
<Skeleton.Node active style={{ width: 200, height: 20 }} />
|
||||
) : (
|
||||
<span className="inline-flex text-sm text-gray-700">
|
||||
Showing {rangeLabel} of {totalCount} results
|
||||
</span>
|
||||
)}
|
||||
|
||||
<AntButton
|
||||
type="default"
|
||||
icon={<SyncOutlined spin={isButtonLoading} />}
|
||||
onClick={handleRefresh}
|
||||
disabled={isButtonLoading}
|
||||
title="Fetch data"
|
||||
>
|
||||
{isButtonLoading ? "Fetching" : "Fetch"}
|
||||
</AntButton>
|
||||
</div>
|
||||
|
||||
<div className="inline-flex items-center gap-2">
|
||||
{isLoading || isFetching ? (
|
||||
{isLoading ? (
|
||||
<Skeleton.Node active style={{ width: 74, height: 20 }} />
|
||||
) : (
|
||||
<span className="text-sm text-gray-700">
|
||||
@@ -623,24 +645,24 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
|
||||
</span>
|
||||
)}
|
||||
|
||||
{isLoading || isFetching ? (
|
||||
{isLoading ? (
|
||||
<Skeleton.Button active size="small" style={{ width: 84, height: 30 }} />
|
||||
) : (
|
||||
<button
|
||||
onClick={() => table.previousPage()}
|
||||
disabled={isLoading || isFetching || !table.getCanPreviousPage()}
|
||||
disabled={isLoading || !table.getCanPreviousPage()}
|
||||
className="px-3 py-1 text-sm border rounded-md hover:bg-gray-50 disabled:opacity-50 disabled:cursor-not-allowed"
|
||||
>
|
||||
Previous
|
||||
</button>
|
||||
)}
|
||||
|
||||
{isLoading || isFetching ? (
|
||||
{isLoading ? (
|
||||
<Skeleton.Button active size="small" style={{ width: 58, height: 30 }} />
|
||||
) : (
|
||||
<button
|
||||
onClick={() => table.nextPage()}
|
||||
disabled={isLoading || isFetching || !table.getCanNextPage()}
|
||||
disabled={isLoading || !table.getCanNextPage()}
|
||||
className="px-3 py-1 text-sm border rounded-md hover:bg-gray-50 disabled:opacity-50 disabled:cursor-not-allowed"
|
||||
>
|
||||
Next
|
||||
@@ -725,7 +747,7 @@ export function VirtualKeysTable({ teams, organizations, onSortChange, currentSo
|
||||
))}
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{isLoading || isFetching ? (
|
||||
{isLoading ? (
|
||||
<TableRow>
|
||||
<TableCell colSpan={columns.length} className="h-8 text-center">
|
||||
<div className="text-center text-gray-500">
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { afterEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { ClaudeCodeModal } from "./ClaudeCodeModal";
|
||||
|
||||
describe("ClaudeCodeModal", () => {
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
});
|
||||
|
||||
it("should render nothing when isOpen is false", () => {
|
||||
renderWithProviders(
|
||||
<ClaudeCodeModal isOpen={false} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
expect(screen.queryByText(/Help us improve your experience/i)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the feedback modal content when isOpen is true", () => {
|
||||
renderWithProviders(
|
||||
<ClaudeCodeModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByText(/Help us improve your experience/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the survey description text", () => {
|
||||
renderWithProviders(
|
||||
<ClaudeCodeModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByText(/your experience using LiteLLM with Claude Code/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should open the Google Form and call onComplete when the feedback button is clicked", async () => {
|
||||
const onComplete = vi.fn();
|
||||
const openSpy = vi.spyOn(window, "open").mockImplementation(() => null);
|
||||
const user = userEvent.setup();
|
||||
|
||||
renderWithProviders(
|
||||
<ClaudeCodeModal isOpen={true} onClose={vi.fn()} onComplete={onComplete} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Open Feedback Form/i }));
|
||||
|
||||
expect(openSpy).toHaveBeenCalledWith(
|
||||
"https://forms.gle/LZeJQ3XytBakckYa9",
|
||||
"_blank",
|
||||
"noopener,noreferrer"
|
||||
);
|
||||
expect(onComplete).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should call onClose when the close button is clicked", async () => {
|
||||
const onClose = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
||||
renderWithProviders(
|
||||
<ClaudeCodeModal isOpen={true} onClose={onClose} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
// The X close button is the first button; the "Open Feedback Form" button is the second
|
||||
const buttons = screen.getAllByRole("button");
|
||||
await user.click(buttons[0]);
|
||||
|
||||
expect(onClose).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,82 @@
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { ClaudeCodePrompt } from "./ClaudeCodePrompt";
|
||||
|
||||
vi.mock("./NudgePrompt", () => ({
|
||||
NudgePrompt: ({
|
||||
title,
|
||||
description,
|
||||
buttonText,
|
||||
onOpen,
|
||||
onDismiss,
|
||||
isVisible,
|
||||
}: {
|
||||
title: string;
|
||||
description: string;
|
||||
buttonText: string;
|
||||
onOpen: () => void;
|
||||
onDismiss: () => void;
|
||||
isVisible: boolean;
|
||||
}) => {
|
||||
if (!isVisible) return null;
|
||||
return (
|
||||
<div>
|
||||
<span>{title}</span>
|
||||
<span>{description}</span>
|
||||
<button onClick={onOpen}>{buttonText}</button>
|
||||
<button onClick={onDismiss}>Dismiss</button>
|
||||
</div>
|
||||
);
|
||||
},
|
||||
}));
|
||||
|
||||
describe("ClaudeCodePrompt", () => {
|
||||
it("should render with the Claude Code Feedback title when visible", () => {
|
||||
renderWithProviders(
|
||||
<ClaudeCodePrompt isVisible={true} onOpen={vi.fn()} onDismiss={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByText("Claude Code Feedback")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the correct description text", () => {
|
||||
renderWithProviders(
|
||||
<ClaudeCodePrompt isVisible={true} onOpen={vi.fn()} onDismiss={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByText(/Help us improve your Claude Code experience/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onOpen when the share feedback button is clicked", async () => {
|
||||
const onOpen = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
||||
renderWithProviders(
|
||||
<ClaudeCodePrompt isVisible={true} onOpen={onOpen} onDismiss={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Share feedback/i }));
|
||||
|
||||
expect(onOpen).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should call onDismiss when the dismiss button is clicked", async () => {
|
||||
const onDismiss = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
||||
renderWithProviders(
|
||||
<ClaudeCodePrompt isVisible={true} onOpen={vi.fn()} onDismiss={onDismiss} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Dismiss/i }));
|
||||
|
||||
expect(onDismiss).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not render when isVisible is false", () => {
|
||||
renderWithProviders(
|
||||
<ClaudeCodePrompt isVisible={false} onOpen={vi.fn()} onDismiss={vi.fn()} />
|
||||
);
|
||||
expect(screen.queryByText("Claude Code Feedback")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,200 @@
|
||||
import { screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { SurveyModal } from "./SurveyModal";
|
||||
|
||||
describe("SurveyModal", () => {
|
||||
beforeEach(() => {
|
||||
vi.spyOn(global, "fetch").mockResolvedValue(new Response());
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
});
|
||||
|
||||
it("should render nothing when isOpen is false", () => {
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={false} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
expect(
|
||||
screen.queryByText(/Are you using LiteLLM at your company\?/i)
|
||||
).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render step 1 when the modal is opened", () => {
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
expect(
|
||||
screen.getByText(/Are you using LiteLLM at your company\?/i)
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should disable the Next button until a step 1 choice is made", () => {
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByRole("button", { name: /Next/i })).toBeDisabled();
|
||||
});
|
||||
|
||||
it("should enable the Next button after selecting Yes", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /We use it for work/i }));
|
||||
|
||||
expect(screen.getByRole("button", { name: /Next/i })).not.toBeDisabled();
|
||||
});
|
||||
|
||||
it("should navigate to the company name step when Yes is selected and Next is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /We use it for work/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
|
||||
expect(
|
||||
screen.getByText(/What company are you using LiteLLM at\?/i)
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should skip the company name step when No is selected and go straight to step 3", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Personal project/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
|
||||
expect(screen.getByText(/When did you start using LiteLLM\?/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 5 total steps when using at a company", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /We use it for work/i }));
|
||||
|
||||
expect(screen.getByText(/Step 1 of 5/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 4 total steps when not using at a company", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Personal project/i }));
|
||||
|
||||
expect(screen.getByText(/Step 1 of 4/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should navigate back to step 1 from step 3 when No was previously selected", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Personal project/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Back/i }));
|
||||
|
||||
expect(
|
||||
screen.getByText(/Are you using LiteLLM at your company\?/i)
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
describe("when step 4 (reasons) is reached", () => {
|
||||
async function navigateToStep4(user: ReturnType<typeof userEvent.setup>) {
|
||||
// No path: step 1 → 3 → 4
|
||||
await user.click(screen.getByRole("button", { name: /Personal project/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
await user.click(screen.getByRole("radio", { name: /Less than a month ago/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
}
|
||||
|
||||
it("should show a text input when the Other reason is selected", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await navigateToStep4(user);
|
||||
await user.click(screen.getByRole("button", { name: /Something else not listed above/i }));
|
||||
|
||||
expect(screen.getByPlaceholderText(/Please specify/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should keep the Next button disabled when Other is selected but the text field is empty", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await navigateToStep4(user);
|
||||
await user.click(screen.getByRole("button", { name: /Something else not listed above/i }));
|
||||
|
||||
expect(screen.getByRole("button", { name: /Next/i })).toBeDisabled();
|
||||
});
|
||||
|
||||
it("should enable Next when a standard reason is selected", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
await navigateToStep4(user);
|
||||
await user.click(
|
||||
screen.getByRole("button", { name: /Stars, contributors, forks, community support/i })
|
||||
);
|
||||
|
||||
expect(screen.getByRole("button", { name: /Next/i })).not.toBeDisabled();
|
||||
});
|
||||
});
|
||||
|
||||
it("should call onComplete after successfully submitting the form", async () => {
|
||||
const onComplete = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={vi.fn()} onComplete={onComplete} />
|
||||
);
|
||||
|
||||
// Navigate through the No path: step 1 → 3 → 4 → 5 → submit
|
||||
await user.click(screen.getByRole("button", { name: /Personal project/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
await user.click(screen.getByRole("radio", { name: /Less than a month ago/i }));
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
await user.click(
|
||||
screen.getByRole("button", { name: /Stars, contributors, forks, community support/i })
|
||||
);
|
||||
await user.click(screen.getByRole("button", { name: /Next/i }));
|
||||
// Step 5: email is optional
|
||||
await user.click(screen.getByRole("button", { name: /Submit/i }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(onComplete).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
it("should call onClose when the close button is clicked", async () => {
|
||||
const onClose = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<SurveyModal isOpen={true} onClose={onClose} onComplete={vi.fn()} />
|
||||
);
|
||||
|
||||
// X close button is the first button in the modal header
|
||||
const buttons = screen.getAllByRole("button");
|
||||
await user.click(buttons[0]);
|
||||
|
||||
expect(onClose).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,82 @@
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { SurveyPrompt } from "./SurveyPrompt";
|
||||
|
||||
vi.mock("./NudgePrompt", () => ({
|
||||
NudgePrompt: ({
|
||||
title,
|
||||
description,
|
||||
buttonText,
|
||||
onOpen,
|
||||
onDismiss,
|
||||
isVisible,
|
||||
}: {
|
||||
title: string;
|
||||
description: string;
|
||||
buttonText: string;
|
||||
onOpen: () => void;
|
||||
onDismiss: () => void;
|
||||
isVisible: boolean;
|
||||
}) => {
|
||||
if (!isVisible) return null;
|
||||
return (
|
||||
<div>
|
||||
<span>{title}</span>
|
||||
<span>{description}</span>
|
||||
<button onClick={onOpen}>{buttonText}</button>
|
||||
<button onClick={onDismiss}>Dismiss</button>
|
||||
</div>
|
||||
);
|
||||
},
|
||||
}));
|
||||
|
||||
describe("SurveyPrompt", () => {
|
||||
it("should render with the Quick feedback title when visible", () => {
|
||||
renderWithProviders(
|
||||
<SurveyPrompt isVisible={true} onOpen={vi.fn()} onDismiss={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByText("Quick feedback")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the correct description text", () => {
|
||||
renderWithProviders(
|
||||
<SurveyPrompt isVisible={true} onOpen={vi.fn()} onDismiss={vi.fn()} />
|
||||
);
|
||||
expect(screen.getByText(/Help us improve LiteLLM/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onOpen when the share feedback button is clicked", async () => {
|
||||
const onOpen = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
||||
renderWithProviders(
|
||||
<SurveyPrompt isVisible={true} onOpen={onOpen} onDismiss={vi.fn()} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Share feedback/i }));
|
||||
|
||||
expect(onOpen).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should call onDismiss when the dismiss button is clicked", async () => {
|
||||
const onDismiss = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
||||
renderWithProviders(
|
||||
<SurveyPrompt isVisible={true} onOpen={vi.fn()} onDismiss={onDismiss} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Dismiss/i }));
|
||||
|
||||
expect(onDismiss).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not render when isVisible is false", () => {
|
||||
renderWithProviders(
|
||||
<SurveyPrompt isVisible={false} onOpen={vi.fn()} onDismiss={vi.fn()} />
|
||||
);
|
||||
expect(screen.queryByText("Quick feedback")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user