Merge pull request #24188 from BerriAI/main

merge main 0319
This commit is contained in:
Sameer Kankute 2026-03-20 11:03:54 +05:30 • committed by GitHub
commit 784f9431ad
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
62 changed files with 4844 additions and 1207 deletions

View file

@ -28,9 +28,12 @@ jobs:
find . -type d -name "__pycache__" -exec rm -rf {} + || true
find . -name "*.pyc" -delete || true
- name: Check poetry.lock is up to date
run: |
poetry check --lock || (echo "❌ poetry.lock is out of sync with pyproject.toml. Run 'poetry lock' locally and commit the result." && exit 1)
- name: Install dependencies
run: |
poetry lock
poetry install --with dev
- name: Check Black formatting

View file

@ -14,12 +14,12 @@ repos:
types: [python]
files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py
exclude: ^litellm/__init__.py$
# - id: black
# name: black
# entry: poetry run black
# language: system
# types: [python]
# files: (litellm/|litellm_proxy_extras/|enterprise/).*\.py
- id: black
name: black
entry: poetry run black
language: system
types: [python]
files: (litellm/|litellm_proxy_extras/).*\.py
- repo: https://github.com/pycqa/flake8
rev: 7.0.0 # The version of flake8 to use
hooks:

View file

@ -401,8 +401,10 @@ router_settings:
| AUTH_STRATEGY | Strategy used for authentication (e.g., OAuth, API key)
| AUTO_REDIRECT_UI_LOGIN_TO_SSO | Flag to enable automatic redirect of UI login page to SSO when SSO is configured. Default is **false**
| AUDIO_SPEECH_CHUNK_SIZE | Chunk size for audio speech processing. Default is 1024
| ANTHROPIC_API_KEY | API key for Anthropic service
| ANTHROPIC_API_KEY | API key for Anthropic service. Uses `x-api-key` header for authentication.
| ANTHROPIC_AUTH_TOKEN | Alternative auth token for Anthropic service. Uses `Authorization: Bearer` header instead of `x-api-key`. Used as fallback when `ANTHROPIC_API_KEY` is not set.
| ANTHROPIC_API_BASE | Base URL for Anthropic API. Default is https://api.anthropic.com
| ANTHROPIC_BASE_URL | Alternative to `ANTHROPIC_API_BASE` for setting the Anthropic API base URL. Used as fallback when `ANTHROPIC_API_BASE` is not set.
| ANTHROPIC_TOKEN_COUNTING_BETA_VERSION | Beta version header for Anthropic token counting API. Default is `token-counting-2024-11-01`
| AWS_ACCESS_KEY_ID | Access Key ID for AWS services
| AWS_BATCH_ROLE_ARN | ARN of the AWS IAM role for batch operations

View file

@ -602,6 +602,22 @@ Since you shouldn't use 12.5, round down to **10** to leave a safety buffer. Thi
- Total maximum connections: 8 workers × 10 connections = 80 connections
- This stays safely under your database's 100 connection limit
## LiteLLM License Key (Enterprise)
To enable [LiteLLM Enterprise features](https://docs.litellm.ai/docs/proxy/enterprise), set your license key as an environment variable:
```bash
export LITELLM_LICENSE="eyJ..."
```
The license key is a JWT token provided when you purchase a LiteLLM Enterprise license. Once set, LiteLLM will automatically detect and activate enterprise features.
You can also add it to your `.env` file:
```env
LITELLM_LICENSE="eyJ..."
```
## Extras

View file

@ -47,6 +47,10 @@ pip install litellm==1.82.3
- **FLUX Kontext image editing** — `flux-kontext-pro` and `flux-kontext-max` added to Black Forest Labs, alongside `flux-pro-1.0-fill` and `flux-pro-1.0-expand` for inpainting and outpainting
- **116 new models, 132 deprecated models cleaned up** — Major model map refresh including Mistral Magistral, Dashscope Qwen3 VL, xAI Grok via Azure AI, ZAI GLM-5, Serper Search; removal of OpenAI GPT-3.5/GPT-4 legacy variants, Gemini 1.5, and Vertex AI PaLM2
- **SageMaker Nova provider** — [New `sagemaker_nova` provider for Amazon Nova models on SageMaker](../../docs/providers/aws_sagemaker) - [PR #21542](https://github.com/BerriAI/litellm/pull/21542)
- **Hashicorp Vault secret manager** — Config override backend powered by Hashicorp Vault, with full UI for managing vault-sourced credentials - [PR #22939](https://github.com/BerriAI/litellm/pull/22939), [PR #23036](https://github.com/BerriAI/litellm/pull/23036)
- **Responses API WebSocket streaming** — Real-time WebSocket streaming for the Responses API, including support across all providers - [PR #22559](https://github.com/BerriAI/litellm/pull/22559), [PR #22771](https://github.com/BerriAI/litellm/pull/22771)
- **Org Admin RBAC expansion** — Org Admins can now access team management endpoints, view and invite internal users, and manage team membership without requiring a global admin role - [PR #23085](https://github.com/BerriAI/litellm/pull/23085), [PR #23080](https://github.com/BerriAI/litellm/pull/23080)
- **Guardrail mode defaults and tag-based modes** — Set a default guardrail mode list globally, and specify a list of modes in tag-based guardrail configs - [PR #22676](https://github.com/BerriAI/litellm/pull/22676), [PR #23020](https://github.com/BerriAI/litellm/pull/23020)
- **Secret redaction in logs** — API keys, tokens, and credentials automatically scrubbed from all proxy log output. Enabled by default; opt out with `LITELLM_DISABLE_REDACT_SECRETS=true` - [PR #23668](https://github.com/BerriAI/litellm/pull/23668)
- **Streaming stability fix** — Critical fix for `RuntimeError: Cannot send a request, as the client has been closed.` crashes after ~1 hour in production - [PR #22926](https://github.com/BerriAI/litellm/pull/22926)
@ -54,7 +58,7 @@ pip install litellm==1.82.3
## New Providers and Endpoints
### New Providers (5 new providers)
### New Providers (7 new providers)
| Provider | Supported LiteLLM Endpoints | Description |
| -------- | --------------------------- | ----------- |
@ -63,6 +67,8 @@ pip install litellm==1.82.3
| [Black Forest Labs](../../docs/providers/black_forest_labs) (`black_forest_labs/`) | `/images/generations`, `/images/edits` | FLUX image generation and editing — Kontext Pro/Max, Pro 1.0 Fill/Expand |
| [Serper](../../docs/providers/serper) (`serper/`) | `/search` | Web search via Serper API |
| [SageMaker Nova](../../docs/providers/aws_sagemaker) (`sagemaker_nova/`) | `/chat/completions` | Amazon Nova models via SageMaker endpoint |
| [Google Search API](../../docs/providers/google_search) (`google_search/`) | `/search` | Google Search API integration - [PR #22752](https://github.com/BerriAI/litellm/pull/22752) |
| [Bedrock Mantle](../../docs/providers/bedrock) (`bedrock_mantle/`) | `/chat/completions` | Amazon Bedrock via Mantle — alternative auth and routing path for Bedrock models - [PR #22866](https://github.com/BerriAI/litellm/pull/22866) |
---
@ -238,22 +244,92 @@ pip install litellm==1.82.3
- **[Responses API](../../docs/response_api)**
- Handle `response.failed`, `response.incomplete`, and `response.cancelled` terminal event types in background streaming — previously only `response.completed` was handled - [PR #23492](https://github.com/BerriAI/litellm/pull/23492)
- WebSocket streaming support for Responses API — real-time streaming via WebSocket for all providers - [PR #22559](https://github.com/BerriAI/litellm/pull/22559), [PR #22771](https://github.com/BerriAI/litellm/pull/22771)
- WebRTC support for real-time audio/video communication - [PR #23446](https://github.com/BerriAI/litellm/pull/23446)
- Responses API support for OpenAI-compatible JSON providers (`openai_like`) - [PR #21398](https://github.com/BerriAI/litellm/pull/21398)
- Route `gpt-5.4+` calls using both tools and reasoning to the Responses API automatically - [PR #23577](https://github.com/BerriAI/litellm/pull/23577)
#### Bug Fixes
- **[Anthropic Files API](../../docs/providers/anthropic)**
- Full Anthropic Files API support — upload, retrieve, list, and delete files; use file references in messages - [PR #16594](https://github.com/BerriAI/litellm/pull/16594)
- **[Mistral](../../docs/providers/mistral)**
- Voxtral audio transcription support — `mistral/voxtral-mini-*` and `mistral/voxtral-*` for audio transcription via Mistral - [PR #22801](https://github.com/BerriAI/litellm/pull/22801)
- **[OpenAI](../../docs/providers/openai)**
- `litellm.acount_tokens()` public API — async token counting with full OpenAI provider support - [PR #22809](https://github.com/BerriAI/litellm/pull/22809)
- Normalize `reasoning_effort` dict to string for chat completion API - [PR #22981](https://github.com/BerriAI/litellm/pull/22981)
- **[OpenRouter](../../docs/providers/openrouter)**
- Image edit support for OpenRouter models - [PR #22403](https://github.com/BerriAI/litellm/pull/22403)
- **[Google Vertex AI](../../docs/providers/vertex)**
- VIDEO modality token usage tracking in `completion_tokens_details` - [PR #22550](https://github.com/BerriAI/litellm/pull/22550)
- **Images API**
- `input_fidelity` parameter for image edit API - [PR #23201](https://github.com/BerriAI/litellm/pull/23201)
- **General**
- Per-request `enable_json_schema_validation` flag for thread-safe JSON schema validation - [PR #21233](https://github.com/BerriAI/litellm/pull/21233)
- Model cost aliases expansion — define aliases in the cost map that inherit pricing from a parent model - [PR #23314](https://github.com/BerriAI/litellm/pull/23314), [PR #23457](https://github.com/BerriAI/litellm/pull/23457)
- Wildcards model support for the Files API - [PR #22740](https://github.com/BerriAI/litellm/pull/22740)
#### Bugs
- **[Anthropic](../../docs/providers/anthropic)**
- Preserve native tool format (web_search, bash, tool_search, etc.) when guardrails convert tools for the Anthropic Messages API - [PR #23526](https://github.com/BerriAI/litellm/pull/23526)
- Enforce `type: "object"` on tool input schemas in `_map_tool_helper` — fixes tool call failures for strict-schema providers - [PR #23103](https://github.com/BerriAI/litellm/pull/23103)
- Deduplicate `tool_result` messages by `tool_call_id` — prevents duplicate tool result errors in multi-turn conversations - [PR #23104](https://github.com/BerriAI/litellm/pull/23104)
- Map `reasoning_effort` to `output_config` for Claude 4.6 models - [PR #22220](https://github.com/BerriAI/litellm/pull/22220)
- **[Google Gemini](../../docs/providers/gemini)**
- Correct streaming `finish_reason` for tool calls — was incorrectly returning `null` instead of `tool_calls` - [PR #21577](https://github.com/BerriAI/litellm/pull/21577)
- Preserve `$ref` in JSON Schema for Gemini 2.0+ — schema references were being stripped, breaking structured output - [PR #21597](https://github.com/BerriAI/litellm/pull/21597)
- Handle `minimal` `reasoning_effort` param for Gemini 3.1 models - [PR #22920](https://github.com/BerriAI/litellm/pull/22920)
- **[Google Vertex AI](../../docs/providers/vertex)**
- Pass through native Gemini `imageConfig` params for image generation - [PR #21585](https://github.com/BerriAI/litellm/pull/21585)
- Prevent content truncation when `finish_reason` races ahead of content chunks in streaming - [PR #22692](https://github.com/BerriAI/litellm/pull/22692)
- Strip LiteLLM-internal keys from `extra_body` before merging to Gemini request body - [PR #23131](https://github.com/BerriAI/litellm/pull/23131)
- Drop unsupported `output_config` parameter from all Vertex AI requests - [PR #22884](https://github.com/BerriAI/litellm/pull/22884)
- Skip schema transforms for Gemini 2.0+ tool parameters — avoids breaking native Gemini schema handling - [PR #23265](https://github.com/BerriAI/litellm/pull/23265)
- **[OpenRouter](../../docs/providers/openrouter)**
- Pattern-based fix for native model double-stripping when provider prefix matches model name - [PR #22320](https://github.com/BerriAI/litellm/pull/22320)
- Use provider-reported usage in streaming responses when `stream_options` is not set - [PR #21592](https://github.com/BerriAI/litellm/pull/21592)
- **[AWS Bedrock](../../docs/providers/bedrock)**
- Extract region and model ID from `bedrock/{region}/{model}` path format - [PR #22546](https://github.com/BerriAI/litellm/pull/22546)
- Strip `scope` from `cache_control` for Anthropic messages on Bedrock and Azure AI - [PR #22867](https://github.com/BerriAI/litellm/pull/22867)
- Populate `completion_tokens_details` in Responses API responses - [PR #23243](https://github.com/BerriAI/litellm/pull/23243)
- **[Azure AI](../../docs/providers/azure_ai)**
- Resolve `api_base` from environment variable in Document Intelligence OCR - [PR #21581](https://github.com/BerriAI/litellm/pull/21581)
- **[Moonshot / Kimi](../../docs/providers/openai_compatible)**
- Auto-fill `reasoning_content` for Moonshot Kimi reasoning models - [PR #23580](https://github.com/BerriAI/litellm/pull/23580)
- Preserve `image_url` blocks in multimodal messages for Moonshot - [PR #21595](https://github.com/BerriAI/litellm/pull/21595)
- **[HuggingFace](../../docs/providers/huggingface)**
- Forward `extra_headers` to HuggingFace embedding API - [PR #23525](https://github.com/BerriAI/litellm/pull/23525)
- **Token Counting / Cost**
- Fix `count_tokens` to include system prompts and tools in token counting API requests - [PR #22301](https://github.com/BerriAI/litellm/pull/22301)
- Pass all custom pricing fields to `register_model` in `completion()` and `embedding()` - [PR #22552](https://github.com/BerriAI/litellm/pull/22552)
- **Tools / Function Calling**
- Gracefully repair truncated JSON in tool call arguments — prevents crashes on malformed tool responses - [PR #22503](https://github.com/BerriAI/litellm/pull/22503)
- Fix `output_item.done` for function calls not emitting `finish_reason` in streaming - [PR #22553](https://github.com/BerriAI/litellm/pull/22553)
- Preserve thinking block order with multiple web searches - [PR #23093](https://github.com/BerriAI/litellm/pull/23093)
- **General**
- Normalize `content_filtered` finish reason across providers - [PR #23564](https://github.com/BerriAI/litellm/pull/23564)
- Unify `finish_reason` mapping to OpenAI-compatible values across all providers - [PR #22138](https://github.com/BerriAI/litellm/pull/22138)
- Fix custom cost tracking on deployments for `/v1/messages` and `/v1/responses` - [PR #23647](https://github.com/BerriAI/litellm/pull/23647)
- Fix per-request custom pricing when `router_model_id` has no pricing data — now falls back to model name
- Fix batch list showing stale `validating` status after completion - [PR #22982](https://github.com/BerriAI/litellm/pull/22982)
- Fix batch retrieve returning raw `output_file_id` when `model_id` is missing - [PR #23194](https://github.com/BerriAI/litellm/pull/23194)
- Encode batch IDs when `x-litellm-model` header is used - [PR #22653](https://github.com/BerriAI/litellm/pull/22653)
- Map `reasoning` to `reasoning_content` in streaming Delta for gpt-oss providers - [PR #22803](https://github.com/BerriAI/litellm/pull/22803)
---
@ -264,10 +340,31 @@ pip install litellm==1.82.3
- **Virtual Keys**
- Add Organization dropdown to Create/Edit Key form — `organization_id` is now a first-class field in Key Ownership - [PR #23595](https://github.com/BerriAI/litellm/pull/23595)
- Allow setting `organization_id` on `/key/update` — keys can be assigned or moved to a different organization after creation - [PR #23557](https://github.com/BerriAI/litellm/pull/23557)
- Manual Spend Reset for virtual keys from the UI — admins can reset key spend to zero on demand - [PR #22715](https://github.com/BerriAI/litellm/pull/22715)
- BYOK (Bring Your Own Key) — client-side provider API key takes precedence over proxy key for Anthropic `/v1/messages` - [PR #22964](https://github.com/BerriAI/litellm/pull/22964)
- UI login session duration configurable via `LITELLM_UI_SESSION_DURATION` environment variable - [PR #22182](https://github.com/BerriAI/litellm/pull/22182)
- Auto-redirect UI login to SSO via `auto_redirect_ui_login_to_sso: true` in config.yaml - [PR #23367](https://github.com/BerriAI/litellm/pull/23367)
- **Access Control (RBAC)**
- Org Admins can now access team management endpoints — `/team/new`, `/team/update`, `/team/delete`, `/team/member_add`, `/team/member_delete` - [PR #23085](https://github.com/BerriAI/litellm/pull/23085), [PR #23095](https://github.com/BerriAI/litellm/pull/23095)
- Org Admins can view and invite internal users — full user management without requiring global admin role - [PR #23080](https://github.com/BerriAI/litellm/pull/23080)
- Allow Admin Viewers to access Audit Logs — view-only admin role now includes audit log access - [PR #23419](https://github.com/BerriAI/litellm/pull/23419)
- RBAC for Vector Stores and Agents — key/team-level access control for vector store and agent resources - [PR #22858](https://github.com/BerriAI/litellm/pull/22858)
- User filter scope (`scope_user_search_to_org`) is now opt-in — previously default-on, causing unintended restriction - [PR #23057](https://github.com/BerriAI/litellm/pull/23057)
- **Vector Stores**
- Vector Store management endpoints — retrieve, list, update, and delete vector stores via `/v1/vector_stores/*` - [PR #23435](https://github.com/BerriAI/litellm/pull/23435)
- **Teams**
- Batch expiry setting for teams — configure a default expiry duration for all team keys - [PR #22705](https://github.com/BerriAI/litellm/pull/22705)
- Team Admin can reset key spend - [PR #22725](https://github.com/BerriAI/litellm/pull/22725)
- **Internal Users**
- Add/Remove Team Membership directly from the Internal Users info page — includes searchable dropdown and role selector; no longer requires navigating to each team - [PR #23638](https://github.com/BerriAI/litellm/pull/23638)
- **Models**
- Attach knowledge base to model via UI - [PR #22656](https://github.com/BerriAI/litellm/pull/22656)
- **Default Team Settings**
- Modernize page to antd (consistent with rest of app) - [PR #23614](https://github.com/BerriAI/litellm/pull/23614)
- Fix: default team params (budget, duration, tpm, rpm, permissions) now correctly applied on `/team/new` - [PR #23614](https://github.com/BerriAI/litellm/pull/23614)
@ -290,6 +387,13 @@ pip install litellm==1.82.3
- Fix Public Model Hub not showing config-defined models after save - [PR #23501](https://github.com/BerriAI/litellm/pull/23501)
- Fix fallback popup model dropdown z-index issue - [PR #23516](https://github.com/BerriAI/litellm/pull/23516)
- Fix double-counting bug in org/team key limit checks on `/key/update`
- Fix invite link allowing multiple password resets for the same link - [PR #22462](https://github.com/BerriAI/litellm/pull/22462)
- Fix key expiry default duration not being applied when `duration` is not set - [PR #22956](https://github.com/BerriAI/litellm/pull/22956)
- Fix all proxy models not including model access groups in key creation - [PR #23236](https://github.com/BerriAI/litellm/pull/23236)
- Fix admin viewers unable to see all organizations - [PR #22940](https://github.com/BerriAI/litellm/pull/22940)
- Fix Audit Logs UI: added server-side pagination, filtering, and drawer view - [PR #22476](https://github.com/BerriAI/litellm/pull/22476)
- Fix virtual keys in teams view not applying the team filter correctly - [PR #23065](https://github.com/BerriAI/litellm/pull/23065)
- Fix team expiry enforcement validation - [PR #22728](https://github.com/BerriAI/litellm/pull/22728)
---
@ -297,6 +401,13 @@ pip install litellm==1.82.3
### Logging
- **[Helicone](../../docs/observability/helicone_integration)**
- Add Gemini and Vertex AI support to HeliconeLogger — routes Gemini and Vertex AI requests through the correct Helicone provider URL - [PR #19288](https://github.com/BerriAI/litellm/pull/19288)
- Fix correct provider URL for Vertex AI Gemini models - [PR #22603](https://github.com/BerriAI/litellm/pull/22603)
- **[Langfuse](../../docs/proxy/logging#langfuse)**
- Fix failure path kwargs inconsistency causing dropped traces on failed requests - [PR #22390](https://github.com/BerriAI/litellm/pull/22390)
- **[Vantage](https://vantage.sh)**
- Add Vantage integration for FOCUS 1.2 CSV export — export LiteLLM proxy spend data as FinOps Open Cost & Usage Specification reports, with time-windowed filenames to prevent overwrites - [PR #23333](https://github.com/BerriAI/litellm/pull/23333)
@ -305,7 +416,10 @@ pip install litellm==1.82.3
### Guardrails
No major guardrail changes in this release.
- **Guardrail mode default list** — Configure a default list of guardrail modes applied globally when no per-request mode is specified - [PR #22676](https://github.com/BerriAI/litellm/pull/22676)
- **Tag-based guardrail mode lists** — Specify a list of modes in tag-based guardrail configs instead of a single mode - [PR #23020](https://github.com/BerriAI/litellm/pull/23020)
- **Fix presidio PII token leak** — Edge case where Anthropic handle in Presidio caused PII data exposure in token response - [PR #22627](https://github.com/BerriAI/litellm/pull/22627)
- **Fix OTEL orphaned guardrail traces** — Span redundancy and missing response IDs in OpenTelemetry guardrail traces - [PR #23001](https://github.com/BerriAI/litellm/pull/23001)
### Prompt Management
@ -313,7 +427,32 @@ No major prompt management changes in this release.
### Secret Managers
No major secret manager changes in this release.
- **[Hashicorp Vault](../../docs/secret_managers)** — Full Hashicorp Vault integration as a config override backend — secrets defined in Vault are fetched at startup and override `config.yaml` values. UI support for managing vault-sourced credentials included - [PR #22939](https://github.com/BerriAI/litellm/pull/22939), [PR #23036](https://github.com/BerriAI/litellm/pull/23036)
---
## MCP Gateway
#### Features
- **Token authentication for MCP servers** — configure `auth_type: "bearer"` per MCP server to require token-based auth on tool calls - [PR #23260](https://github.com/BerriAI/litellm/pull/23260)
- **Team-scoped MCP server filtering** — keys created under a team only see MCP servers available to that team - [PR #23323](https://github.com/BerriAI/litellm/pull/23323)
- **Per-server health recheck in UI** — trigger a health check for individual MCP servers without reloading all servers - [PR #23328](https://github.com/BerriAI/litellm/pull/23328)
#### Bugs
- Fix MCP server URL and tools management issues causing tool discovery to fail - [PR #22751](https://github.com/BerriAI/litellm/pull/22751)
- Fix MCP server health checks triggering on server deletion - [PR #23063](https://github.com/BerriAI/litellm/pull/23063)
---
## Spend Tracking, Budgets and Rate Limiting
- **Fix budget-linked keys never having spend reset** — Keys linked to budget objects were not having their spend reset on the configured reset interval - [PR #20688](https://github.com/BerriAI/litellm/pull/20688)
- **Flex pricing support** — Add `flex_pricing` field to cost map for providers that offer dynamic pricing tiers - [PR #22992](https://github.com/BerriAI/litellm/pull/22992)
- **Fix spend log cleanup** — Resolved lock tracking, integer retention, and skip-log-level issues in spend log cleanup job - [PR #22687](https://github.com/BerriAI/litellm/pull/22687)
- **Fix WebSearch spend log deduplication** — WebSearch interception was failing with thinking enabled; fixed along with spend log dedup - [PR #22679](https://github.com/BerriAI/litellm/pull/22679)
- **Fix TypeError when request has no API key** — Spend tracking was throwing unhandled exception when API key was absent from request - [PR #23363](https://github.com/BerriAI/litellm/pull/23363)
---
@ -323,6 +462,10 @@ No major secret manager changes in this release.
- **Fix OOM / Prisma connection loss** on large installs — unbounded managed-object poll was exhausting Prisma connections after ~60–70 minutes on instances with 336K+ queued response rows - [PR #23472](https://github.com/BerriAI/litellm/pull/23472)
- **Centralize logging kwarg updates** — root cause fix migrating all logging updates to a single function, eliminating kwarg inconsistencies across logging paths - [PR #23659](https://github.com/BerriAI/litellm/pull/23659)
- **Fix tiktoken cache for non-root offline containers** — tiktoken cache now works correctly in offline environments running as non-root users - [PR #23498](https://github.com/BerriAI/litellm/pull/23498)
- **Block proxy startup when Redis transaction buffer has no Redis** — prevents silent data loss when `use_redis_transaction_buffer: true` is set without a Redis connection - [PR #23019](https://github.com/BerriAI/litellm/pull/23019)
- **Fix `InFlightRequestsMiddleware` crash** — undefined kwargs in middleware were causing request failures - [PR #22523](https://github.com/BerriAI/litellm/pull/22523)
- **Fix `BaseModelResponseIterator` crash on non-string stream chunks** — streaming was crashing when providers returned non-string chunk data - [PR #23497](https://github.com/BerriAI/litellm/pull/23497)
- **Fix `SERVER_ROOT_PATH` prefix handling** — strip prefix before checking mapped pass-through routes to prevent double-prefix issues - [PR #23414](https://github.com/BerriAI/litellm/pull/23414)
- **Add CodSpeed continuous performance benchmarks** — automated performance regression tracking on CI - [PR #23676](https://github.com/BerriAI/litellm/pull/23676)
---
@ -342,6 +485,16 @@ No major secret manager changes in this release.
---
## Documentation Updates
- Add Anthropic `/v1/messages` → `/responses` parameter mapping reference - [PR #22893](https://github.com/BerriAI/litellm/pull/22893)
- Update Okta SSO docs and custom SSO handler example - [PR #22786](https://github.com/BerriAI/litellm/pull/22786)
- Add `LITELLM_MAX_BUDGET_PER_SESSION_TTL` to environment variables reference - [PR #23186](https://github.com/BerriAI/litellm/pull/23186)
- Add DB query performance guidelines to `CLAUDE.md` - [PR #23196](https://github.com/BerriAI/litellm/pull/23196)
- Add Gemini Vertex AI PayGo/priority cost tracking docs - [PR #22948](https://github.com/BerriAI/litellm/pull/22948)
---
## New Contributors
* @ryanh-ai made their first contribution in [PR #21542](https://github.com/BerriAI/litellm/pull/21542)
@ -359,14 +512,17 @@ No major secret manager changes in this release.
## Diff Summary
## 03/16/2026
* New Providers: 5
* New Providers: 7
* New Models / Updated Models: 116 new, 132 removed
* LLM API Endpoints: 5
* Management Endpoints / UI: 11
* AI Integrations: 2
* Performance / Reliability: 5
* LLM API Endpoints: 37
* Management Endpoints / UI: 31
* AI Integrations: 8
* MCP Gateway: 5
* Spend Tracking, Budgets and Rate Limiting: 5
* Performance / Loadbalancing / Reliability improvements: 9
* Security: 3
* Database / Proxy Operations: 2
* Documentation Updates: 5
---

View file

@ -524,6 +524,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
optional_params.api_base
or litellm.api_base
or get_secret_str("ANTHROPIC_API_BASE")
or get_secret_str("ANTHROPIC_BASE_URL")
)
api_key = (
optional_params.api_key

View file

@ -1444,6 +1444,7 @@ SENTRY_DENYLIST = [
"credential",
"OPENAI_API_KEY",
"ANTHROPIC_API_KEY",
"ANTHROPIC_AUTH_TOKEN",
"AZURE_API_KEY",
"COHERE_API_KEY",
"REPLICATE_API_KEY",

View file

@ -442,7 +442,9 @@ class DataDogLogger(
verbose_logger.debug("Datadog: Logger - Logging payload = %s", json_payload)
dd_payload = DatadogPayload(
ddsource=get_datadog_source(),
ddtags=",".join(get_datadog_tags(standard_logging_object=standard_logging_object)),
ddtags=",".join(
get_datadog_tags(standard_logging_object=standard_logging_object)
),
hostname=get_datadog_hostname(),
message=json_payload,
service=get_datadog_service(),

View file

@ -5,7 +5,6 @@ import os
import random
import traceback
import types
from litellm._uuid import uuid
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional
@ -14,10 +13,11 @@ from pydantic import BaseModel # type: ignore
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.langsmith_mock_client import (
should_use_langsmith_mock,
create_mock_langsmith_client,
should_use_langsmith_mock,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -110,6 +110,60 @@ class LangsmithLogger(CustomBatchLogger):
LANGSMITH_TENANT_ID=_credentials_tenant_id,
)
def _extract_metadata_fields(
self, metadata: dict, credentials: LangsmithCredentialsObject
):
return {
"project_name": metadata.get(
"project_name", credentials["LANGSMITH_PROJECT"]
),
"run_name": metadata.get("run_name", self.langsmith_default_run_name),
"run_id": metadata.get("id", metadata.get("run_id", None)),
"parent_run_id": metadata.get("parent_run_id", None),
"trace_id": metadata.get("trace_id", None),
"session_id": metadata.get("session_id", None),
"dotted_order": metadata.get("dotted_order", None),
}
def _build_extra_metadata(self, metadata: Dict):
extra_metadata = dict(metadata)
requester_metadata = extra_metadata.get("requester_metadata")
if requester_metadata and isinstance(requester_metadata, dict):
for key in ("session_id", "thread_id", "conversation_id"):
if key in requester_metadata and key not in extra_metadata:
extra_metadata[key] = requester_metadata[key]
return extra_metadata
def _build_outputs_with_usage(
self, payload: StandardLoggingPayload
) -> Dict[str, Any]:
response = payload["response"]
outputs: Dict[str, Any]
if isinstance(response, dict):
outputs = {**response}
else:
outputs = {"output": response}
outputs["usage_metadata"] = {
"input_tokens": payload.get("prompt_tokens", 0),
"output_tokens": payload.get("completion_tokens", 0),
"total_tokens": payload.get("total_tokens", 0),
"total_cost": payload.get("response_cost", 0),
}
return outputs
def _ensure_required_ids(self, data: dict, run_id: Optional[str]):
if "id" not in data or data["id"] is None:
run_id = str(uuid.uuid4())
data["id"] = run_id
if "trace_id" not in data or data["trace_id"] is None:
if run_id is not None and isinstance(run_id, str):
data["trace_id"] = run_id
if "dotted_order" not in data or data["dotted_order"] is None:
if run_id is not None and isinstance(run_id, str):
data["dotted_order"] = self.make_dot_order(run_id=run_id)
def _prepare_log_data(
self,
kwargs,
@ -121,44 +175,28 @@ class LangsmithLogger(CustomBatchLogger):
try:
_litellm_params = kwargs.get("litellm_params", {}) or {}
metadata = _litellm_params.get("metadata", {}) or {}
project_name = metadata.get(
"project_name", credentials["LANGSMITH_PROJECT"]
)
run_name = metadata.get("run_name", self.langsmith_default_run_name)
run_id = metadata.get("id", metadata.get("run_id", None))
parent_run_id = metadata.get("parent_run_id", None)
trace_id = metadata.get("trace_id", None)
session_id = metadata.get("session_id", None)
dotted_order = metadata.get("dotted_order", None)
fields = self._extract_metadata_fields(metadata, credentials)
verbose_logger.debug(
f"Langsmith Logging - project_name: {project_name}, run_name {run_name}"
f"Langsmith Logging - project_name: {fields['project_name']}, run_name {fields['run_name']}"
)
# Ensure everything in the payload is converted to str
payload: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
)
if payload is None:
raise Exception("Error logging request payload. Payload=none.")
metadata = payload[
"metadata"
] # ensure logged metadata is json serializable
extra_metadata = dict(metadata)
requester_metadata = extra_metadata.get("requester_metadata")
if requester_metadata and isinstance(requester_metadata, dict):
for key in ("session_id", "thread_id", "conversation_id"):
if key in requester_metadata and key not in extra_metadata:
extra_metadata[key] = requester_metadata[key]
metadata = payload["metadata"]
extra_metadata = self._build_extra_metadata(dict(metadata))
outputs = self._build_outputs_with_usage(payload)
data = {
"name": run_name,
"run_type": "llm", # this should always be llm, since litellm always logs llm calls. Langsmith allow us to log "chain"
"name": fields["run_name"],
"run_type": "llm",
"inputs": payload,
"outputs": payload["response"],
"session_name": project_name,
"outputs": outputs,
"session_name": fields["project_name"],
"start_time": payload["startTime"],
"end_time": payload["endTime"],
"tags": payload["request_tags"],
@ -168,46 +206,19 @@ class LangsmithLogger(CustomBatchLogger):
if payload["error_str"] is not None and payload["status"] == "failure":
data["error"] = payload["error_str"]
if run_id:
data["id"] = run_id
if parent_run_id:
data["parent_run_id"] = parent_run_id
if trace_id:
data["trace_id"] = trace_id
if session_id:
data["session_id"] = session_id
if dotted_order:
data["dotted_order"] = dotted_order
run_id: Optional[str] = data.get("id") # type: ignore
if "id" not in data or data["id"] is None:
"""
for /batch langsmith requires id, trace_id and dotted_order passed as params
"""
run_id = str(uuid.uuid4())
data["id"] = run_id
if (
"trace_id" not in data
or data["trace_id"] is None
and (run_id is not None and isinstance(run_id, str))
for key in (
"id",
"parent_run_id",
"trace_id",
"session_id",
"dotted_order",
):
data["trace_id"] = run_id
if (
"dotted_order" not in data
or data["dotted_order"] is None
and (run_id is not None and isinstance(run_id, str))
):
data["dotted_order"] = self.make_dot_order(run_id=run_id) # type: ignore
field_key = "run_id" if key == "id" else key
if fields[field_key]:
data[key] = fields[field_key]
self._ensure_required_ids(data, fields["run_id"])
verbose_logger.debug("Langsmith Logging data on langsmith: %s", data)
return data
except Exception:
raise

View file

@ -84,6 +84,8 @@ from litellm.types.llms.openai import (
OpenAIModerationResponse,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponseFailedEvent,
ResponseIncompleteEvent,
ResponsesAPIResponse,
)
from litellm.types.mcp import MCPPostCallResponseObject
@ -516,6 +518,23 @@ class Logging(LiteLLMLoggingBaseClass):
),
)
def get_router_model_id(self) -> Optional[str]:
"""Extract the router deployment model_id from litellm_params.
Checks both litellm_metadata and metadata for model_info.id.
Used by cost calculators to look up custom pricing registered
under the deployment's model_info.id in litellm.model_cost.
"""
if not hasattr(self, "litellm_params"):
return None
for key in ("litellm_metadata", "metadata"):
meta = self.litellm_params.get(key, {}) or {}
info = meta.get("model_info", {}) or {}
model_id = info.get("id")
if model_id is not None:
return model_id
return None
def update_environment_variables(
self,
litellm_params: Dict,
@ -1458,16 +1477,8 @@ class Logging(LiteLLMLoggingBaseClass):
# Fallback: extract router_model_id from litellm_params when not available
# from the result object. ResponsesAPIResponse objects (used by /v1/responses
# streaming) don't carry _hidden_params["model_id"] like ModelResponse does.
if router_model_id is None and hasattr(self, "litellm_params"):
for metadata_key in ("litellm_metadata", "metadata"):
_metadata: dict = (
self.litellm_params.get(metadata_key, {}) or {}
)
_model_info: dict = _metadata.get("model_info", {}) or {}
_model_id = _model_info.get("id")
if _model_id is not None:
router_model_id = _model_id
break
if router_model_id is None:
router_model_id = self.get_router_model_id()
## RESPONSE COST ##
custom_pricing = use_custom_pricing_for_model(
@ -2972,8 +2983,7 @@ class Logging(LiteLLMLoggingBaseClass):
if (
isinstance(callback, CustomLogger)
and is_sync_request
and self.call_type
!= CallTypes.pass_through.value
and self.call_type != CallTypes.pass_through.value
): # custom logger class
callback.log_failure_event(
start_time=start_time,
@ -3321,7 +3331,10 @@ class Logging(LiteLLMLoggingBaseClass):
return result
elif isinstance(result, TextCompletionResponse):
return result
elif isinstance(result, ResponseCompletedEvent):
elif isinstance(
result,
(ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent),
):
## return unified Usage object
if isinstance(result.response.usage, ResponseAPIUsage):
transformed_usage = (
@ -3342,7 +3355,6 @@ class Logging(LiteLLMLoggingBaseClass):
return result.response
else:
return None
return None
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
"""

View file

@ -160,6 +160,7 @@ class CustomStreamWrapper:
self.chunks: List = (
[]
) # keep track of the returned chunks - used for calculating the input/output tokens for stream options
self._repeated_messages_count = 1
self.is_function_call = self.check_is_function_call(logging_obj=logging_obj)
self.created: Optional[int] = None
self._last_returned_hidden_params: Optional[dict] = None
@ -241,7 +242,7 @@ class CustomStreamWrapper:
except Exception as e:
raise e
def safety_checker(self) -> None:
def raise_on_model_repetition(self) -> None:
"""
Fixes - https://github.com/BerriAI/litellm/issues/5158
@ -249,28 +250,35 @@ class CustomStreamWrapper:
Raises - InternalServerError, if LLM enters infinite loop while streaming
"""
if len(self.chunks) >= litellm.REPEATED_STREAMING_CHUNK_LIMIT:
# Get the last n chunks
last_chunks = self.chunks[-litellm.REPEATED_STREAMING_CHUNK_LIMIT :]
if len(self.chunks) < 2:
return
# Extract the relevant content from the chunks
last_contents = [chunk.choices[0].delta.content for chunk in last_chunks]
last_content = self.chunks[-1].choices[0].delta.content
# Check if all extracted contents are identical
if all(content == last_contents[0] for content in last_contents):
if (
last_contents[0] is not None
and isinstance(last_contents[0], str)
and len(last_contents[0]) > 2
): # ignore empty content - https://github.com/BerriAI/litellm/issues/5158#issuecomment-2287156946
# All last n chunks are identical
raise litellm.InternalServerError(
message="The model is repeating the same chunk = {}.".format(
last_contents[0]
),
model="",
llm_provider="",
)
if (
last_content is None
or not isinstance(last_content, str)
or len(last_content) <= 2
): # ignore empty content - https://github.com/BerriAI/litellm/issues/5158#issuecomment-2287156946
self._repeated_messages_count = 1
return
second_to_last_content = self.chunks[-2].choices[0].delta.content
if last_content == second_to_last_content:
self._repeated_messages_count += 1
else:
self._repeated_messages_count = 1
if self._repeated_messages_count >= litellm.REPEATED_STREAMING_CHUNK_LIMIT:
# All last n chunks are identical
raise litellm.InternalServerError(
message="The model is repeating the same chunk = {}.".format(
last_content
),
model="",
llm_provider="",
)
def check_special_tokens(self, chunk: str, finish_reason: Optional[str]):
"""
@ -924,7 +932,7 @@ class CustomStreamWrapper:
if (
is_chunk_non_empty
): # cannot set content of an OpenAI Object to be an empty string
self.safety_checker()
self.raise_on_model_repetition()
hold, model_response_str = self.check_special_tokens(
chunk=completion_obj["content"],
finish_reason=model_response.choices[0].finish_reason,
@ -1894,12 +1902,8 @@ class CustomStreamWrapper:
getattr(complete_streaming_response, "usage"),
)
try:
_cache_copy = complete_streaming_response.model_copy(
deep=True
)
_log_copy = complete_streaming_response.model_copy(
deep=True
)
_cache_copy = complete_streaming_response.model_copy(deep=True)
_log_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_cache_copy = complete_streaming_response.model_copy()
_log_copy = complete_streaming_response.model_copy()
@ -2122,9 +2126,7 @@ class CustomStreamWrapper:
getattr(complete_streaming_response, "usage"),
)
try:
_copy = complete_streaming_response.model_copy(
deep=True
)
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_copy = complete_streaming_response.model_copy()
asyncio.create_task(

View file

@ -42,9 +42,8 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
api_base: Optional[str] = None,
) -> dict:
"""Validate and prepare environment-specific headers and parameters."""
# Resolve api_key from environment if not provided
api_key = api_key or self.anthropic_model_info.get_api_key()
if api_key is None:
auth_header = self.anthropic_model_info.get_auth_header(api_key)
if auth_header is None:
raise ValueError(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params"
)
@ -52,8 +51,8 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
"accept": "application/json",
"anthropic-version": "2023-06-01",
"content-type": "application/json",
"x-api-key": api_key,
}
_headers.update(auth_header)
# Add beta header for message batches
if "anthropic-beta" not in headers:
headers["anthropic-beta"] = "message-batches-2024-09-24"

View file

@ -48,6 +48,10 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
)
from litellm.types.responses.main import (
OutputCodeInterpreterCall,
build_code_interpreter_log_outputs,
)
from litellm.types.utils import (
Delta,
GenericStreamingChunk,
@ -538,6 +542,12 @@ class ModelResponseIterator:
# Accumulate compaction blocks for multi-turn reconstruction
self.compaction_blocks: List[Dict[str, Any]] = []
# Track server tool use inputs and results for code_interpreter_results
self._server_tool_inputs: Dict[str, Any] = {}
self.tool_results: List[Dict[str, Any]] = []
self._current_server_tool_id: Optional[str] = None
self._container_id: Optional[str] = None
def check_empty_tool_call_args(self) -> bool:
"""
Check if the tool call block so far has been an empty string
@ -682,6 +692,39 @@ class ModelResponseIterator:
return content_block_start
def _build_code_interpreter_results(self) -> list:
"""Convert accumulated tool_results to OutputCodeInterpreterCall objects.
Called during streaming to produce provider-neutral code_interpreter_results
alongside the raw tool_results, so the Responses API layer doesn't need
Anthropic-specific knowledge.
Returns the full cumulative list each time (not incremental), matching
how web_search_results works. stream_chunk_builder uses "last value
wins" for list-valued provider_specific_fields keys, so the last
emission must contain every result.
"""
results = []
for tr in self.tool_results:
if tr.get("type") != "bash_code_execution_tool_result":
continue
call_id = tr.get("tool_use_id", "")
content = tr.get("content", {})
log_outputs = build_code_interpreter_log_outputs(content)
tool_input = self._server_tool_inputs.get(call_id, {})
code = tool_input.get("command", "") if isinstance(tool_input, dict) else ""
results.append(
OutputCodeInterpreterCall(
type="code_interpreter_call",
id=call_id,
code=code,
container_id=self._container_id,
status="completed",
outputs=log_outputs,
)
)
return results
def chunk_parser(self, chunk: dict) -> ModelResponseStream: # noqa: PLR0915
try:
type_chunk = chunk.get("type", "") or ""
@ -748,6 +791,23 @@ class ModelResponseIterator:
),
index=self.tool_index,
)
# Track server tool use inputs for code_interpreter_results.
# The initial input in content_block_start is typically {}
# for streaming; the full input arrives via input_json_delta
# and is assembled at content_block_stop.
if (
content_block_start["content_block"]["type"]
== "server_tool_use"
):
self._current_server_tool_id = content_block_start[
"content_block"
]["id"]
tool_input = content_block_start["content_block"].get(
"input", {}
)
self._server_tool_inputs[
self._current_server_tool_id
] = tool_input
# Include caller information if present (for programmatic tool calling)
if "caller" in content_block_start["content_block"]:
caller_data = content_block_start["content_block"]["caller"]
@ -808,10 +868,12 @@ class ModelResponseIterator:
elif content_type != "tool_search_tool_result":
# Handle other tool results (code execution, etc.)
# Skip tool_search_tool_result as it's internal metadata
if not hasattr(self, "tool_results"):
self.tool_results = []
self.tool_results.append(content_block_start["content_block"])
provider_specific_fields["tool_results"] = self.tool_results
# Convert to provider-neutral code_interpreter_results
provider_specific_fields[
"code_interpreter_results"
] = self._build_code_interpreter_results()
elif type_chunk == "content_block_stop":
ContentBlockStop(**chunk) # type: ignore
@ -828,6 +890,26 @@ class ModelResponseIterator:
),
index=self.tool_index,
)
# Update server_tool_inputs with fully assembled input
# from input_json_delta chunks (content_block_start has {})
if (
self.current_content_block_type == "server_tool_use"
and self._current_server_tool_id
):
args = ""
for block in self.content_blocks:
if block["delta"]["type"] == "input_json_delta":
partial_json = block["delta"].get("partial_json")
if isinstance(partial_json, str):
args += partial_json
if args:
try:
self._server_tool_inputs[
self._current_server_tool_id
] = json.loads(args)
except (json.JSONDecodeError, TypeError):
pass
self._current_server_tool_id = None
# Reset response_format tool tracking when block stops
self.is_response_format_tool = False
# Reset current content block type
@ -840,6 +922,17 @@ class ModelResponseIterator:
finish_reason, usage, container = self._handle_message_delta(chunk)
if container:
provider_specific_fields["container"] = container
# Store container_id and re-emit code_interpreter_results
# so stream_chunk_builder's last-value-wins picks up the
# version with container_id populated.
container_id = (
container.get("id") if isinstance(container, dict) else None
)
if container_id and self.tool_results:
self._container_id = container_id
provider_specific_fields[
"code_interpreter_results"
] = self._build_code_interpreter_results()
elif type_chunk == "message_start":
"""
Anthropic

View file

@ -50,6 +50,10 @@ from litellm.types.llms.openai import (
OpenAIMcpServerTool,
OpenAIWebSearchOptions,
)
from litellm.types.responses.main import (
OutputCodeInterpreterCall,
build_code_interpreter_log_outputs,
)
from litellm.types.utils import (
CacheCreationTokenDetails,
CompletionTokensDetailsWrapper,
@ -1522,7 +1526,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
tool_results = []
tool_results.append(content)
elif content.get("thinking", None) is not None:
elif content.get("type") == "thinking":
if thinking_blocks is None:
thinking_blocks = []
thinking_blocks.append(cast(ChatCompletionThinkingBlock, content))
@ -1682,6 +1686,96 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
return usage
def _build_code_by_id_map(
self, tool_calls: List[ChatCompletionToolCallChunk]
) -> Dict[str, str]:
code_by_id: Dict[str, str] = {}
for tc in tool_calls:
try:
args = json.loads(tc.get("function", {}).get("arguments", "{}"))
call_id = tc.get("id")
command = args.get("command", "")
if isinstance(call_id, str):
code_by_id[call_id] = command if isinstance(command, str) else ""
except Exception:
pass
return code_by_id
def _build_code_interpreter_results(
self,
tool_results: List[Any],
code_by_id: Dict[str, str],
container_id: Optional[str],
) -> List[OutputCodeInterpreterCall]:
code_interpreter_results = []
for tr in tool_results:
if tr.get("type") != "bash_code_execution_tool_result":
continue
call_id = tr.get("tool_use_id", "")
content = tr.get("content", {})
log_outputs = build_code_interpreter_log_outputs(content)
code_interpreter_results.append(
OutputCodeInterpreterCall(
type="code_interpreter_call",
id=call_id,
code=code_by_id.get(call_id, ""),
container_id=container_id,
status="completed",
outputs=log_outputs,
)
)
return code_interpreter_results
def _build_provider_specific_fields(
self,
completion_response: dict,
citations: Optional[List[Any]],
thinking_blocks: Optional[
List[
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
]
],
web_search_results: Optional[List[Any]],
tool_results: Optional[List[Any]],
compaction_blocks: Optional[List[Any]],
tool_calls: List[ChatCompletionToolCallChunk],
) -> Dict[str, Any]:
provider_specific_fields: Dict[str, Any] = {
"citations": citations,
"thinking_blocks": thinking_blocks,
}
context_management = completion_response.get("context_management")
if context_management is not None:
provider_specific_fields["context_management"] = context_management
if web_search_results is not None:
provider_specific_fields["web_search_results"] = web_search_results
if tool_results is not None:
provider_specific_fields["tool_results"] = tool_results
container_id = (
completion_response.get("container", {}).get("id")
if isinstance(completion_response.get("container"), dict)
else None
)
code_by_id = self._build_code_by_id_map(tool_calls)
code_interpreter_results = self._build_code_interpreter_results(
tool_results, code_by_id, container_id
)
provider_specific_fields[
"code_interpreter_results"
] = code_interpreter_results
container = completion_response.get("container")
if container is not None:
provider_specific_fields["container"] = container
if compaction_blocks is not None:
provider_specific_fields["compaction_blocks"] = compaction_blocks
return provider_specific_fields
def transform_parsed_response(
self,
completion_response: dict,
@ -1702,98 +1796,73 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
status_code=raw_response.status_code,
headers=response_headers,
)
else:
text_content = ""
citations: Optional[List[Any]] = None
thinking_blocks: Optional[
List[
Union[
ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock
]
]
] = None
reasoning_content: Optional[str] = None
tool_calls: List[ChatCompletionToolCallChunk] = []
(
text_content,
citations,
thinking_blocks,
reasoning_content,
tool_calls,
web_search_results,
tool_results,
compaction_blocks,
) = self.extract_response_content(completion_response=completion_response)
(
text_content,
citations,
thinking_blocks,
reasoning_content,
tool_calls,
web_search_results,
tool_results,
compaction_blocks,
) = self.extract_response_content(completion_response=completion_response)
if (
prefix_prompt is not None
and not text_content.startswith(prefix_prompt)
and not litellm.disable_add_prefix_to_prompt
):
text_content = prefix_prompt + text_content
if (
prefix_prompt is not None
and not text_content.startswith(prefix_prompt)
and not litellm.disable_add_prefix_to_prompt
):
text_content = prefix_prompt + text_content
context_management: Optional[Dict] = completion_response.get(
"context_management"
)
provider_specific_fields = self._build_provider_specific_fields(
completion_response,
citations,
thinking_blocks,
web_search_results,
tool_results,
compaction_blocks,
tool_calls,
)
container: Optional[Dict] = completion_response.get("container")
_message = litellm.Message(
tool_calls=tool_calls,
content=text_content or None,
provider_specific_fields=provider_specific_fields,
thinking_blocks=thinking_blocks,
reasoning_content=reasoning_content,
)
_message.provider_specific_fields = provider_specific_fields
provider_specific_fields: Dict[str, Any] = {
"citations": citations,
"thinking_blocks": thinking_blocks,
}
if context_management is not None:
provider_specific_fields["context_management"] = context_management
if web_search_results is not None:
provider_specific_fields["web_search_results"] = web_search_results
if tool_results is not None:
provider_specific_fields["tool_results"] = tool_results
if container is not None:
provider_specific_fields["container"] = container
if compaction_blocks is not None:
provider_specific_fields["compaction_blocks"] = compaction_blocks
json_mode_message = self._transform_response_for_json_mode(
json_mode=json_mode,
tool_calls=tool_calls,
)
if json_mode_message is not None:
completion_response["stop_reason"] = "stop"
_message = json_mode_message
_message = litellm.Message(
tool_calls=tool_calls,
content=text_content or None,
provider_specific_fields=provider_specific_fields,
thinking_blocks=thinking_blocks,
reasoning_content=reasoning_content,
)
_message.provider_specific_fields = provider_specific_fields
model_response.choices[0].message = _message
model_response._hidden_params["original_response"] = completion_response[
"content"
]
model_response.choices[0].finish_reason = cast(
OpenAIChatCompletionFinishReason,
map_finish_reason(completion_response["stop_reason"]),
)
## HANDLE JSON MODE - anthropic returns single function call
json_mode_message = self._transform_response_for_json_mode(
json_mode=json_mode,
tool_calls=tool_calls,
)
if json_mode_message is not None:
completion_response["stop_reason"] = "stop"
_message = json_mode_message
model_response.choices[0].message = _message # type: ignore
model_response._hidden_params["original_response"] = completion_response[
"content"
] # allow user to access raw anthropic tool calling response
model_response.choices[0].finish_reason = cast(
OpenAIChatCompletionFinishReason,
map_finish_reason(completion_response["stop_reason"]),
)
## CALCULATING USAGE
usage = self.calculate_usage(
usage_object=completion_response["usage"],
reasoning_content=reasoning_content,
completion_response=completion_response,
speed=speed,
)
setattr(model_response, "usage", usage) # type: ignore
setattr(model_response, "usage", usage)
model_response.created = int(time.time())
model_response.model = completion_response["model"]
_hidden_params["provider_specific_fields"] = provider_specific_fields
model_response._hidden_params = _hidden_params
return model_response

View file

@ -359,9 +359,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
Returns:
List of beta header strings
"""
from litellm.types.llms.anthropic import (
ANTHROPIC_EFFORT_BETA_HEADER,
)
from litellm.types.llms.anthropic import ANTHROPIC_EFFORT_BETA_HEADER
betas = []
@ -390,7 +388,8 @@ class AnthropicModelInfo(BaseLLMModelInfo):
def get_anthropic_headers(
self,
api_key: str,
api_key: Optional[str] = None,
auth_token: Optional[str] = None,
anthropic_version: Optional[str] = None,
computer_tool_used: Optional[str] = None,
prompt_caching_set: bool = False,
@ -451,7 +450,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-dangerous-direct-browser-access"] = "true"
betas.add(ANTHROPIC_OAUTH_BETA_HEADER)
else:
elif auth_token and not api_key:
headers["authorization"] = f"Bearer {auth_token}"
elif api_key:
headers["x-api-key"] = api_key
if user_anthropic_beta_headers is not None:
@ -485,9 +486,14 @@ class AnthropicModelInfo(BaseLLMModelInfo):
headers, api_key = optionally_handle_anthropic_oauth(
headers=headers, api_key=api_key
)
api_key = AnthropicModelInfo.get_api_key(api_key)
# Resolve auth_token from ANTHROPIC_AUTH_TOKEN if api_key is not set
auth_token: Optional[str] = None
if api_key is None:
auth_token = AnthropicModelInfo.get_auth_token()
if api_key is None and auth_token is None:
raise litellm.AuthenticationError(
message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` in your environment vars",
message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` in your environment vars",
llm_provider="anthropic",
model=model,
)
@ -519,6 +525,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
prompt_caching_set=prompt_caching_set,
pdf_used=pdf_used,
api_key=api_key,
auth_token=auth_token,
file_id_used=file_id_used,
web_search_tool_used=web_search_tool_used,
is_vertex_request=optional_params.get("is_vertex_request", False),
@ -543,6 +550,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
return (
api_base
or get_secret_str("ANTHROPIC_API_BASE")
or get_secret_str("ANTHROPIC_BASE_URL")
or "https://api.anthropic.com"
)
@ -552,6 +560,35 @@ class AnthropicModelInfo(BaseLLMModelInfo):
return api_key or get_secret_str("ANTHROPIC_API_KEY")
@staticmethod
def get_auth_token(auth_token: Optional[str] = None) -> Optional[str]:
"""Get auth token from ANTHROPIC_AUTH_TOKEN env var.
Unlike api_key (which uses X-Api-Key header), auth_token uses
Authorization: Bearer header, matching the official Anthropic SDK behavior.
"""
from litellm.secret_managers.main import get_secret_str
return auth_token or get_secret_str("ANTHROPIC_AUTH_TOKEN")
@staticmethod
def get_auth_header(api_key: Optional[str] = None) -> Optional[dict]:
"""Resolve Anthropic credentials and return the appropriate auth header dict.
Checks ANTHROPIC_API_KEY first (-> x-api-key), then
ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer).
Returns None if neither is available.
"""
resolved_key = AnthropicModelInfo.get_api_key(api_key)
if resolved_key is not None:
if is_anthropic_oauth_key(resolved_key):
return {"authorization": f"Bearer {resolved_key}"}
return {"x-api-key": resolved_key}
auth_token = AnthropicModelInfo.get_auth_token()
if auth_token is not None:
return {"authorization": f"Bearer {auth_token}"}
return None
@staticmethod
def get_base_model(model: Optional[str] = None) -> Optional[str]:
return model.replace("anthropic/", "") if model else None
@ -560,14 +597,16 @@ class AnthropicModelInfo(BaseLLMModelInfo):
self, api_key: Optional[str] = None, api_base: Optional[str] = None
) -> List[str]:
api_base = AnthropicModelInfo.get_api_base(api_base)
api_key = AnthropicModelInfo.get_api_key(api_key)
if api_base is None or api_key is None:
auth_header = AnthropicModelInfo.get_auth_header(api_key)
if api_base is None or auth_header is None:
raise ValueError(
"ANTHROPIC_API_BASE or ANTHROPIC_API_KEY is not set. Please set the environment variable, to query Anthropic's `/models` endpoint."
"ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN is not set. Please set the environment variable, to query Anthropic's `/models` endpoint."
)
headers = {"anthropic-version": "2023-06-01"}
headers.update(auth_header)
response = litellm.module_level_client.get(
url=f"{api_base}/v1/models",
headers={"x-api-key": api_key, "anthropic-version": "2023-06-01"},
headers=headers,
)
try:

View file

@ -23,7 +23,6 @@ from ...common_utils import (
optionally_handle_anthropic_oauth,
)
DEFAULT_ANTHROPIC_API_BASE = "https://api.anthropic.com"
DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01"
@ -127,7 +126,9 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
api_base = api_base or DEFAULT_ANTHROPIC_API_BASE
api_base = (
AnthropicModelInfo.get_api_base(api_base) or "https://api.anthropic.com"
)
if not api_base.endswith("/v1/messages"):
api_base = f"{api_base}/v1/messages"
return api_base
@ -142,17 +143,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> Tuple[dict, Optional[str]]:
import os
# Check for Anthropic OAuth token in Authorization header
headers, api_key = optionally_handle_anthropic_oauth(
headers=headers, api_key=api_key
)
if api_key is None:
api_key = os.getenv("ANTHROPIC_API_KEY")
if "x-api-key" not in headers and "authorization" not in headers and api_key:
headers["x-api-key"] = api_key
if "x-api-key" not in headers and "authorization" not in headers:
auth_header = AnthropicModelInfo.get_auth_header(api_key)
if auth_header is not None:
headers.update(auth_header)
if "anthropic-version" not in headers:
headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION
if "content-type" not in headers:

View file

@ -8,10 +8,8 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.openai import (
FileContentRequest,
HttpxBinaryResponseContent,
@ -85,9 +83,9 @@ class AnthropicFilesHandler:
# Get Anthropic API credentials
api_base = self.anthropic_model_info.get_api_base(api_base)
api_key = api_key or self.anthropic_model_info.get_api_key()
auth_header = self.anthropic_model_info.get_auth_header(api_key)
if not api_key:
if auth_header is None:
raise ValueError("Missing Anthropic API Key")
# Construct the Anthropic batch results URL
@ -97,8 +95,8 @@ class AnthropicFilesHandler:
headers = {
"accept": "application/json",
"anthropic-version": "2023-06-01",
"x-api-key": api_key,
}
headers.update(auth_header)
# Make the request to Anthropic
async_client = get_async_httpx_client(llm_provider=LlmProviders.ANTHROPIC)

View file

@ -94,14 +94,14 @@ class AnthropicFilesConfig(BaseFilesConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
api_key = AnthropicModelInfo.get_api_key(api_key)
if not api_key:
auth_header = AnthropicModelInfo.get_auth_header(api_key)
if auth_header is None:
raise ValueError(
"Anthropic API key is required. Set ANTHROPIC_API_KEY environment variable or pass api_key parameter."
"Anthropic API key is required. Set ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN environment variable or pass api_key parameter."
)
headers.update(
{
"x-api-key": api_key,
**auth_header,
"anthropic-version": "2023-06-01",
"anthropic-beta": ANTHROPIC_FILES_BETA_HEADER,
}

View file

@ -35,17 +35,18 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
"""Add Anthropic-specific headers"""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
# Get API key
# Get API key from litellm_params if available
api_key = None
if litellm_params:
if litellm_params is not None:
api_key = litellm_params.api_key
api_key = AnthropicModelInfo.get_api_key(api_key)
if not api_key:
raise ValueError("ANTHROPIC_API_KEY is required for Skills API")
auth_header = AnthropicModelInfo.get_auth_header(api_key)
if auth_header is None:
raise ValueError(
"ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN is required for Skills API"
)
# Add required headers
headers["x-api-key"] = api_key
headers.update(auth_header)
headers["anthropic-version"] = "2023-06-01"
# Add beta header for skills API

View file

@ -37104,5 +37104,157 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"volcengine/doubao-seed-2-0-pro-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 4.6e-07,
"output_cost_per_token": 2.3e-06,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 7e-07,
"output_cost_per_token": 3.5e-06,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 1.4e-06,
"output_cost_per_token": 7e-06,
"range": [
128000.0,
256000.0
]
}
]
},
"volcengine/doubao-seed-2-0-lite-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 8.7e-08,
"output_cost_per_token": 5.2e-07,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 1.3e-07,
"output_cost_per_token": 7.8e-07,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 2.6e-07,
"output_cost_per_token": 1.6e-06,
"range": [
128000.0,
256000.0
]
}
]
},
"volcengine/doubao-seed-2-0-mini-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 2.9e-08,
"output_cost_per_token": 2.9e-07,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 5.8e-08,
"output_cost_per_token": 5.8e-07,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 1.2e-07,
"output_cost_per_token": 1.2e-06,
"range": [
128000.0,
256000.0
]
}
]
},
"volcengine/doubao-seed-2-0-code-preview-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 4.6e-07,
"output_cost_per_token": 2.3e-06,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 7e-07,
"output_cost_per_token": 3.5e-06,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 1.4e-06,
"output_cost_per_token": 7e-06,
"range": [
128000.0,
256000.0
]
}
]
}
}

View file

@ -2955,7 +2955,9 @@ class LiteLLM_ErrorLogs(LiteLLMPydanticObjectBase):
endTime: Union[str, datetime, None]
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "rotated"]
AUDIT_ACTIONS = Literal[
"created", "updated", "deleted", "blocked", "unblocked", "rotated"
]
class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):

View file

@ -413,7 +413,9 @@ async def common_checks( # noqa: PLR0915
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
team_model_aliases=valid_token.team_model_aliases
if valid_token
else None,
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
@ -482,7 +484,9 @@ async def common_checks( # noqa: PLR0915
)
# 3.1. If organization is in budget
with tracer.trace("litellm.proxy.auth.common_checks.organization_max_budget_check"):
with tracer.trace(
"litellm.proxy.auth.common_checks.organization_max_budget_check"
):
await _organization_max_budget_check(
valid_token=valid_token,
team_object=team_object,

View file

@ -1438,10 +1438,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
): # user set proxy max budget
cache_key = "{}:spend".format(litellm_proxy_admin_name)
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
global_proxy_spend = (
await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
)
if global_proxy_spend is not None:

View file

@ -41,10 +41,8 @@ from litellm.proxy._experimental.mcp_server.db import (
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.auth.auth_checks import (
_cache_key_object,
_delete_cache_key_object,
can_team_access_model,
get_key_object,
get_org_object,
get_project_object,
get_team_object,
@ -1656,7 +1654,7 @@ async def _get_and_validate_existing_key(
LiteLLM_VerificationToken: The existing key row
Raises:
HTTPException: If key is not found
ProxyException: 404 if key is not found
"""
if prisma_client is None:
raise HTTPException(
@ -1664,16 +1662,18 @@ async def _get_and_validate_existing_key(
detail={"error": "Database not connected"},
)
existing_key_row = await prisma_client.get_data(
token=token,
table_name="key",
query_type="find_unique",
hashed_token = _hash_token_if_needed(token=token)
existing_key_row = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
if existing_key_row is None:
raise HTTPException(
status_code=404,
detail={"error": f"Key not found: {token}"},
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
return existing_key_row
@ -2111,19 +2111,11 @@ async def update_key_fn(
key = data_json.pop("key")
# get the row from db
if prisma_client is None:
raise Exception("Not connected to DB!")
existing_key_row = await prisma_client.get_data(
token=data.key, table_name="key", query_type="find_unique"
existing_key_row = await _get_and_validate_existing_key(
token=data.key,
prisma_client=prisma_client,
)
if existing_key_row is None:
raise HTTPException(
status_code=404,
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
await _validate_update_key_data(
data=data,
existing_key_row=existing_key_row,
@ -2158,6 +2150,8 @@ async def update_key_fn(
)
_data = {**non_default_values, "token": key}
if prisma_client is None:
raise Exception("Not connected to DB!")
response = await prisma_client.update_data(token=key, data=_data)
# Delete - key from cache, since it's been updated!
@ -2330,6 +2324,8 @@ async def bulk_update_keys(
error_message = error_detail.get("error", str(e))
else:
error_message = str(error_detail)
elif isinstance(e, ProxyException):
error_message = e.message
else:
error_message = str(e)
@ -4945,18 +4941,19 @@ async def block_key(
route="/key/block",
)
if litellm.store_audit_logs is True:
# make an audit log for key update
record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
# Check if the key exists before trying to block it
existing_record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
if existing_record is None:
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if record is None:
raise ProxyException(
message=f"Key {data.key} not found",
type=ProxyErrorTypes.bad_request_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -4970,7 +4967,7 @@ async def block_key(
object_id=hashed_token,
action="blocked",
updated_values="{}",
before_value=record.model_dump_json(),
before_value=existing_record.model_dump_json(),
)
)
)
@ -4979,24 +4976,9 @@ async def block_key(
where={"token": hashed_token}, data={"blocked": True} # type: ignore
)
## UPDATE KEY CACHE
### get cached object ###
key_object = await get_key_object(
## UPDATE KEY CACHE - invalidate so next read re-fetches from DB
await _delete_cache_key_object(
hashed_token=hashed_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
)
### update cached object ###
key_object.blocked = True
### store cached object ###
await _cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=key_object,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
@ -5068,18 +5050,19 @@ async def unblock_key(
route="/key/unblock",
)
if litellm.store_audit_logs is True:
# make an audit log for key update
record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
# Check if the key exists before trying to unblock it
existing_record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
if existing_record is None:
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if record is None:
raise ProxyException(
message=f"Key {data.key} not found",
type=ProxyErrorTypes.bad_request_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -5091,9 +5074,9 @@ async def unblock_key(
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=hashed_token,
action="blocked",
action="unblocked",
updated_values="{}",
before_value=record.model_dump_json(),
before_value=existing_record.model_dump_json(),
)
)
)
@ -5102,24 +5085,9 @@ async def unblock_key(
where={"token": hashed_token}, data={"blocked": False} # type: ignore
)
## UPDATE KEY CACHE
### get cached object ###
key_object = await get_key_object(
## UPDATE KEY CACHE - invalidate so next read re-fetches from DB
await _delete_cache_key_object(
hashed_token=hashed_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
)
### update cached object ###
key_object.blocked = False
### store cached object ###
await _cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=key_object,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)

View file

@ -3334,7 +3334,9 @@ def _convert_teams_to_response_models(
use_deleted_table: bool,
) -> List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]:
"""Convert raw Prisma team rows to response models."""
team_list: List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]] = []
team_list: List[
Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]
] = []
for team in teams:
try:
team_dict = team.model_dump()

View file

@ -22,6 +22,7 @@ from litellm.constants import (
ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS,
BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES,
)
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
from litellm.proxy.auth.route_checks import RouteChecks
@ -585,7 +586,11 @@ async def anthropic_proxy_route(
"""
[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)
"""
base_target_url = os.getenv("ANTHROPIC_API_BASE") or "https://api.anthropic.com"
base_target_url = (
os.getenv("ANTHROPIC_API_BASE")
or os.getenv("ANTHROPIC_BASE_URL")
or "https://api.anthropic.com"
)
encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction
@ -606,10 +611,11 @@ async def anthropic_proxy_route(
is_streaming_request = await is_streaming_request_fn(request)
## CREATE PASS-THROUGH
auth_header = AnthropicModelInfo.get_auth_header(anthropic_api_key or None)
endpoint_func = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"x-api-key": "{}".format(anthropic_api_key)},
custom_headers=auth_header if auth_header is not None else {},
_forward_headers=True,
is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path

View file

@ -7,6 +7,7 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
from litellm.llms.anthropic import get_anthropic_config
from litellm.llms.anthropic.chat.handler import (
ModelResponseIterator as AnthropicModelResponseIterator,
@ -124,10 +125,21 @@ class AnthropicPassthroughLoggingHandler:
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/"):
model_for_cost = f"{custom_llm_provider}/{model}"
router_model_id = logging_obj.get_router_model_id()
custom_pricing = use_custom_pricing_for_model(
litellm_params=(
logging_obj.litellm_params
if hasattr(logging_obj, "litellm_params")
else None
)
)
response_cost = litellm.completion_cost(
completion_response=litellm_model_response,
model=model_for_cost,
custom_llm_provider=custom_llm_provider,
custom_pricing=custom_pricing,
router_model_id=router_model_id,
)
kwargs["response_cost"] = response_cost
@ -319,9 +331,7 @@ class AnthropicPassthroughLoggingHandler:
import base64
from litellm._uuid import uuid
from litellm.llms.anthropic.batches.transformation import (
AnthropicBatchesConfig,
)
from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig
from litellm.types.utils import Choices, SpecialEnums
try:

View file

@ -12,6 +12,7 @@ import click
import httpx
from dotenv import load_dotenv
import litellm
from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY
from litellm.secret_managers.main import get_secret_bool
@ -387,7 +388,7 @@ class ProxyInitializationHelpers:
@click.option("--api_base", default=None, help="API base URL.")
@click.option(
"--api_version",
default="2024-07-01-preview",
default=litellm.AZURE_DEFAULT_API_VERSION,
help="For azure - pass in the api version.",
)
@click.option(

View file

@ -1878,28 +1878,33 @@ class ProxyLogging:
)
input: Union[list, str, dict] = ""
normalized_call_type: Optional[str] = None
if "messages" in request_data and isinstance(
request_data["messages"], list
):
input = request_data["messages"]
litellm_logging_obj.model_call_details["messages"] = input
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
litellm_logging_obj.call_type = CallTypes.acompletion.value
normalized_call_type = CallTypes.acompletion.value
elif "prompt" in request_data and isinstance(request_data["prompt"], str):
input = request_data["prompt"]
litellm_logging_obj.model_call_details["prompt"] = input
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
litellm_logging_obj.call_type = CallTypes.atext_completion.value
normalized_call_type = CallTypes.atext_completion.value
elif "input" in request_data and isinstance(request_data["input"], list):
input = request_data["input"]
litellm_logging_obj.model_call_details["input"] = input
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
litellm_logging_obj.call_type = CallTypes.aembedding.value
normalized_call_type = CallTypes.aembedding.value
if normalized_call_type is not None:
litellm_logging_obj.call_type = normalized_call_type
litellm_logging_obj.model_call_details[
"call_type"
] = normalized_call_type
# Pass-through endpoints are logged via the callback loop's
# async_post_call_failure_hook — skip pre_call and failure handlers.
if litellm_logging_obj.call_type == CallTypes.pass_through.value:
return
litellm_logging_obj.pre_call(
input=input,
api_key="",

View file

@ -107,6 +107,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
self._reasoning_done_emitted = False
self._reasoning_item_id: Optional[str] = None
self._accumulated_reasoning_content_parts: List[str] = []
self._accumulated_provider_specific_fields: Dict[str, Any] = {}
def _get_or_assign_tool_output_index(self, call_id: str) -> int:
existing = self._tool_output_index_by_call_id.get(call_id)
@ -479,16 +480,36 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
event.__dict__["sequence_number"] = self._sequence_number
return event
def create_litellm_model_response(
self,
) -> Optional[ModelResponse]:
return cast(
def _merge_provider_specific_fields(self, src: dict) -> None:
"""Merge provider_specific_fields using last-value-wins for lists.
List-valued keys (web_search_results, tool_results,
code_interpreter_results, etc.) are emitted cumulatively — each
emission contains the full list so far. Using "last value wins"
matches stream_chunk_builder's semantics and avoids quadratic
growth from repeated extend calls.
"""
for key, val in src.items():
self._accumulated_provider_specific_fields[key] = val
def create_litellm_model_response(self) -> Optional[ModelResponse]:
response = cast(
Optional[ModelResponse],
stream_chunk_builder(
chunks=self.collected_chat_completion_chunks,
logging_obj=self.litellm_logging_obj,
),
)
if response is not None and self._accumulated_provider_specific_fields:
if (
not hasattr(response, "_hidden_params")
or response._hidden_params is None
):
response._hidden_params = {}
response._hidden_params.setdefault("provider_specific_fields", {}).update(
self._accumulated_provider_specific_fields
)
return response
@staticmethod
def _snapshot_chunk_for_stream_chunk_builder(
@ -853,6 +874,17 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
if chunk is not None:
chunk = cast(ModelResponseStream, chunk)
self._ensure_output_item_for_chunk(chunk)
# Accumulate provider_specific_fields from chunk and delta
for src in (
getattr(chunk, "provider_specific_fields", None),
getattr(
chunk.choices[0].delta if chunk.choices else None,
"provider_specific_fields",
None,
),
):
if src and isinstance(src, dict):
self._merge_provider_specific_fields(src)
# Proceed to transformation
self.collected_chat_completion_chunks.append(
self._snapshot_chunk_for_stream_chunk_builder(chunk)
@ -964,6 +996,17 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
try:
chunk = self.litellm_custom_stream_wrapper.__next__()
self._ensure_output_item_for_chunk(chunk)
# Accumulate provider_specific_fields from chunk and delta
for src in (
getattr(chunk, "provider_specific_fields", None),
getattr(
chunk.choices[0].delta if chunk.choices else None,
"provider_specific_fields",
None,
),
):
if src and isinstance(src, dict):
self._merge_provider_specific_fields(src)
# Emit any just-queued output_item event
if self._pending_response_events:
return self._pending_response_events.pop(0)

View file

@ -42,6 +42,7 @@ from litellm.types.llms.openai import (
from litellm.types.responses.main import (
GenericResponseOutputItem,
GenericResponseOutputItemContentAnnotation,
OutputCodeInterpreterCall,
OutputFunctionToolCall,
OutputImageGenerationCall,
OutputText,
@ -1696,6 +1697,7 @@ class LiteLLMCompletionResponsesConfig:
) -> List[
Union[
GenericResponseOutputItem,
OutputCodeInterpreterCall,
OutputFunctionToolCall,
OutputImageGenerationCall,
ResponseFunctionToolCall,
@ -1704,6 +1706,7 @@ class LiteLLMCompletionResponsesConfig:
responses_output: List[
Union[
GenericResponseOutputItem,
OutputCodeInterpreterCall,
OutputFunctionToolCall,
OutputImageGenerationCall,
ResponseFunctionToolCall,
@ -1725,8 +1728,63 @@ class LiteLLMCompletionResponsesConfig:
chat_completion_response=chat_completion_response
)
)
# Convert server-side tool results (e.g. Anthropic code execution)
# into code_interpreter_call output items, replacing the corresponding
# function_call items so the output matches OpenAI's native shape.
tool_result_items = (
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(
chat_completion_response
)
)
if tool_result_items:
result_by_id = {item.id: item for item in tool_result_items}
replaced_ids = set(result_by_id.keys())
responses_output = [
(
result_by_id[getattr(item, "call_id", None)]
if (
getattr(item, "type", None) == "function_call"
and getattr(item, "call_id", None) in replaced_ids
)
else item
)
for item in responses_output
]
return responses_output
@staticmethod
def _extract_tool_result_output_items(
chat_completion_response: ModelResponse,
) -> list:
"""Extract pre-built code_interpreter_call output items from provider_specific_fields.
Provider transformers (e.g. Anthropic) convert their native tool results
into OutputCodeInterpreterCall objects and store them in
provider_specific_fields["code_interpreter_results"]. This method
simply retrieves them — no provider-specific parsing here.
"""
output_items: list = []
for choice in chat_completion_response.choices or []:
message = getattr(choice, "message", None)
if not message:
continue
psf = getattr(message, "provider_specific_fields", None)
if not psf or not isinstance(psf, dict):
continue
results = psf.get("code_interpreter_results")
if results and isinstance(results, list):
for item in results:
# In the streaming path, items are plain dicts after
# model_dump() in stream_chunk_builder. Reconstruct
# Pydantic objects so responses_output has a uniform type.
if isinstance(item, dict):
output_items.append(OutputCodeInterpreterCall(**item))
else:
output_items.append(item)
return output_items
@staticmethod
def _extract_reasoning_output_items(
chat_completion_response: ModelResponse,

View file

@ -166,11 +166,12 @@ class BaseResponsesAPIStreamingIterator:
)
setattr(item, "encrypted_content", wrapped_content)
# Store the completed response
if (
openai_responses_api_chunk
and getattr(openai_responses_api_chunk, "type", None)
== ResponsesAPIStreamEvents.RESPONSE_COMPLETED
# Store the completed response (also for incomplete/failed so logging still fires)
_chunk_type = getattr(openai_responses_api_chunk, "type", None)
if openai_responses_api_chunk and _chunk_type in (
ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
ResponsesAPIStreamEvents.RESPONSE_FAILED,
):
self.completed_response = openai_responses_api_chunk
# Add cost to usage object if include_cost_in_streaming_usage is True
@ -195,10 +196,12 @@ class BaseResponsesAPIStreamingIterator:
if cost is not None:
setattr(usage_obj, "cost", cost)
except Exception:
# If cost calculation fails, continue without cost
pass
self._handle_logging_completed_response()
if _chunk_type == ResponsesAPIStreamEvents.RESPONSE_FAILED:
self._handle_logging_failed_response()
else:
self._handle_logging_completed_response()
return openai_responses_api_chunk
@ -216,6 +219,32 @@ class BaseResponsesAPIStreamingIterator:
"""Base implementation - should be overridden by subclasses"""
pass
def _handle_logging_failed_response(self):
"""
Handle logging for RESPONSE_FAILED events by routing to failure handlers.
Unlike _handle_logging_completed_response (which calls success handlers),
this constructs an exception from the response error and routes to
async_failure_handler / failure_handler so logging integrations correctly
record the call as failed.
"""
response_obj = (
getattr(self.completed_response, "response", None)
if self.completed_response
else None
)
error_info = getattr(response_obj, "error", None) if response_obj else None
error_message = "Response failed"
if isinstance(error_info, dict):
error_message = error_info.get("message", str(error_info))
exception = litellm.APIError(
status_code=500,
message=error_message,
llm_provider=self.custom_llm_provider or "",
model=self.model or "",
)
self._handle_failure(exception)
async def _call_post_streaming_deployment_hook(self, chunk):
"""
Allow callbacks to modify streaming chunks before returning (parity with chat).

View file

@ -3874,14 +3874,23 @@ class Router:
The response from the handler function
"""
handler_name = original_function.__name__
metadata_variable_name = _get_router_metadata_variable_name(
function_name="generic_api_call"
)
try:
verbose_router_logger.debug(
f"Inside _generic_api_call() - handler: {handler_name}, model: {model}; kwargs: {kwargs}"
)
self._update_kwargs_before_fallbacks(
model=model,
kwargs=kwargs,
metadata_variable_name=metadata_variable_name,
)
deployment = self.get_available_deployment(
model=model,
messages=kwargs.get("messages", None),
specific_deployment=kwargs.pop("specific_deployment", None),
request_kwargs=kwargs,
)
self._update_kwargs_with_deployment(
deployment=deployment, kwargs=kwargs, function_name="generic_api_call"

View file

@ -84,6 +84,7 @@ from typing_extensions import Annotated, Dict, Required, TypedDict, override
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
from litellm.types.responses.main import (
GenericResponseOutputItem,
OutputCodeInterpreterCall,
OutputFunctionToolCall,
OutputImageGenerationCall,
)
@ -1242,6 +1243,7 @@ class ResponsesAPIResponse(BaseLiteLLMOpenAIResponseObject):
List[
Union[
GenericResponseOutputItem,
OutputCodeInterpreterCall,
OutputFunctionToolCall,
OutputImageGenerationCall,
ResponseFunctionToolCall,
@ -1308,13 +1310,16 @@ class ResponsesAPIResponse(BaseLiteLLMOpenAIResponseObject):
if not isinstance(serialized, list):
return serialized
return [
{
k: v
for k, v in item.items()
if v is not None or k not in ("status", "content", "encrypted_content")
}
if isinstance(item, dict) and item.get("type") == "reasoning"
else item
(
{
k: v
for k, v in item.items()
if v is not None
or k not in ("status", "content", "encrypted_content")
}
if isinstance(item, dict) and item.get("type") == "reasoning"
else item
)
for item in serialized
]

View file

@ -49,6 +49,42 @@ class OutputImageGenerationCall(BaseLiteLLMOpenAIResponseObject):
result: Optional[str] # Base64 encoded image data (without data:image prefix)
class OutputCodeInterpreterCallLog(BaseLiteLLMOpenAIResponseObject):
"""Log output from a code interpreter call"""
type: Literal["logs"]
logs: str
class OutputCodeInterpreterCall(BaseLiteLLMOpenAIResponseObject):
"""A code interpreter / code execution call output"""
type: Literal["code_interpreter_call"]
id: str
code: Optional[str]
container_id: Optional[str]
status: Literal["in_progress", "completed", "incomplete", "failed"]
outputs: Optional[List[OutputCodeInterpreterCallLog]]
def build_code_interpreter_log_outputs(
content: Any,
) -> Optional[List[OutputCodeInterpreterCallLog]]:
"""Convert Anthropic bash_code_execution stdout/stderr to log outputs.
Shared by streaming (handler.py) and non-streaming (transformation.py) paths.
"""
if not isinstance(content, dict):
return None
parts = []
if content.get("stdout"):
parts.append(content["stdout"])
if content.get("stderr"):
parts.append(f"STDERR: {content['stderr']}")
logs = "".join(parts)
return [OutputCodeInterpreterCallLog(type="logs", logs=logs)] if logs else None
class GenericResponseOutputItem(BaseLiteLLMOpenAIResponseObject):
"""
Generic response API output item

View file

@ -6183,7 +6183,10 @@ def validate_environment( # noqa: PLR0915
["AZURE_API_BASE", "AZURE_API_VERSION", "AZURE_API_KEY"]
)
elif custom_llm_provider == "anthropic":
if "ANTHROPIC_API_KEY" in os.environ:
if (
"ANTHROPIC_API_KEY" in os.environ
or "ANTHROPIC_AUTH_TOKEN" in os.environ
):
keys_in_environment = True
else:
missing_keys.append("ANTHROPIC_API_KEY")
@ -6422,7 +6425,10 @@ def validate_environment( # noqa: PLR0915
missing_keys.append("OPENAI_API_KEY")
## anthropic
elif model in litellm.anthropic_models:
if "ANTHROPIC_API_KEY" in os.environ:
if (
"ANTHROPIC_API_KEY" in os.environ
or "ANTHROPIC_AUTH_TOKEN" in os.environ
):
keys_in_environment = True
else:
missing_keys.append("ANTHROPIC_API_KEY")
@ -8616,9 +8622,7 @@ class ProviderConfigManager:
return ManusFilesConfig()
elif LlmProviders.ANTHROPIC == provider:
from litellm.llms.anthropic.files.transformation import (
AnthropicFilesConfig,
)
from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig
return AnthropicFilesConfig()
return None

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.82.4"
version = "1.82.5"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -184,7 +184,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.82.4"
version = "1.82.5"
version_files = [
"pyproject.toml:^version"
]

View file

@ -30,9 +30,11 @@ from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterat
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponseFailedEvent,
ResponseIncompleteEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
OutputTextDeltaEvent
OutputTextDeltaEvent,
)
@ -429,3 +431,155 @@ class TestBaseResponsesAPIStreamingIterator:
mock_logging_obj.async_failure_handler.assert_not_called()
mock_logging_obj.failure_handler.assert_not_called()
def test_process_chunk_response_failed_calls_failure_handler(self):
"""
Test that a RESPONSE_FAILED event routes to failure handlers,
not success handlers. Failed responses represent genuine LLM-level
errors and should be logged as failures.
"""
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_lines = Mock()
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.async_failure_handler = Mock()
mock_logging_obj.failure_handler = Mock()
mock_logging_obj.async_success_handler = Mock()
mock_logging_obj.success_handler = Mock()
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_responses_api_response = Mock(spec=ResponsesAPIResponse)
mock_responses_api_response.id = "resp_failed_123"
mock_responses_api_response.error = {
"type": "server_error",
"message": "The model encountered an error",
}
mock_responses_api_response.usage = None
mock_failed_event = Mock(spec=ResponseFailedEvent)
mock_failed_event.type = ResponsesAPIStreamEvents.RESPONSE_FAILED
mock_failed_event.response = mock_responses_api_response
mock_config.transform_streaming_response.return_value = mock_failed_event
iterator = ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-4",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
litellm_metadata={"model_info": {"id": "model_123"}},
custom_llm_provider="openai",
)
test_chunk_data = {
"type": "response.failed",
"response": {
"id": "resp_failed_123",
"error": {
"type": "server_error",
"message": "The model encountered an error",
},
},
}
with patch.object(
ResponsesAPIRequestUtils,
"_update_responses_api_response_id_with_model_id",
return_value=mock_responses_api_response,
), patch(
"litellm.responses.streaming_iterator.run_async_function"
) as mock_run_async, patch(
"litellm.responses.streaming_iterator.executor"
) as mock_executor:
result = iterator._process_chunk(json.dumps(test_chunk_data))
assert result is not None
assert result.type == ResponsesAPIStreamEvents.RESPONSE_FAILED
assert iterator.completed_response == result
# Failure handler should have been called via _handle_failure
mock_run_async.assert_called_once()
call_kwargs = mock_run_async.call_args
assert (
call_kwargs[1]["async_function"]
== mock_logging_obj.async_failure_handler
)
mock_executor.submit.assert_called_once()
submit_args = mock_executor.submit.call_args
assert submit_args[0][0] == mock_logging_obj.failure_handler
def test_process_chunk_response_incomplete_calls_success_handler(self):
"""
Test that a RESPONSE_INCOMPLETE event routes to success handlers.
Incomplete responses (e.g. max_output_tokens reached) are still valid
responses with usage data — analogous to finish_reason='length' in chat.
"""
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_lines = Mock()
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.async_failure_handler = Mock()
mock_logging_obj.failure_handler = Mock()
mock_logging_obj.async_success_handler = Mock()
mock_logging_obj.success_handler = Mock()
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_responses_api_response = Mock(spec=ResponsesAPIResponse)
mock_responses_api_response.id = "resp_incomplete_123"
mock_responses_api_response.incomplete_details = {
"reason": "max_output_tokens"
}
mock_responses_api_response.usage = None
mock_incomplete_event = Mock(spec=ResponseIncompleteEvent)
mock_incomplete_event.type = ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE
mock_incomplete_event.response = mock_responses_api_response
mock_config.transform_streaming_response.return_value = mock_incomplete_event
iterator = ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-4",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
litellm_metadata={"model_info": {"id": "model_123"}},
custom_llm_provider="openai",
)
test_chunk_data = {
"type": "response.incomplete",
"response": {
"id": "resp_incomplete_123",
"incomplete_details": {"reason": "max_output_tokens"},
},
}
with patch.object(
ResponsesAPIRequestUtils,
"_update_responses_api_response_id_with_model_id",
return_value=mock_responses_api_response,
), patch(
"asyncio.create_task"
) as mock_create_task, patch(
"litellm.responses.streaming_iterator.executor"
) as mock_executor:
result = iterator._process_chunk(json.dumps(test_chunk_data))
assert result is not None
assert result.type == ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE
assert iterator.completed_response == result
# Success handler should have been called (via _handle_logging_completed_response)
mock_create_task.assert_called_once()
mock_executor.submit.assert_called_once()
# Failure handlers should NOT have been called
mock_logging_obj.async_failure_handler.assert_not_called()
mock_logging_obj.failure_handler.assert_not_called()

View file

@ -2278,6 +2278,75 @@ async def test_post_call_failure_hook_auth_error_llm_api_route():
mock_handle_logging.assert_called_once()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"request_data, route, expected_call_type",
[
(
{"model": "bad-model", "messages": [{"role": "user", "content": "hello"}]},
"/v1/chat/completions",
"acompletion",
),
(
{"model": "bad-model", "prompt": "hello"},
"/v1/completions",
"atext_completion",
),
(
{"model": "bad-model", "input": ["hello"]},
"/v1/embeddings",
"aembedding",
),
],
)
async def test_handle_logging_proxy_only_error_syncs_normalized_call_type(
request_data, route, expected_call_type
):
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.utils import ProxyLogging
cache = DualCache()
proxy_logging = ProxyLogging(user_api_key_cache=cache)
captured_logging_obj = {}
original_function_setup = litellm.utils.function_setup
def _capture_function_setup(*args, **kwargs):
logging_obj, data = original_function_setup(*args, **kwargs)
captured_logging_obj["logging_obj"] = logging_obj
return logging_obj, data
with patch(
"litellm.proxy.utils.litellm.utils.function_setup",
side_effect=_capture_function_setup,
), patch.object(
Logging, "async_failure_handler", new=AsyncMock(return_value=None)
), patch.object(
Logging, "failure_handler", return_value=None
), patch(
"litellm.proxy.utils.threading.Thread"
) as mock_thread:
mock_thread.return_value.start = Mock()
await proxy_logging._handle_logging_proxy_only_error(
request_data=request_data,
user_api_key_dict=UserAPIKeyAuth(
api_key="test_key",
user_id="test_user",
token="test_token",
request_route=route,
),
route=route,
original_exception=HTTPException(status_code=400, detail="bad request"),
)
logging_obj = captured_logging_obj["logging_obj"]
assert logging_obj.call_type == expected_call_type
assert logging_obj.model_call_details["call_type"] == expected_call_type
@pytest.mark.asyncio
async def test_during_call_hook_parallel_execution():
"""

View file

@ -1568,6 +1568,91 @@ def test_handle_clientside_credential_with_deployment_model_name(model_list):
print("✓ _handle_clientside_credential test passed!")
def test_sync_generic_api_call_preserves_requested_model_group_in_logs():
router = Router(
model_list=[
{
"model_name": "claude-sonnet-4-6",
"litellm_params": {
"model": "bedrock/global.anthropic.claude-sonnet-4-6",
"aws_access_key_id": "test-access-key",
"aws_secret_access_key": "test-secret-key",
"aws_region_name": "us-west-2",
},
}
]
)
try:
captured_kwargs = {}
def mock_original_function(**kwargs):
captured_kwargs.update(kwargs)
return {"status": "ok"}
response = router._generic_api_call_with_fallbacks(
model="claude-sonnet-4-6",
original_function=mock_original_function,
)
assert response == {"status": "ok"}
assert (
captured_kwargs["model"] == "bedrock/global.anthropic.claude-sonnet-4-6"
)
assert (
captured_kwargs["litellm_metadata"]["model_group"] == "claude-sonnet-4-6"
)
assert (
captured_kwargs["litellm_metadata"]["deployment"]
== "bedrock/global.anthropic.claude-sonnet-4-6"
)
finally:
router.discard()
def test_sync_generic_api_call_uses_request_kwargs_for_deployment_selection():
router = Router(
model_list=[
{
"model_name": "regional-model",
"litellm_params": {
"model": "anthropic/us-model",
"api_key": "test-api-key",
"region_name": "us",
},
},
{
"model_name": "regional-model",
"litellm_params": {
"model": "anthropic/eu-model",
"api_key": "test-api-key",
"region_name": "eu",
},
},
],
enable_pre_call_checks=True,
)
try:
captured_kwargs = {}
def mock_original_function(**kwargs):
captured_kwargs.update(kwargs)
return {"status": "ok"}
response = router._generic_api_call_with_fallbacks(
model="regional-model",
original_function=mock_original_function,
messages=[{"role": "user", "content": "Hello from Europe"}],
allowed_model_region="eu",
)
assert response == {"status": "ok"}
assert captured_kwargs["model"] == "anthropic/eu-model"
finally:
router.discard()
@pytest.mark.parametrize(
"function_name, expected_metadata_key",
[

View file

@ -132,3 +132,57 @@ class TestLangsmithLoggerInit:
assert (
logger.sampling_rate >= 0.0
), f"sampling_rate should be non-negative, got {logger.sampling_rate}"
class TestLangsmithPrepareLogData:
"""Regression test for #24001: _prepare_log_data must inject
usage_metadata into outputs so LangSmith's Cost column is populated."""
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_outputs_contain_usage_metadata(self, mock_create_task):
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
)
payload = {
"id": "test-id",
"response": {"choices": [{"message": {"content": "hi"}}]},
"metadata": {},
"startTime": 1.0,
"endTime": 2.0,
"request_tags": [],
"error_str": None,
"status": "success",
"response_cost": 0.0042,
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
}
kwargs = {
"litellm_params": {"metadata": {}},
"standard_logging_object": payload,
}
credentials = {
"LANGSMITH_API_KEY": "test-key",
"LANGSMITH_PROJECT": "test-project",
"LANGSMITH_BASE_URL": "https://api.smith.langchain.com",
}
data = logger._prepare_log_data(
kwargs=kwargs,
response_obj=None,
start_time=1.0,
end_time=2.0,
credentials=credentials,
)
assert "usage_metadata" in data["outputs"]
um = data["outputs"]["usage_metadata"]
assert um["total_cost"] == 0.0042
assert um["input_tokens"] == 100
assert um["output_tokens"] == 50
assert um["total_tokens"] == 150

View file

@ -11,7 +11,8 @@ sys.path.insert(
import time
from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
from litellm.litellm_core_utils.litellm_logging import \
Logging as LitellmLogging
from litellm.litellm_core_utils.litellm_logging import set_callbacks
from litellm.types.utils import ModelResponse, TextCompletionResponse
@ -139,7 +140,8 @@ def test_sentry_environment():
def test_use_custom_pricing_for_model():
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
from litellm.litellm_core_utils.litellm_logging import \
use_custom_pricing_for_model
litellm_params = {
"custom_llm_provider": "azure",
@ -154,7 +156,8 @@ def test_use_custom_pricing_for_model_via_litellm_metadata():
Generic API call routes (/messages, /responses) store model_info
under litellm_metadata, not metadata. Regression test for #23185.
"""
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
from litellm.litellm_core_utils.litellm_logging import \
use_custom_pricing_for_model
litellm_params = {
"litellm_metadata": {
@ -170,7 +173,8 @@ def test_use_custom_pricing_for_model_via_litellm_metadata():
def test_use_custom_pricing_not_detected_litellm_metadata_no_pricing():
"""Should return False when litellm_metadata.model_info has no pricing keys."""
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
from litellm.litellm_core_utils.litellm_logging import \
use_custom_pricing_for_model
litellm_params = {
"litellm_metadata": {
@ -186,7 +190,8 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata():
does not carry _hidden_params (e.g. ResponsesAPIResponse from /v1/responses
streaming). Regression test for custom pricing on streaming responses."""
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import ResponsesAPIResponse
custom_model_id = "gpt-5-custom-pricing"
@ -256,6 +261,121 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata():
litellm.model_cost.pop(custom_model_id, None)
class TestGetRouterModelId:
"""Tests for the get_router_model_id helper method."""
def test_returns_id_from_litellm_metadata(self, logging_obj):
"""Should extract model_info.id from litellm_metadata."""
logging_obj.litellm_params = {
"litellm_metadata": {
"model_info": {"id": "custom-deploy-1"},
},
}
assert logging_obj.get_router_model_id() == "custom-deploy-1"
def test_returns_id_from_metadata(self, logging_obj):
"""Should fall back to metadata when litellm_metadata has no model_info."""
logging_obj.litellm_params = {
"metadata": {
"model_info": {"id": "custom-deploy-2"},
},
}
assert logging_obj.get_router_model_id() == "custom-deploy-2"
def test_prefers_litellm_metadata_over_metadata(self, logging_obj):
"""litellm_metadata should take priority over metadata."""
logging_obj.litellm_params = {
"litellm_metadata": {
"model_info": {"id": "from-litellm-meta"},
},
"metadata": {
"model_info": {"id": "from-meta"},
},
}
assert logging_obj.get_router_model_id() == "from-litellm-meta"
def test_returns_none_when_no_model_info(self, logging_obj):
"""Should return None when no model_info is present."""
logging_obj.litellm_params = {"api_base": ""}
assert logging_obj.get_router_model_id() is None
def test_returns_none_when_no_litellm_params(self):
"""Should return None when litellm_params is not set."""
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
obj = LiteLLMLoggingObj(
model="test",
messages=[],
stream=False,
call_type="completion",
start_time=time.time(),
litellm_call_id="x",
function_id="x",
)
# litellm_params exists but is empty by default
assert obj.get_router_model_id() is None
class TestAnthropicPassthroughCustomPricing:
"""Verify the Anthropic pass-through handler forwards custom pricing."""
def test_completion_cost_receives_custom_pricing_args(self):
"""_create_anthropic_response_logging_payload should pass
custom_pricing and router_model_id to litellm.completion_cost
when the logging object carries custom pricing in model_info."""
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import \
AnthropicPassthroughLoggingHandler
logging_obj = LiteLLMLoggingObj(
model="claude-sonnet-4-20250514",
messages=[{"role": "user", "content": "Hi"}],
stream=False,
call_type="anthropic_messages",
start_time=time.time(),
litellm_call_id="test-456",
function_id="test-fn",
)
logging_obj.update_environment_variables(
model="claude-sonnet-4-20250514",
user="",
optional_params={},
litellm_params={
"api_base": "",
"litellm_metadata": {
"model_info": {
"id": "claude-custom-pricing",
"input_cost_per_token": 0.5,
"output_cost_per_token": 1.5,
},
},
},
)
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
mock_response = ModelResponse()
mock_response.usage = {"prompt_tokens": 10, "completion_tokens": 5} # type: ignore
with patch("litellm.completion_cost", return_value=42.0) as mock_cost:
AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
litellm_model_response=mock_response,
model="claude-sonnet-4-20250514",
kwargs={},
start_time=time.time(),
end_time=time.time(),
logging_obj=logging_obj,
)
mock_cost.assert_called_once()
call_kwargs = mock_cost.call_args
assert call_kwargs.kwargs.get("custom_pricing") is True
assert call_kwargs.kwargs.get("router_model_id") == "claude-custom-pricing"
class TestUpdateFromKwargs:
"""Tests for the update_from_kwargs convenience wrapper."""
@ -321,9 +441,8 @@ class TestUpdateFromKwargs:
def test_custom_pricing_detected_via_litellm_metadata(self, logging_obj):
"""Custom pricing in litellm_metadata.model_info should set custom_pricing flag."""
from litellm.litellm_core_utils.litellm_logging import (
use_custom_pricing_for_model,
)
from litellm.litellm_core_utils.litellm_logging import \
use_custom_pricing_for_model
lm_meta = {
"model_info": {
@ -382,7 +501,8 @@ async def test_datadog_logger_not_shadowed_by_llm_obs(monkeypatch):
monkeypatch.setenv("DD_SITE", "us5.datadoghq.com")
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
from litellm.integrations.datadog.datadog_llm_obs import \
DataDogLLMObsLogger
from litellm.litellm_core_utils import litellm_logging as logging_module
logging_module._in_memory_loggers.clear()
@ -423,7 +543,8 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch):
) # no trailing slash on purpose
# Import after env vars are set (important if module-level caching exists)
from litellm.integrations.opentelemetry import OpenTelemetry # logger class
from litellm.integrations.opentelemetry import \
OpenTelemetry # logger class
from litellm.litellm_core_utils import litellm_logging as logging_module
logging_module._in_memory_loggers.clear()
@ -752,7 +873,8 @@ def test_success_handler_runs_guardrail_logging_hook_when_enabled(logging_obj):
def test_get_user_agent_tags():
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
tags = StandardLoggingPayloadSetup._get_user_agent_tags(
proxy_server_request={
@ -767,7 +889,8 @@ def test_get_user_agent_tags():
def test_get_request_tags():
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
tags = StandardLoggingPayloadSetup._get_request_tags(
litellm_params={"metadata": {"tags": ["test-tag"]}},
@ -794,7 +917,8 @@ def test_get_request_tags_from_metadata_and_litellm_metadata():
4. No tags in either
5. None values for metadata/litellm_metadata
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Test case 1: Tags in metadata only
tags = StandardLoggingPayloadSetup._get_request_tags(
@ -875,7 +999,8 @@ def test_get_request_tags_does_not_mutate_original_tags():
would cause User-Agent tags to be duplicated because the function was mutating
the original tags list instead of creating a copy.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Create metadata with original tags
original_tags = ["custom-tag-1", "custom-tag-2"]
@ -935,7 +1060,8 @@ def test_get_request_tags_does_not_mutate_original_tags():
def test_get_extra_header_tags():
"""Test the _get_extra_header_tags method with various scenarios."""
import litellm
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Store original value to restore later
original_extra_headers = getattr(litellm, "extra_spend_tag_headers", None)
@ -1156,7 +1282,8 @@ async def test_e2e_generate_cold_storage_object_key_successful():
from datetime import datetime, timezone
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Create test data
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
@ -1198,7 +1325,8 @@ async def test_e2e_generate_cold_storage_object_key_with_custom_logger_s3_path()
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Create test data
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
@ -1249,7 +1377,8 @@ async def test_e2e_generate_cold_storage_object_key_with_logger_no_s3_path():
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Create test data
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
@ -1296,7 +1425,8 @@ async def test_e2e_generate_cold_storage_object_key_not_configured():
from unittest.mock import patch
import litellm
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Create test data
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
@ -1320,7 +1450,8 @@ def test_get_final_response_obj_with_empty_response_obj_and_list_init():
When response_obj is empty (falsy), the method should return init_response_obj if it's a list.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Create test objects
class TestObject1:
@ -1356,7 +1487,8 @@ def test_get_usage_as_dict():
"""
Test get_usage_as_dict returns usage as plain dict from response_obj or combined_usage_object.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
from litellm.types.utils import Usage
# Test case 1: None response_obj returns empty usage dict
@ -1394,7 +1526,8 @@ def test_append_system_prompt_messages():
"""
Test append_system_prompt_messages prepends system message from kwargs to messages list.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Test case 1: system in kwargs with existing messages
kwargs = {"system": "You are a helpful assistant"}
@ -1465,7 +1598,8 @@ async def test_async_success_handler_sets_standard_logging_object_for_pass_throu
from datetime import datetime
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.types.utils import StandardPassThroughResponseObject
# Create a logging object for a pass-through endpoint
@ -1546,7 +1680,8 @@ async def test_async_success_handler_prevents_reprocessing_for_pass_through_endp
from datetime import datetime
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.types.utils import StandardPassThroughResponseObject
# Create a logging object for a pass-through endpoint
@ -1622,7 +1757,8 @@ async def test_async_success_handler_sets_standard_logging_object_for_streaming_
from datetime import datetime
from unittest.mock import patch
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.types.utils import StandardPassThroughResponseObject
# Create a logging object for a streaming pass-through endpoint
@ -1678,7 +1814,8 @@ def test_get_error_information_error_code_priority():
Test get_error_information prioritizes 'code' attribute over 'status_code' attribute
and handles edge cases like empty strings and "None" string values.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.litellm_logging import \
StandardLoggingPayloadSetup
# Test case 1: Exception with 'code' attribute (ProxyException style)
class ProxyException(Exception):
@ -1871,7 +2008,8 @@ async def test_async_success_handler_preserves_response_cost_for_pass_through_en
by pass-through handlers (Gemini/Vertex)."""
from datetime import datetime
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.types.utils import ModelResponse, Usage
logging_obj = LiteLLMLoggingObj(

View file

@ -1340,6 +1340,121 @@ def test_is_chunk_non_empty_with_valid_tool_calls(
)
def _make_chunk(content: Optional[str]) -> ModelResponseStream:
return ModelResponseStream(
id="test",
created=1741037890,
model="test-model",
choices=[StreamingChoices(index=0, delta=Delta(content=content))],
)
def _build_chunks(pattern: list[str], N: int) -> list[ModelResponseStream]:
"""
Build a list of chunks based on a pattern specification.
"""
chunks = []
for i, p in enumerate(pattern):
if p == "same":
chunks.append(_make_chunk("same_chunk"))
elif p == "diff":
chunks.append(_make_chunk(f"chunk_{i}"))
else:
chunks.append(_make_chunk(p))
return chunks
_REPETITION_TEST_CASES = [
# Basic cases
pytest.param(
["same"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
True,
id="all_identical_raises",
),
pytest.param(
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1),
False,
id="below_threshold_no_raise",
),
pytest.param(
[None] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
False,
id="none_content_no_raise",
),
pytest.param(
[""] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
False,
id="empty_content_no_raise",
),
# Short content (len <= 2) should not raise
pytest.param(
["##"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
False,
id="short_content_2chars_no_raise",
),
pytest.param(
["{"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
False,
id="short_content_1char_no_raise",
),
pytest.param(
["ab"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
False,
id="short_content_2chars_ab_no_raise",
),
# All different chunks
pytest.param(
["diff"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
False,
id="all_different_no_raise",
),
# One chunk different at various positions
pytest.param(
["different_first"] + ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1),
False,
id="first_chunk_different_no_raise",
),
pytest.param(
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1) + ["different_last"],
False,
id="last_chunk_different_no_raise",
),
pytest.param(
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT // 2 + 1) + ["different_mid"] + ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - litellm.REPEATED_STREAMING_CHUNK_LIMIT // 2 + 1),
False,
id="middle_chunk_different_no_raise",
),
pytest.param(
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 2) + ["diff", "diff"],
False,
id="last_two_different_no_raise",
),
pytest.param(
["diff"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT + ["same"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT + ["diff"],
True,
id="in_between_same_and_diff_raise",
),
]
@pytest.mark.parametrize("chunks_pattern,should_raise", _REPETITION_TEST_CASES)
def test_raise_on_model_repetition(
initialized_custom_stream_wrapper: CustomStreamWrapper,
chunks_pattern: list,
should_raise: bool,
):
wrapper = initialized_custom_stream_wrapper
chunks = _build_chunks(chunks_pattern, len(chunks_pattern))
if should_raise:
with pytest.raises(litellm.InternalServerError) as exc_info:
for chunk in chunks:
wrapper.chunks.append(chunk)
wrapper.raise_on_model_repetition()
assert "repeating the same chunk" in str(exc_info.value)
else:
for chunk in chunks:
wrapper.chunks.append(chunk)
wrapper.raise_on_model_repetition()
def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj):
"""
Test that provider-reported usage from a post-finish_reason chunk

View file

@ -6,6 +6,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
)
from litellm.types.responses.main import OutputCodeInterpreterCall
def test_redacted_thinking_content_block_delta():
@ -479,14 +480,22 @@ def test_partial_json_chunk_accumulation():
# First partial chunk should return None (still accumulating)
result1 = iterator._parse_sse_data(f"data:{partial_chunk_1}")
assert result1 is None, "First partial chunk should return None while accumulating"
assert iterator.chunk_type == "accumulated_json", "Should switch to accumulated_json mode"
assert iterator.accumulated_json == partial_chunk_1, "Should have accumulated first part"
assert (
iterator.chunk_type == "accumulated_json"
), "Should switch to accumulated_json mode"
assert (
iterator.accumulated_json == partial_chunk_1
), "Should have accumulated first part"
# Second partial chunk should complete the JSON and return a parsed result
result2 = iterator._parse_sse_data(f"data:{partial_chunk_2}")
assert result2 is not None, "Second chunk should return parsed result"
assert iterator.accumulated_json == "", "Buffer should be cleared after successful parse"
assert result2.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result2.choices[0].delta.content}'"
assert (
iterator.accumulated_json == ""
), "Buffer should be cleared after successful parse"
assert (
result2.choices[0].delta.content == "Hello"
), f"Expected 'Hello', got '{result2.choices[0].delta.content}'"
def test_complete_json_chunk_no_accumulation():
@ -503,7 +512,9 @@ def test_complete_json_chunk_no_accumulation():
assert result is not None, "Complete chunk should return parsed result immediately"
assert iterator.chunk_type == "valid_json", "Should remain in valid_json mode"
assert iterator.accumulated_json == "", "Buffer should remain empty"
assert result.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result.choices[0].delta.content}'"
assert (
result.choices[0].delta.content == "Hello"
), f"Expected 'Hello', got '{result.choices[0].delta.content}'"
def test_multiple_partial_chunks_accumulation():
@ -620,7 +631,9 @@ def test_web_search_tool_result_no_extra_tool_calls():
# Should have exactly 2 tool calls:
# 1. From content_block_start (server_tool_use) with id and name
# 2. From content_block_delta with the actual query
assert len(tool_calls_emitted) == 2, f"Expected 2 tool calls, got {len(tool_calls_emitted)}"
assert (
len(tool_calls_emitted) == 2
), f"Expected 2 tool calls, got {len(tool_calls_emitted)}"
# First tool call should have the id and name
assert tool_calls_emitted[0]["id"] == "srvtoolu_01ABC123"
@ -722,7 +735,10 @@ def test_web_search_tool_result_captured_in_provider_specific_fields():
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '{"query": "otter facts"}'},
"delta": {
"type": "input_json_delta",
"partial_json": '{"query": "otter facts"}',
},
},
# 4. content_block_stop for server_tool_use
{"type": "content_block_stop", "index": 0},
@ -822,7 +838,10 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields():
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '{"url": "https://example.com"}'},
"delta": {
"type": "input_json_delta",
"partial_json": '{"url": "https://example.com"}',
},
},
# 4. content_block_stop for server_tool_use
{"type": "content_block_stop", "index": 0},
@ -946,7 +965,7 @@ def test_web_fetch_tool_result_no_extra_tool_calls():
def test_container_in_provider_specific_fields_streaming():
"""
Test that container is captured in provider_specific_fields for streaming responses.
When container with skills is used, the container field should be present in
the provider_specific_fields of the message_delta chunk.
"""
@ -1025,7 +1044,9 @@ def test_container_in_provider_specific_fields_streaming():
]
# Verify container was captured
assert container_field is not None, "container should be captured in provider_specific_fields"
assert (
container_field is not None
), "container should be captured in provider_specific_fields"
assert (
container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p"
), "container id should match"
@ -1033,18 +1054,14 @@ def test_container_in_provider_specific_fields_streaming():
container_field["expires_at"] == "2025-12-16T04:57:16.913181Z"
), "expires_at should match"
assert len(container_field["skills"]) == 1, "Should have 1 skill"
assert (
container_field["skills"][0]["skill_id"] == "pptx"
), "skill_id should be pptx"
assert (
container_field["skills"][0]["version"] == "20251013"
), "version should match"
assert container_field["skills"][0]["skill_id"] == "pptx", "skill_id should be pptx"
assert container_field["skills"][0]["version"] == "20251013", "version should match"
def test_container_in_provider_specific_fields_non_streaming():
"""
Test that container is captured in provider_specific_fields for non-streaming responses.
When container with skills is used in non-streaming, the container field should be
present in the provider_specific_fields of the response.
"""
@ -1106,7 +1123,7 @@ def test_container_in_provider_specific_fields_non_streaming():
def test_container_absent_when_not_provided():
"""
Test that container is not added to provider_specific_fields when not provided.
This ensures we don't add empty or None container fields.
"""
iterator = ModelResponseIterator(
@ -1133,3 +1150,434 @@ def test_container_absent_when_not_provided():
assert (
"container" not in model_response.choices[0].delta.provider_specific_fields
), "container should not be present when not provided in delta"
def test_streaming_code_execution_produces_code_interpreter_results():
"""
Test that bash_code_execution_tool_result content blocks in streaming
produce code_interpreter_results in provider_specific_fields, so the
Responses API layer can use them without Anthropic-specific knowledge.
"""
chunks = [
{
"type": "message_start",
"message": {
"id": "msg_01XYZ",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 100, "output_tokens": 1},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "text",
"text": "",
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "Running code..."},
},
{"type": "content_block_stop", "index": 0},
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01ABC",
"name": "bash_code_execution",
"input": {"command": "echo hello"},
},
},
{"type": "content_block_stop", "index": 1},
{
"type": "content_block_start",
"index": 2,
"content_block": {
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_01ABC",
"content": {
"type": "bash_code_execution_result",
"stdout": "hello\n",
"stderr": "",
"return_code": 0,
},
},
},
{"type": "content_block_stop", "index": 2},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
iterator = ModelResponseIterator(None, sync_stream=True)
found_code_interpreter_results = False
for chunk in chunks:
parsed = iterator.chunk_parser(chunk)
psf = None
if parsed.choices and parsed.choices[0].delta:
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
if psf and "code_interpreter_results" in psf:
found_code_interpreter_results = True
results = psf["code_interpreter_results"]
assert len(results) == 1
assert isinstance(results[0], OutputCodeInterpreterCall)
assert results[0].type == "code_interpreter_call"
assert results[0].id == "srvtoolu_01ABC"
assert results[0].code == "echo hello"
assert results[0].outputs is not None
assert len(results[0].outputs) == 1
assert results[0].outputs[0].logs == "hello\n"
assert found_code_interpreter_results, (
"code_interpreter_results should appear in provider_specific_fields "
"when bash_code_execution_tool_result is streamed"
)
def test_streaming_multiple_code_executions_no_duplicates():
"""
Test that multiple code executions in a single streaming response emit
cumulative code_interpreter_results on each chunk (matching stream_chunk_builder's
"last value wins" contract). The final emission must contain ALL results.
"""
chunks = [
{
"type": "message_start",
"message": {
"id": "msg_01XYZ",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 100, "output_tokens": 1},
},
},
# First code execution
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01AAA",
"name": "bash_code_execution",
"input": {"command": "echo first"},
},
},
{"type": "content_block_stop", "index": 0},
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_01AAA",
"content": {
"type": "bash_code_execution_result",
"stdout": "first\n",
"stderr": "",
"return_code": 0,
},
},
},
{"type": "content_block_stop", "index": 1},
# Second code execution
{
"type": "content_block_start",
"index": 2,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01BBB",
"name": "bash_code_execution",
"input": {"command": "echo second"},
},
},
{"type": "content_block_stop", "index": 2},
{
"type": "content_block_start",
"index": 3,
"content_block": {
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_01BBB",
"content": {
"type": "bash_code_execution_result",
"stdout": "second\n",
"stderr": "",
"return_code": 0,
},
},
},
{"type": "content_block_stop", "index": 3},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
iterator = ModelResponseIterator(None, sync_stream=True)
# Collect each emission of code_interpreter_results
emissions = []
for chunk in chunks:
parsed = iterator.chunk_parser(chunk)
psf = None
if parsed.choices and parsed.choices[0].delta:
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
if psf and "code_interpreter_results" in psf:
emissions.append(psf["code_interpreter_results"])
# Should have 2 emissions (one per tool_result block)
assert len(emissions) == 2, f"Expected 2 emissions, got {len(emissions)}"
# First emission: cumulative list with 1 result
assert len(emissions[0]) == 1
assert emissions[0][0].id == "srvtoolu_01AAA"
assert emissions[0][0].code == "echo first"
assert emissions[0][0].outputs[0].logs == "first\n"
# Second (final) emission: cumulative list with BOTH results
# This is what stream_chunk_builder will pick as "last value wins"
assert len(emissions[1]) == 2, (
f"Expected final emission to have 2 results, got {len(emissions[1])}. "
f"IDs: {[r.id for r in emissions[1]]}"
)
assert emissions[1][0].id == "srvtoolu_01AAA"
assert emissions[1][0].code == "echo first"
assert emissions[1][0].outputs[0].logs == "first\n"
assert emissions[1][1].id == "srvtoolu_01BBB"
assert emissions[1][1].code == "echo second"
assert emissions[1][1].outputs[0].logs == "second\n"
def test_streaming_code_execution_input_assembled_from_deltas():
"""
In real Anthropic streaming, content_block_start for server_tool_use has
input: {}. The actual input arrives via input_json_delta deltas and must
be assembled at content_block_stop so the code field is populated.
This test uses realistic chunk shapes (empty input in start, partial JSON
in deltas) to exercise the input assembly path.
"""
chunks = [
{
"type": "message_start",
"message": {
"id": "msg_01XYZ",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 100, "output_tokens": 1},
},
},
# server_tool_use with empty input (real streaming behaviour)
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01AAA",
"name": "code_execution",
"input": {},
},
},
# Input arrives via deltas, split across two chunks
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "input_json_delta",
"partial_json": '{"comma',
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "input_json_delta",
"partial_json": 'nd": "echo hello"}',
},
},
{"type": "content_block_stop", "index": 0},
# Tool result
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_01AAA",
"content": {
"type": "bash_code_execution_result",
"stdout": "hello\n",
"stderr": "",
"return_code": 0,
},
},
},
{"type": "content_block_stop", "index": 1},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
iterator = ModelResponseIterator(None, sync_stream=True)
code_results = None
for chunk in chunks:
parsed = iterator.chunk_parser(chunk)
psf = None
if parsed.choices and parsed.choices[0].delta:
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
if psf and "code_interpreter_results" in psf:
code_results = psf["code_interpreter_results"]
# The code field must contain the assembled input, not be empty
assert code_results is not None, "No code_interpreter_results emitted"
assert len(code_results) == 1
assert code_results[0].id == "srvtoolu_01AAA"
assert code_results[0].code == "echo hello"
assert code_results[0].outputs[0].logs == "hello\n"
def test_empty_output_produces_null_outputs():
"""
When both stdout and stderr are empty, outputs should be None
(matching OpenAI's native behavior) rather than [{logs: ""}].
"""
chunks = [
{
"type": "message_start",
"message": {
"id": "msg_01XYZ",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 100, "output_tokens": 1},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01AAA",
"name": "bash_code_execution",
"input": {"command": "true"},
},
},
{"type": "content_block_stop", "index": 0},
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_01AAA",
"content": {
"type": "bash_code_execution_result",
"stdout": "",
"stderr": "",
"return_code": 0,
},
},
},
{"type": "content_block_stop", "index": 1},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
iterator = ModelResponseIterator(None, sync_stream=True)
code_results = None
for chunk in chunks:
parsed = iterator.chunk_parser(chunk)
psf = None
if parsed.choices and parsed.choices[0].delta:
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
if psf and "code_interpreter_results" in psf:
code_results = psf["code_interpreter_results"]
assert code_results is not None, "No code_interpreter_results emitted"
assert len(code_results) == 1
assert code_results[0].id == "srvtoolu_01AAA"
assert (
code_results[0].outputs is None
), f"Expected outputs=None for empty execution, got {code_results[0].outputs}"
def test_non_bash_tool_result_skipped():
"""
Tool result types other than bash_code_execution_tool_result (e.g.
text_editor_code_execution_tool_result) should be skipped and NOT
produce code_interpreter_call items.
"""
chunks = [
{
"type": "message_start",
"message": {
"id": "msg_01XYZ",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 100, "output_tokens": 1},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01AAA",
"name": "text_editor",
"input": {"command": "view", "path": "/tmp/test.py"},
},
},
{"type": "content_block_stop", "index": 0},
# text_editor result — should NOT become a code_interpreter_call
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "text_editor_code_execution_tool_result",
"tool_use_id": "srvtoolu_01AAA",
"content": [
{"type": "text", "text": "file contents here"},
],
},
},
{"type": "content_block_stop", "index": 1},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
iterator = ModelResponseIterator(None, sync_stream=True)
code_results = None
for chunk in chunks:
parsed = iterator.chunk_parser(chunk)
psf = None
if parsed.choices and parsed.choices[0].delta:
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
if psf and "code_interpreter_results" in psf:
code_results = psf["code_interpreter_results"]
# code_interpreter_results should be emitted but empty (no bash results)
assert (
code_results is not None
), "Expected code_interpreter_results key to be emitted"
assert (
len(code_results) == 0
), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}"

View file

@ -0,0 +1,268 @@
"""
Tests for the Responses API _extract_tool_result_output_items path,
the non-streaming _hidden_params propagation of code_interpreter_results,
and mock end-to-end streaming integration.
"""
from unittest.mock import MagicMock
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
from litellm.main import stream_chunk_builder
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
from litellm.types.responses.main import (
OutputCodeInterpreterCall,
OutputCodeInterpreterCallLog,
)
from litellm.types.utils import Choices, Message, ModelResponse
def _make_model_response(code_interpreter_results=None, provider_specific_fields=None):
"""Helper to build a ModelResponse with provider_specific_fields on the message."""
psf = provider_specific_fields or {}
if code_interpreter_results is not None:
psf["code_interpreter_results"] = code_interpreter_results
msg = Message(content="test", provider_specific_fields=psf if psf else None)
choice = Choices(index=0, message=msg, finish_reason="stop")
resp = ModelResponse()
resp.choices = [choice]
return resp
def test_extract_tool_result_output_items_from_pydantic_objects():
"""Non-streaming path: code_interpreter_results are Pydantic OutputCodeInterpreterCall objects."""
items = [
OutputCodeInterpreterCall(
type="code_interpreter_call",
id="srvtoolu_01AAA",
code="echo hello",
container_id=None,
status="completed",
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="hello\n")],
),
OutputCodeInterpreterCall(
type="code_interpreter_call",
id="srvtoolu_01BBB",
code="echo world",
container_id=None,
status="completed",
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="world\n")],
),
]
resp = _make_model_response(code_interpreter_results=items)
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
assert len(result) == 2
assert result[0].id == "srvtoolu_01AAA"
assert result[1].id == "srvtoolu_01BBB"
def test_extract_tool_result_output_items_from_dicts():
"""Streaming path: after model_dump(), code_interpreter_results are plain dicts.
_extract_tool_result_output_items reconstructs them as Pydantic objects."""
items = [
{
"type": "code_interpreter_call",
"id": "srvtoolu_01AAA",
"code": "echo hello",
"container_id": None,
"status": "completed",
"outputs": [{"type": "logs", "logs": "hello\n"}],
},
]
resp = _make_model_response(code_interpreter_results=items)
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
assert len(result) == 1
assert isinstance(result[0], OutputCodeInterpreterCall)
assert result[0].id == "srvtoolu_01AAA"
def test_extract_tool_result_output_items_empty():
"""No code_interpreter_results → empty list."""
resp = _make_model_response()
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
assert result == []
def test_extract_tool_result_output_items_no_provider_specific_fields():
"""Message with no provider_specific_fields → empty list."""
msg = Message(content="test")
choice = Choices(index=0, message=msg, finish_reason="stop")
resp = ModelResponse()
resp.choices = [choice]
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
assert result == []
def test_in_place_substitution_preserves_ordering():
"""
function_call items matching code_interpreter_results should be replaced
in-place, preserving the original output ordering.
Simulates: [message, function_call(exec1), function_call(regular), function_call(exec2)]
Expected: [message, code_interpreter_call(exec1), function_call(regular), code_interpreter_call(exec2)]
"""
code_results = [
OutputCodeInterpreterCall(
type="code_interpreter_call",
id="srvtoolu_01AAA",
code="echo first",
container_id=None,
status="completed",
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="first\n")],
),
OutputCodeInterpreterCall(
type="code_interpreter_call",
id="srvtoolu_01CCC",
code="echo third",
container_id=None,
status="completed",
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="third\n")],
),
]
resp = _make_model_response(code_interpreter_results=code_results)
# Build a mock responses_output list with interleaved items
class MockItem:
def __init__(self, type, call_id=None):
self.type = type
self.call_id = call_id
msg_item = MockItem(type="message")
fc_exec1 = MockItem(type="function_call", call_id="srvtoolu_01AAA")
fc_regular = MockItem(type="function_call", call_id="srvtoolu_01BBB")
fc_exec2 = MockItem(type="function_call", call_id="srvtoolu_01CCC")
responses_output = [msg_item, fc_exec1, fc_regular, fc_exec2]
# Apply the same logic as _transform_chat_completion_choices_to_responses_output
tool_result_items = (
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
)
if tool_result_items:
result_by_id = {
(item.get("id") if isinstance(item, dict) else item.id): item
for item in tool_result_items
}
replaced_ids = set(result_by_id.keys())
responses_output = [
(
result_by_id[getattr(item, "call_id", None)]
if (
getattr(item, "type", None) == "function_call"
and getattr(item, "call_id", None) in replaced_ids
)
else item
)
for item in responses_output
]
# Verify ordering: message, code_interpreter(AAA), function_call(BBB), code_interpreter(CCC)
assert len(responses_output) == 4
assert responses_output[0].type == "message"
assert responses_output[1].type == "code_interpreter_call"
assert responses_output[1].id == "srvtoolu_01AAA"
assert responses_output[2].type == "function_call"
assert responses_output[2].call_id == "srvtoolu_01BBB"
assert responses_output[3].type == "code_interpreter_call"
assert responses_output[3].id == "srvtoolu_01CCC"
def test_end_to_end_streaming_chunks_to_code_interpreter_output():
"""
Mock end-to-end test: Anthropic SSE chunks → ModelResponseIterator →
stream_chunk_builder → _extract_tool_result_output_items → final output
with code_interpreter_call items replacing function_call items.
This exercises the full streaming data flow without a live server.
"""
# Realistic Anthropic streaming chunks for a single code execution
raw_chunks = [
{
"type": "message_start",
"message": {
"id": "msg_01XYZ",
"type": "message",
"role": "assistant",
"content": [],
"usage": {"input_tokens": 100, "output_tokens": 1},
},
},
{
"type": "content_block_start",
"index": 0,
"content_block": {
"type": "server_tool_use",
"id": "srvtoolu_01AAA",
"name": "bash_code_execution",
"input": {},
},
},
{
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "input_json_delta",
"partial_json": '{"command": "echo e2e_test"}',
},
},
{"type": "content_block_stop", "index": 0},
{
"type": "content_block_start",
"index": 1,
"content_block": {
"type": "bash_code_execution_tool_result",
"tool_use_id": "srvtoolu_01AAA",
"content": {
"type": "bash_code_execution_result",
"stdout": "e2e_test\n",
"stderr": "",
"return_code": 0,
},
},
},
{"type": "content_block_stop", "index": 1},
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": 50},
},
]
# Step 1: Parse chunks through ModelResponseIterator (Anthropic handler)
iterator = ModelResponseIterator(None, sync_stream=True)
parsed_chunks = []
for chunk in raw_chunks:
parsed = iterator.chunk_parser(chunk)
d = parsed.model_dump()
# In production, CustomStreamWrapper sets the model on each chunk;
# stream_chunk_builder requires it.
d["model"] = "claude-sonnet-4-20250514"
parsed_chunks.append(d)
# Step 2: Assemble via stream_chunk_builder (simulates end-of-stream)
assembled = stream_chunk_builder(chunks=parsed_chunks)
assert assembled is not None
# Verify stream_chunk_builder picked up code_interpreter_results via last-value-wins
psf = assembled.choices[0].message.provider_specific_fields
assert psf is not None
assert "code_interpreter_results" in psf
code_results = psf["code_interpreter_results"]
assert len(code_results) == 1
# After model_dump + stream_chunk_builder, results are plain dicts
assert code_results[0]["id"] == "srvtoolu_01AAA"
assert code_results[0]["code"] == "echo e2e_test"
# Step 3: Extract via _extract_tool_result_output_items (Responses API layer)
tool_result_items = (
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(assembled)
)
assert len(tool_result_items) == 1
item = tool_result_items[0]
# Items are reconstructed as Pydantic OutputCodeInterpreterCall objects
assert isinstance(item, OutputCodeInterpreterCall)
assert item.type == "code_interpreter_call"
assert item.id == "srvtoolu_01AAA"
assert item.code == "echo e2e_test"
assert item.outputs[0].logs == "e2e_test\n"

View file

@ -1,20 +1,27 @@
"""
Tests for Anthropic OAuth token handling in common_utils.
Tests for Anthropic authentication and environment variable handling in common_utils.
Verifies that OAuth tokens (sk-ant-oat*) are sent via Authorization: Bearer
instead of x-api-key, per Anthropic's OAuth specification.
Verifies that:
- OAuth tokens (sk-ant-oat*) produce Authorization: Bearer headers with OAuth beta flags.
- Regular API keys produce x-api-key headers.
- ANTHROPIC_AUTH_TOKEN produces Authorization: Bearer headers,
matching the official Anthropic SDK behavior.
- ANTHROPIC_BASE_URL is used as a fallback for base URL resolution.
- ANTHROPIC_API_KEY / ANTHROPIC_API_BASE take precedence over their aliases.
"""
import os
import sys
from unittest.mock import patch
sys.path.insert(
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))
)
# Fake OAuth token for testing (not a real secret)
# Fake tokens for testing (not real secrets)
FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef"
FAKE_REGULAR_KEY = "sk-ant-api03-regular-key-for-testing-123456789"
FAKE_AUTH_TOKEN = "sk-ant-aut01-fake-auth-token-for-testing-123456789"
class TestOptionallyHandleAnthropicOAuth:
@ -697,3 +704,430 @@ class TestProxyOAuthHeaderForwarding:
assert cleaned["authorization"] == oauth_token
# Proxy key must be stripped
assert "x-litellm-api-key" not in cleaned
class TestGetAnthropicHeadersWithAuthToken:
"""Tests for get_anthropic_headers with auth_token parameter."""
def test_auth_token_uses_bearer_header(self):
"""auth_token should produce Authorization: Bearer header."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
headers = config.get_anthropic_headers(
api_key=None,
auth_token=FAKE_AUTH_TOKEN,
computer_tool_used=False,
prompt_caching_set=False,
pdf_used=False,
is_vertex_request=False,
)
assert headers["authorization"] == f"Bearer {FAKE_AUTH_TOKEN}"
assert "x-api-key" not in headers
# auth_token should NOT set OAuth-specific flags
assert "anthropic-dangerous-direct-browser-access" not in headers
def test_auth_token_includes_standard_headers(self):
"""auth_token path should include standard Anthropic headers."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
headers = config.get_anthropic_headers(
api_key=None,
auth_token=FAKE_AUTH_TOKEN,
computer_tool_used=False,
prompt_caching_set=False,
pdf_used=False,
is_vertex_request=False,
)
assert headers["anthropic-version"] == "2023-06-01"
assert headers["accept"] == "application/json"
assert headers["content-type"] == "application/json"
def test_api_key_takes_precedence_over_auth_token(self):
"""When both api_key and auth_token are provided, api_key wins."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
headers = config.get_anthropic_headers(
api_key=FAKE_REGULAR_KEY,
auth_token=FAKE_AUTH_TOKEN,
computer_tool_used=False,
prompt_caching_set=False,
pdf_used=False,
is_vertex_request=False,
)
assert headers["x-api-key"] == FAKE_REGULAR_KEY
assert "authorization" not in headers
class TestValidateEnvironmentAuthToken:
"""Tests for validate_environment with auth_token resolution."""
def test_auth_token_env_var_produces_bearer_header(self):
"""validate_environment should use Bearer auth when only ANTHROPIC_AUTH_TOKEN is set."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN},
clear=True,
):
headers = config.validate_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert headers["authorization"] == f"Bearer {FAKE_AUTH_TOKEN}"
assert "x-api-key" not in headers
assert "anthropic-dangerous-direct-browser-access" not in headers
def test_api_key_param_takes_precedence_over_auth_token_env_var(self):
"""validate_environment should prefer explicit api_key over ANTHROPIC_AUTH_TOKEN."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN},
clear=True,
):
headers = config.validate_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=FAKE_REGULAR_KEY,
api_base=None,
)
assert headers["x-api-key"] == FAKE_REGULAR_KEY
assert "authorization" not in headers
def test_raises_when_no_credentials(self):
"""validate_environment should raise when neither API key nor auth token is available."""
from unittest.mock import patch as mock_patch
import pytest
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
with mock_patch.dict("os.environ", {}, clear=True):
with pytest.raises(
Exception, match="ANTHROPIC_API_KEY.*ANTHROPIC_AUTH_TOKEN"
):
config.validate_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
def test_resolves_api_key_from_env_when_param_is_none(self):
"""validate_environment should resolve ANTHROPIC_API_KEY from env when api_key param is None."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
config = AnthropicModelInfo()
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY},
clear=True,
):
headers = config.validate_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert headers["x-api-key"] == FAKE_REGULAR_KEY
assert "authorization" not in headers
class TestGetAuthToken:
"""Tests for AnthropicModelInfo.get_auth_token() static method."""
def test_returns_env_var_value(self):
"""get_auth_token returns the ANTHROPIC_AUTH_TOKEN env var value."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ", {"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, clear=True
):
assert AnthropicModelInfo.get_auth_token() == FAKE_AUTH_TOKEN
def test_returns_none_when_not_set(self):
"""get_auth_token returns None when ANTHROPIC_AUTH_TOKEN is not set."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict("os.environ", {}, clear=True):
assert AnthropicModelInfo.get_auth_token() is None
def test_explicit_param_takes_precedence(self):
"""Explicit auth_token param takes precedence over env var."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
explicit_token = "sk-ant-aut01-explicit-token-override-123456789"
assert AnthropicModelInfo.get_auth_token(explicit_token) == explicit_token
class TestGetAuthHeader:
"""Tests for AnthropicModelInfo.get_auth_header() centralized helper."""
def test_returns_x_api_key_when_api_key_provided(self):
"""Explicit api_key param should return x-api-key header."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
result = AnthropicModelInfo.get_auth_header(api_key=FAKE_REGULAR_KEY)
assert result == {"x-api-key": FAKE_REGULAR_KEY}
def test_returns_x_api_key_from_env(self):
"""ANTHROPIC_API_KEY env var should return x-api-key header."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY},
clear=True,
):
result = AnthropicModelInfo.get_auth_header()
assert result == {"x-api-key": FAKE_REGULAR_KEY}
def test_returns_bearer_from_auth_token_env(self):
"""ANTHROPIC_AUTH_TOKEN env var should return Authorization: Bearer header."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN},
clear=True,
):
result = AnthropicModelInfo.get_auth_header()
assert result == {"authorization": f"Bearer {FAKE_AUTH_TOKEN}"}
def test_api_key_takes_precedence_over_auth_token(self):
"""ANTHROPIC_API_KEY should take precedence over ANTHROPIC_AUTH_TOKEN."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{
"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY,
"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN,
},
clear=True,
):
result = AnthropicModelInfo.get_auth_header()
assert result == {"x-api-key": FAKE_REGULAR_KEY}
def test_explicit_api_key_overrides_env_auth_token(self):
"""Explicit api_key param should override ANTHROPIC_AUTH_TOKEN env var."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN},
clear=True,
):
result = AnthropicModelInfo.get_auth_header(api_key=FAKE_REGULAR_KEY)
assert result == {"x-api-key": FAKE_REGULAR_KEY}
def test_returns_none_when_no_credentials(self):
"""Should return None when neither api_key nor auth_token is available."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict("os.environ", {}, clear=True):
result = AnthropicModelInfo.get_auth_header()
assert result is None
def test_oauth_token_uses_bearer_not_x_api_key(self):
"""OAuth token (sk-ant-oat*) should return Authorization: Bearer, not x-api-key."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
result = AnthropicModelInfo.get_auth_header(api_key=FAKE_OAUTH_TOKEN)
assert result == {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
def test_oauth_token_from_env_uses_bearer(self):
"""OAuth token in ANTHROPIC_API_KEY env var should return Authorization: Bearer."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_API_KEY": FAKE_OAUTH_TOKEN},
clear=True,
):
result = AnthropicModelInfo.get_auth_header()
assert result == {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"}
class TestGetApiBaseFallbackChain:
"""Tests for AnthropicModelInfo.get_api_base() fallback to ANTHROPIC_BASE_URL."""
def test_explicit_param_takes_precedence(self):
"""Explicit api_base param takes precedence over all env vars."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert (
AnthropicModelInfo.get_api_base("https://explicit.example.com")
== "https://explicit.example.com"
)
def test_defaults_to_anthropic_api(self):
"""get_api_base returns the default Anthropic API base when no env vars are set."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict("os.environ", {}, clear=True):
assert AnthropicModelInfo.get_api_base() == "https://api.anthropic.com"
def test_api_base_env_preferred_over_base_url_env(self):
"""ANTHROPIC_API_BASE takes precedence over ANTHROPIC_BASE_URL."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{
"ANTHROPIC_API_BASE": "https://api-base.example.com",
"ANTHROPIC_BASE_URL": "https://base-url.example.com",
},
clear=True,
):
assert AnthropicModelInfo.get_api_base() == "https://api-base.example.com"
def test_falls_back_to_base_url_env(self):
"""get_api_base falls back to ANTHROPIC_BASE_URL when ANTHROPIC_API_BASE is not set."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_BASE_URL": "https://base-url.example.com"},
clear=True,
):
assert AnthropicModelInfo.get_api_base() == "https://base-url.example.com"
class TestPassthroughAuthToken:
"""Tests for passthrough messages endpoint with ANTHROPIC_AUTH_TOKEN."""
def test_passthrough_auth_token_uses_bearer_header(self):
"""Passthrough endpoint should use Bearer auth when only ANTHROPIC_AUTH_TOKEN is set."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
config = AnthropicMessagesConfig()
with mock_patch.dict(
"os.environ", {"ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, clear=True
):
updated_headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert updated_headers["authorization"] == f"Bearer {FAKE_AUTH_TOKEN}"
assert "x-api-key" not in updated_headers
assert "anthropic-dangerous-direct-browser-access" not in updated_headers
def test_passthrough_api_key_takes_precedence(self):
"""Passthrough endpoint should prefer ANTHROPIC_API_KEY over ANTHROPIC_AUTH_TOKEN."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
config = AnthropicMessagesConfig()
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY, "ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN},
clear=True,
):
updated_headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="claude-sonnet-4-5-20250929",
messages=[{"role": "user", "content": "Hello"}],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY
assert "authorization" not in updated_headers
def test_passthrough_get_complete_url_honours_base_url_env(self):
"""get_complete_url should use ANTHROPIC_BASE_URL when api_base is None."""
from unittest.mock import patch as mock_patch
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
config = AnthropicMessagesConfig()
with mock_patch.dict(
"os.environ",
{"ANTHROPIC_BASE_URL": "https://custom.example.com"},
clear=True,
):
url = config.get_complete_url(
api_base=None,
api_key=FAKE_REGULAR_KEY,
model="claude-sonnet-4-5-20250929",
optional_params={},
litellm_params={},
)
assert url == "https://custom.example.com/v1/messages"

View file

@ -1458,10 +1458,6 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
return_value=mock_key_record
)
# Mock get_key_object and _cache_key_object functions
mock_key_object = MagicMock()
mock_key_object.blocked = True # Initially blocked
# Mock hash_token function
def mock_hash_token(token):
if token == "sk-test123456789":
@ -1482,19 +1478,12 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
) # Disable audit logs for simpler test
# Mock get_key_object and _cache_key_object
async def mock_get_key_object(**kwargs):
return mock_key_object
async def mock_cache_key_object(**kwargs):
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_key_object",
mock_get_key_object,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
mock_cache_key_object,
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,
)
# Create mock request and user auth
@ -1519,11 +1508,9 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
)
assert result == mock_key_record
assert mock_key_object.blocked == False # Should be updated to unblocked
# Reset mocks for second test
mock_prisma_client.db.litellm_verificationtoken.update.reset_mock()
mock_key_object.blocked = True # Reset to blocked state
# Test Case 2: Using already hashed token
hashed_token_request = BlockKeyRequest(key=test_hashed_token)
@ -1541,7 +1528,6 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
)
assert result == mock_key_record
assert mock_key_object.blocked == False # Should be updated to unblocked
@pytest.mark.asyncio
@ -1579,6 +1565,249 @@ async def test_unblock_key_invalid_key_format(monkeypatch):
assert "Invalid key format" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_block_key_nonexistent_key_returns_404(monkeypatch):
"""
Test that block_key returns 404 (not misleading 401) when the key
doesn't exist in the database, even when the caller is authenticated
as a proxy admin.
Previously, block_key would call get_key_object() for cache refresh,
which raised a 401 ProxyException with 'Authentication Error' — making
it look like an auth failure when it was really a missing-key error.
"""
from litellm.proxy._types import BlockKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import block_key
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
# find_unique returns None → key does not exist
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
def mock_hash_token(token):
return "abcd1234" * 8 # 64-char hex
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = BlockKeyRequest(key="sk-does-not-exist-key")
with pytest.raises(ProxyException) as exc_info:
await block_key(
data=data,
http_request=mock_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert exc_info.value.code == "404"
assert "not found" in str(exc_info.value.message).lower()
# Must NOT contain "Authentication Error"
assert "Authentication Error" not in str(exc_info.value.message)
# update should never be called since the key doesn't exist
mock_prisma_client.db.litellm_verificationtoken.update.assert_not_called()
@pytest.mark.asyncio
async def test_unblock_key_nonexistent_key_returns_404(monkeypatch):
"""
Test that unblock_key returns 404 (not misleading 401) when the key
doesn't exist in the database.
"""
from litellm.proxy._types import BlockKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
unblock_key,
)
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
# find_unique returns None → key does not exist
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
def mock_hash_token(token):
return "abcd1234" * 8
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = BlockKeyRequest(key="sk-does-not-exist-key")
with pytest.raises(ProxyException) as exc_info:
await unblock_key(
data=data,
http_request=mock_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert exc_info.value.code == "404"
assert "not found" in str(exc_info.value.message).lower()
assert "Authentication Error" not in str(exc_info.value.message)
mock_prisma_client.db.litellm_verificationtoken.update.assert_not_called()
@pytest.mark.asyncio
async def test_update_key_nonexistent_key_returns_404(monkeypatch):
"""
Test that update_key_fn returns 404 (not misleading 401) when the body
key doesn't exist in the database, even when the caller is authenticated
as a proxy admin via the Authorization header.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
# find_unique returns None → key does not exist
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = UpdateKeyRequest(key="sk-does-not-exist-key")
with pytest.raises(ProxyException) as exc_info:
await update_key_fn(
request=mock_request,
data=data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert exc_info.value.code == "404"
assert "not found" in str(exc_info.value.message).lower()
assert "Authentication Error" not in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_block_key_existing_key_succeeds(monkeypatch):
"""
Test that block_key successfully blocks an existing key and
invalidates the cache entry.
"""
from litellm.proxy._types import BlockKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import block_key
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
test_hashed_token = "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
mock_key_record = MagicMock()
mock_key_record.token = test_hashed_token
mock_key_record.blocked = False
mock_key_record.model_dump_json.return_value = (
f'{{"token": "{test_hashed_token}", "blocked": false}}'
)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key_record
)
mock_updated_record = MagicMock()
mock_updated_record.token = test_hashed_token
mock_updated_record.blocked = True
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
return_value=mock_updated_record
)
def mock_hash_token(token):
if token.startswith("sk-"):
return test_hashed_token
return token
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
# Mock _delete_cache_key_object
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,
)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = BlockKeyRequest(key="sk-test123456789")
result = await block_key(
data=data,
http_request=mock_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify the key was found and updated
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once_with(
where={"token": test_hashed_token}
)
mock_prisma_client.db.litellm_verificationtoken.update.assert_called_once_with(
where={"token": test_hashed_token}, data={"blocked": True}
)
assert result == mock_updated_record
@pytest.mark.asyncio
async def test_validate_key_team_change_with_member_permissions():
"""
@ -4871,14 +5100,16 @@ async def test_validate_max_budget():
async def test_get_and_validate_existing_key():
"""
Test _get_and_validate_existing_key helper function.
Tests:
1. Successfully retrieve existing key
2. Key not found raises HTTPException
2. Key not found raises ProxyException
3. Database not connected raises HTTPException
"""
from fastapi import HTTPException
from litellm.proxy._types import ProxyException
# Test Case 1: Successfully retrieve existing key
mock_prisma_client = AsyncMock()
mock_key = LiteLLM_VerificationToken(
@ -4887,39 +5118,49 @@ async def test_get_and_validate_existing_key():
models=["gpt-4"],
team_id=None,
)
mock_prisma_client.get_data = AsyncMock(return_value=mock_key)
result = await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=mock_prisma_client,
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key
)
assert result == mock_key
mock_prisma_client.get_data.assert_called_once_with(
token="test-key-123",
table_name="key",
query_type="find_unique",
)
# Test Case 2: Key not found raises HTTPException
mock_prisma_client.get_data = AsyncMock(return_value=None)
with pytest.raises(HTTPException) as exc_info:
await _get_and_validate_existing_key(
token="non-existent-key",
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
return_value="hashed-test-key-123",
):
result = await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=mock_prisma_client,
)
assert exc_info.value.status_code == 404
assert "Key not found" in str(exc_info.value.detail)
assert result == mock_key
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once_with(
where={"token": "hashed-test-key-123"}
)
# Test Case 2: Key not found raises ProxyException
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
return_value="hashed-non-existent-key",
):
with pytest.raises(ProxyException) as exc_info:
await _get_and_validate_existing_key(
token="non-existent-key",
prisma_client=mock_prisma_client,
)
assert str(exc_info.value.code) == "404"
assert "Key not found" in exc_info.value.message
# Test Case 3: Database not connected raises HTTPException
with pytest.raises(HTTPException) as exc_info:
await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=None,
)
assert exc_info.value.status_code == 500
assert "Database not connected" in str(exc_info.value.detail)
@ -4960,75 +5201,82 @@ async def test_process_single_key_update():
"tags": ["production"],
}
mock_prisma_client.get_data = AsyncMock(return_value=existing_key)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=existing_key
)
mock_updated_key_obj = MagicMock()
mock_updated_key_obj.model_dump.return_value = updated_key_data
mock_prisma_client.update_data = AsyncMock(
return_value={"data": mock_updated_key_obj}
)
# Mock prepare_key_update_data
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
# Mock TeamMemberPermissionChecks
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
) as mock_permission_check:
mock_permission_check.return_value = None
# Mock _delete_cache_key_object
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache:
mock_delete_cache.return_value = None
# Mock hash_token (imported from litellm.proxy._types)
with patch(
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.return_value = "hashed-test-key-123"
# Mock KeyManagementEventHooks
# Mock _hash_token_if_needed
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
return_value="hashed-test-key-123",
):
# Create update request
key_update_item = BulkUpdateKeyRequestItem(
key="test-key-123",
max_budget=100.0,
tags=["production"],
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call the function
result = await _process_single_key_update(
key_update_item=key_update_item,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
llm_router=mock_llm_router,
)
# Verify results
assert result is not None
assert "token" not in result # Token should be removed
assert result.get("max_budget") == 100.0
assert result.get("tags") == ["production"]
# Verify mocks were called
mock_prisma_client.get_data.assert_called_once()
mock_prisma_client.update_data.assert_called_once()
mock_delete_cache.assert_called_once()
# Mock KeyManagementEventHooks
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create update request
key_update_item = BulkUpdateKeyRequestItem(
key="test-key-123",
max_budget=100.0,
tags=["production"],
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call the function
result = await _process_single_key_update(
key_update_item=key_update_item,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
llm_router=mock_llm_router,
)
# Verify results
assert result is not None
assert "token" not in result # Token should be removed
assert result.get("max_budget") == 100.0
assert result.get("tags") == ["production"]
# Verify mocks were called
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once()
mock_prisma_client.update_data.assert_called_once()
mock_delete_cache.assert_called_once()
@pytest.mark.asyncio
@ -5090,7 +5338,7 @@ async def test_bulk_update_keys_success(monkeypatch):
"tags": ["staging"],
}
mock_prisma_client.get_data = AsyncMock(
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
side_effect=[existing_key_1, existing_key_2]
)
mock_updated_key_1_obj = MagicMock()
@ -5103,7 +5351,7 @@ async def test_bulk_update_keys_success(monkeypatch):
{"data": mock_updated_key_2_obj},
]
)
# Patch dependencies
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
@ -5115,7 +5363,7 @@ async def test_bulk_update_keys_success(monkeypatch):
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
# Mock helper functions
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
@ -5124,7 +5372,7 @@ async def test_bulk_update_keys_success(monkeypatch):
{"max_budget": 100.0, "tags": ["production"]},
{"max_budget": 200.0, "tags": ["staging"]},
]
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
):
@ -5135,45 +5383,49 @@ async def test_bulk_update_keys_success(monkeypatch):
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.side_effect = ["hashed-key-1", "hashed-key-2"]
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
side_effect=["hashed-key-1", "hashed-key-2"],
):
# Create request
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="test-key-2",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 2
assert len(response.failed_updates) == 0
assert response.successful_updates[0].key == "test-key-1"
assert response.successful_updates[1].key == "test-key-2"
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create request
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="test-key-2",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 2
assert len(response.failed_updates) == 0
assert response.successful_updates[0].key == "test-key-1"
assert response.successful_updates[1].key == "test-key-2"
@pytest.mark.asyncio
@ -5218,7 +5470,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
}
# First key exists, second key doesn't exist
mock_prisma_client.get_data = AsyncMock(
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
side_effect=[existing_key_1, None] # Second key not found
)
mock_updated_key_1_obj = MagicMock()
@ -5226,7 +5478,9 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
mock_prisma_client.update_data = AsyncMock(
return_value={"data": mock_updated_key_1_obj}
)
# Mock get_data for the error handler path (used to fetch key_info on failure)
mock_prisma_client.get_data = AsyncMock(return_value=None)
# Patch dependencies
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
@ -5238,13 +5492,13 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
# Mock helper functions
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
):
@ -5255,46 +5509,50 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.return_value = "hashed-key-1"
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
side_effect=["hashed-key-1", "hashed-non-existent-key"],
):
# Create request with one valid and one invalid key
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="non-existent-key",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 1
assert len(response.failed_updates) == 1
assert response.successful_updates[0].key == "test-key-1"
assert response.failed_updates[0].key == "non-existent-key"
assert "Key not found" in response.failed_updates[0].failed_reason
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create request with one valid and one invalid key
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="non-existent-key",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 1
assert len(response.failed_updates) == 1
assert response.successful_updates[0].key == "test-key-1"
assert response.failed_updates[0].key == "non-existent-key"
assert "Key not found" in response.failed_updates[0].failed_reason
@pytest.mark.parametrize(
@ -7379,19 +7637,12 @@ def _setup_block_unblock_mocks(monkeypatch, mock_key_team_id=None):
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
async def mock_get_key_object(**kwargs):
return mock_key_object
async def mock_cache_key_object(**kwargs):
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_key_object",
mock_get_key_object,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
mock_cache_key_object,
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,
)
return mock_prisma_client, test_hashed_token
@ -7638,16 +7889,9 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
async def mock_cache_key_object(**kwargs):
pass
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
mock_cache_key_object,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,

View file

@ -280,6 +280,47 @@ class TestProxyInitializationHelpers:
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
mock_uvicorn_run.assert_called_once()
@patch("uvicorn.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
@patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False)
def test_proxy_default_api_version_uses_azure_default(
self, mock_should_update, mock_setup_db, mock_atexit_register, mock_uvicorn_run
):
"""Proxy default api_version should match litellm.AZURE_DEFAULT_API_VERSION for consistency."""
from click.testing import CliRunner
import litellm
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=MagicMock(),
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")}
with patch.dict(os.environ, clean_env, clear=True), patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
), patch(
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
) as mock_get_args:
mock_get_args.return_value = {
"app": "litellm.proxy.proxy_server:app",
"host": "localhost",
"port": 8000,
}
result = runner.invoke(run_server, ["--local", "--skip_server_startup"])
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
mock_proxy_module.save_worker_config.assert_called_once()
call_kwargs = mock_proxy_module.save_worker_config.call_args[1]
assert call_kwargs["api_version"] == litellm.AZURE_DEFAULT_API_VERSION
@patch("uvicorn.run")
@patch("builtins.print")
def test_keepalive_timeout_flag(self, mock_print, mock_uvicorn_run):

View file

@ -0,0 +1,146 @@
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import { flexRender, getCoreRowModel, useReactTable } from "@tanstack/react-table";
import { getAgentHubTableColumns, AgentHubData } from "./AgentHubTableColumns";
const mockAgent: AgentHubData = {
agent_id: "agent-1",
protocolVersion: "1.0",
name: "Test Agent",
description: "A test agent for unit testing",
url: "https://agent.example.com",
version: "2.0",
capabilities: { streaming: true, caching: false },
defaultInputModes: ["text"],
defaultOutputModes: ["text", "image"],
skills: [
{ id: "s1", name: "Skill One", description: "First skill" },
{ id: "s2", name: "Skill Two", description: "Second skill" },
{ id: "s3", name: "Skill Three", description: "Third skill" },
],
is_public: true,
};
function TestTable({
data,
publicPage = false,
showModal = vi.fn(),
copyToClipboard = vi.fn(),
}: {
data: AgentHubData[];
publicPage?: boolean;
showModal?: ReturnType<typeof vi.fn>;
copyToClipboard?: ReturnType<typeof vi.fn>;
}) {
const columns = getAgentHubTableColumns(showModal, copyToClipboard, publicPage);
const table = useReactTable({ data, columns, getCoreRowModel: getCoreRowModel() });
return (
<table>
<thead>
{table.getHeaderGroups().map((hg) => (
<tr key={hg.id}>
{hg.headers.map((h) => (
<th key={h.id}>{flexRender(h.column.columnDef.header, h.getContext())}</th>
))}
</tr>
))}
</thead>
<tbody>
{table.getRowModel().rows.map((row) => (
<tr key={row.id}>
{row.getVisibleCells().map((cell) => (
<td key={cell.id}>{flexRender(cell.column.columnDef.cell, cell.getContext())}</td>
))}
</tr>
))}
</tbody>
</table>
);
}
describe("AgentHubTableColumns", () => {
it("should render", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("Test Agent")).toBeInTheDocument();
});
it("should display the agent description", () => {
render(<TestTable data={[mockAgent]} />);
// Description appears in both the description column and the mobile view within agent name column
expect(screen.getAllByText("A test agent for unit testing").length).toBeGreaterThanOrEqual(1);
});
it("should display the version with a 'v' prefix", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("v2.0")).toBeInTheDocument();
});
it("should display the protocol version", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("1.0")).toBeInTheDocument();
});
it("should show skill count with correct pluralization", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("3 skills")).toBeInTheDocument();
});
it("should show first two skills and '+1' for overflow", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("Skill One")).toBeInTheDocument();
expect(screen.getByText("Skill Two")).toBeInTheDocument();
expect(screen.getByText("+1")).toBeInTheDocument();
});
it("should show only true capabilities as badges", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("streaming")).toBeInTheDocument();
expect(screen.queryByText("caching")).not.toBeInTheDocument();
});
it("should display I/O modes", () => {
render(<TestTable data={[mockAgent]} />);
// "In:" and "Out:" are in <span> children; getByText with exact:false
// matches against the element's full textContent across child nodes
expect(screen.getByText((_, el) =>
el?.tagName === "P" && el.textContent === "In: text"
)).toBeInTheDocument();
expect(screen.getByText((_, el) =>
el?.tagName === "P" && el.textContent === "Out: text, image"
)).toBeInTheDocument();
});
it("should display 'Yes' badge for public agents", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("Yes")).toBeInTheDocument();
});
it("should display 'No' badge for non-public agents", () => {
const privateAgent = { ...mockAgent, is_public: false };
render(<TestTable data={[privateAgent]} />);
expect(screen.getByText("No")).toBeInTheDocument();
});
it("should display a Details button", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByRole("button", { name: /details|info/i })).toBeInTheDocument();
});
it("should show '-' when agent has no capabilities", () => {
const noCapAgent = { ...mockAgent, capabilities: {} };
render(<TestTable data={[noCapAgent]} />);
// The dash is rendered in the capabilities column
expect(screen.getByText("-")).toBeInTheDocument();
});
it("should show singular 'skill' for one skill", () => {
const oneSkillAgent = {
...mockAgent,
skills: [{ id: "s1", name: "Only Skill", description: "One" }],
};
render(<TestTable data={[oneSkillAgent]} />);
expect(screen.getByText("1 skill")).toBeInTheDocument();
});
});

View file

@ -194,7 +194,6 @@ export const getAgentHubTableColumns = (
return publicA - publicB;
},
cell: ({ row }) => {
console.log(`CHECKPOINT 1: ${JSON.stringify(row.original)}`);
const agent = row.original;
return agent.is_public === true ? (

View file

@ -0,0 +1,73 @@
import { renderWithProviders, screen } from "../../../tests/test-utils";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import UsageExportHeader from "./UsageExportHeader";
import type { EntitySpendData } from "./types";
vi.mock("./EntityUsageExportModal", () => ({
default: ({ isOpen, onClose }: { isOpen: boolean; onClose: () => void }) =>
isOpen ? (
<div data-testid="export-modal">
<button onClick={onClose}>Close</button>
</div>
) : null,
}));
const defaultProps = {
dateValue: { from: new Date("2025-01-01"), to: new Date("2025-01-31") },
entityType: "team" as const,
spendData: {
results: [],
metadata: {
total_spend: 0,
total_api_requests: 0,
total_successful_requests: 0,
total_failed_requests: 0,
total_tokens: 0,
},
} satisfies EntitySpendData,
};
describe("UsageExportHeader", () => {
it("should render", () => {
renderWithProviders(<UsageExportHeader {...defaultProps} />);
expect(screen.getByRole("button", { name: /export data/i })).toBeInTheDocument();
});
it("should open the export modal when the export button is clicked", async () => {
const user = userEvent.setup();
renderWithProviders(<UsageExportHeader {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /export data/i }));
expect(screen.getByTestId("export-modal")).toBeInTheDocument();
});
it("should close the export modal when onClose is called", async () => {
const user = userEvent.setup();
renderWithProviders(<UsageExportHeader {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /export data/i }));
await user.click(screen.getByRole("button", { name: /close/i }));
expect(screen.queryByTestId("export-modal")).not.toBeInTheDocument();
});
it("should not show filter dropdown when showFilters is false", () => {
renderWithProviders(<UsageExportHeader {...defaultProps} showFilters={false} />);
expect(screen.queryByText(/filter/i)).not.toBeInTheDocument();
});
it("should show filter dropdown when showFilters is true and options provided", () => {
renderWithProviders(
<UsageExportHeader
{...defaultProps}
showFilters
filterLabel="Team"
filterPlaceholder="Select teams"
filterOptions={[
{ label: "Team A", value: "team-a" },
{ label: "Team B", value: "team-b" },
]}
onFiltersChange={vi.fn()}
/>,
);
expect(screen.getByText("Team")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,98 @@
import { render, screen, act } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import { GuardrailConfig } from "./GuardrailConfig";
describe("GuardrailConfig", () => {
const defaultProps = {
guardrailName: "Content Safety",
guardrailType: "Content Safety",
provider: "bedrock",
};
afterEach(() => {
vi.useRealTimers();
});
it("should render", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByText("Parameters")).toBeInTheDocument();
});
it("should display the guardrail name in the parameters description", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByText(/Configure Content Safety behavior/)).toBeInTheDocument();
});
// Note: Version history entries are hardcoded placeholders in the component.
// These assertions will need updating when wired to real API data.
it("should show version history when 'View history' is clicked", async () => {
const user = userEvent.setup();
render(<GuardrailConfig {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /view history/i }));
expect(screen.getByText("Initial configuration")).toBeInTheDocument();
expect(screen.getByText("Added custom categories list")).toBeInTheDocument();
});
it("should toggle version history text between View/Hide", async () => {
const user = userEvent.setup();
render(<GuardrailConfig {...defaultProps} />);
const button = screen.getByRole("button", { name: /view history/i });
await user.click(button);
expect(screen.getByRole("button", { name: /hide history/i })).toBeInTheDocument();
});
it("should show custom code textarea when custom code override is toggled on", async () => {
const user = userEvent.setup();
render(<GuardrailConfig {...defaultProps} />);
// Walk up from "Custom Code Override" heading to find the enclosing section,
// then locate the switch within it
const heading = screen.getByText("Custom Code Override");
let container = heading.parentElement;
let customCodeSwitch: Element | null = null;
while (container && !customCodeSwitch) {
customCodeSwitch = container.querySelector('[role="switch"]');
container = container.parentElement;
}
if (!customCodeSwitch) {
throw new Error("Could not find the Custom Code Override switch via DOM traversal");
}
await user.click(customCodeSwitch);
expect(screen.getByPlaceholderText(/async def evaluate/)).toBeInTheDocument();
});
it("should hide custom code textarea when custom code override is off", () => {
render(<GuardrailConfig {...defaultProps} />);
// There's an input for categories, but no textarea
expect(screen.queryByPlaceholderText(/async def evaluate/)).not.toBeInTheDocument();
});
it("should show the re-run button in idle state", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByRole("button", { name: /re-run on failing logs/i })).toBeInTheDocument();
});
it("should show loading state when re-run is clicked", async () => {
vi.useFakeTimers({ shouldAdvanceTime: true });
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
render(<GuardrailConfig {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /re-run on failing logs/i }));
expect(screen.getByText(/Running on 10 samples/)).toBeInTheDocument();
});
it("should show success message after re-run completes", async () => {
vi.useFakeTimers({ shouldAdvanceTime: true });
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
render(<GuardrailConfig {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /re-run on failing logs/i }));
await act(async () => { vi.advanceTimersByTime(2500); });
expect(screen.getByText(/7\/10 would now pass/)).toBeInTheDocument();
});
it("should display the Revert and Save buttons", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByRole("button", { name: /revert/i })).toBeInTheDocument();
// The component's hardcoded default version is "v3", so Save shows "v4"
expect(screen.getByRole("button", { name: /save as v\d+/i })).toBeInTheDocument();
});
});

View file

@ -138,4 +138,18 @@ describe("DocsMenu", () => {
await user.click(button);
expect(button).toHaveAttribute("aria-expanded", "true");
});
it("should close menu when clicking outside", async () => {
const user = userEvent.setup();
renderWithProviders(
<div>
<DocsMenu items={items} />
<button>Outside</button>
</div>,
);
await user.click(screen.getByRole("button", { name: /docs/i }));
expect(screen.getByText("Custom pricing")).toBeInTheDocument();
await user.click(screen.getByRole("button", { name: /outside/i }));
expect(screen.queryByText("Custom pricing")).not.toBeInTheDocument();
});
});

View file

@ -26,6 +26,10 @@ const PERMISSION_OPTIONS = [
"/key/unblock",
"/key/bulk_update",
"/key/{key_id}/reset_spend",
"/key/info",
"/key/list",
"/key/aliases",
"/team/daily/activity",
];
interface SettingRowProps {

View file

@ -13,6 +13,7 @@ import {
CreditCardOutlined,
DatabaseOutlined,
ExperimentOutlined,
ExportOutlined,
FileTextOutlined,
FolderOutlined,
KeyOutlined,
@ -400,7 +401,7 @@ const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapse
onClick={(e) => e.stopPropagation()}
style={{ color: "inherit", textDecoration: "none" }}
>
{label}
{label} <ExportOutlined style={{ fontSize: 10, marginLeft: 4 }} />
</a>
);
}

View file

@ -0,0 +1,296 @@
import { render, screen } from "@testing-library/react";
import { describe, it, expect, vi } from "vitest";
import ChatMessageBubble from "./ChatMessageBubble";
import { EndpointType } from "./mode_endpoint_mapping";
import { MessageType } from "./types";
// Mock child components to isolate bubble rendering logic
vi.mock("react-markdown", () => ({
default: ({ children }: { children: string }) => <div data-testid="react-markdown">{children}</div>,
}));
vi.mock("react-syntax-highlighter", () => ({
Prism: ({ children }: { children: string }) => <pre data-testid="syntax-highlighter">{children}</pre>,
}));
vi.mock("react-syntax-highlighter/dist/esm/styles/prism", () => ({
coy: {},
}));
vi.mock("./ReasoningContent", () => ({
default: ({ reasoningContent }: { reasoningContent: string }) => (
<div data-testid="reasoning-content">{reasoningContent}</div>
),
}));
vi.mock("./MCPEventsDisplay", () => ({
default: ({ events }: { events: unknown[] }) => (
<div data-testid="mcp-events-display">{events.length} events</div>
),
}));
vi.mock("./SearchResultsDisplay", () => ({
SearchResultsDisplay: ({ searchResults }: { searchResults: unknown[] }) => (
<div data-testid="search-results-display">{searchResults.length} results</div>
),
}));
vi.mock("./ResponseMetrics", () => ({
default: ({ timeToFirstToken }: { timeToFirstToken?: number }) => (
<div data-testid="response-metrics">TTFT: {timeToFirstToken}</div>
),
}));
vi.mock("./A2AMetrics", () => ({
default: ({ a2aMetadata }: { a2aMetadata: unknown }) => (
<div data-testid="a2a-metrics">A2A</div>
),
}));
vi.mock("./CodeInterpreterOutput", () => ({
default: ({ code }: { code: string }) => <div data-testid="code-interpreter-output">{code}</div>,
}));
vi.mock("./AudioRenderer", () => ({
default: ({ message }: { message: MessageType }) => (
<div data-testid="audio-renderer">{typeof message.content === "string" ? message.content : ""}</div>
),
}));
vi.mock("./ResponsesImageRenderer", () => ({
default: () => <div data-testid="responses-image-renderer" />,
}));
vi.mock("./ChatImageRenderer", () => ({
default: () => <div data-testid="chat-image-renderer" />,
}));
const defaultProps = {
isLastMessage: false,
endpointType: EndpointType.CHAT,
mcpEvents: [],
codeInterpreterResult: null,
accessToken: "test-token",
};
describe("ChatMessageBubble", () => {
it("should render a user message with right-aligned text", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "user", content: "Hello" }}
/>,
);
expect(screen.getByText("user")).toBeInTheDocument();
expect(screen.getByText("Hello")).toBeInTheDocument();
});
it("should render an assistant message with left-aligned text", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "Hi there" }}
/>,
);
expect(screen.getByText("assistant")).toBeInTheDocument();
expect(screen.getByText("Hi there")).toBeInTheDocument();
});
it("should show model badge for assistant messages when model is provided", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "Reply", model: "gpt-4" }}
/>,
);
expect(screen.getByText("gpt-4")).toBeInTheDocument();
});
it("should not show model badge for user messages even when model is set", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "user", content: "Hello", model: "gpt-4" }}
/>,
);
expect(screen.queryByText("gpt-4")).not.toBeInTheDocument();
});
it("should render markdown content via ReactMarkdown", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "**bold text**" }}
/>,
);
expect(screen.getByTestId("react-markdown")).toHaveTextContent("**bold text**");
});
it("should render an image when isImage is true", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "https://example.com/img.png", isImage: true }}
/>,
);
expect(screen.getByAltText("Generated image")).toHaveAttribute("src", "https://example.com/img.png");
});
it("should render AudioRenderer when isAudio is true", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "audio-url", isAudio: true }}
/>,
);
expect(screen.getByTestId("audio-renderer")).toBeInTheDocument();
});
it("should show ReasoningContent when reasoningContent is present", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "answer", reasoningContent: "thinking..." }}
/>,
);
expect(screen.getByTestId("reasoning-content")).toHaveTextContent("thinking...");
});
it("should show MCP events on the last assistant message for RESPONSES endpoint", () => {
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={true}
endpointType={EndpointType.RESPONSES}
mcpEvents={mcpEvents as any}
message={{ role: "assistant", content: "response" }}
/>,
);
expect(screen.getByTestId("mcp-events-display")).toHaveTextContent("1 events");
});
it("should show MCP events on the last assistant message for CHAT endpoint", () => {
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={true}
endpointType={EndpointType.CHAT}
mcpEvents={mcpEvents as any}
message={{ role: "assistant", content: "response" }}
/>,
);
expect(screen.getByTestId("mcp-events-display")).toHaveTextContent("1 events");
});
it("should not show MCP events when isLastMessage is false", () => {
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={false}
endpointType={EndpointType.RESPONSES}
mcpEvents={mcpEvents as any}
message={{ role: "assistant", content: "response" }}
/>,
);
expect(screen.queryByTestId("mcp-events-display")).not.toBeInTheDocument();
});
it("should show SearchResultsDisplay when searchResults are present", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "found results",
searchResults: [{ object: "search", search_query: "q", data: [] }],
}}
/>,
);
expect(screen.getByTestId("search-results-display")).toBeInTheDocument();
});
it("should show ResponseMetrics when usage data is present and no a2aMetadata", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "response",
timeToFirstToken: 150,
usage: { completionTokens: 10, promptTokens: 5, totalTokens: 15 },
}}
/>,
);
expect(screen.getByTestId("response-metrics")).toBeInTheDocument();
});
it("should show A2AMetrics when a2aMetadata is present instead of ResponseMetrics", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "agent response",
timeToFirstToken: 100,
a2aMetadata: { taskId: "task-1", status: { state: "completed" } },
}}
/>,
);
expect(screen.getByTestId("a2a-metrics")).toBeInTheDocument();
expect(screen.queryByTestId("response-metrics")).not.toBeInTheDocument();
});
it("should show CodeInterpreterOutput on the last assistant message for RESPONSES endpoint", () => {
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={true}
endpointType={EndpointType.RESPONSES}
codeInterpreterResult={{
code: "print('hello')",
containerId: "container-1",
annotations: [],
}}
message={{ role: "assistant", content: "result" }}
/>,
);
expect(screen.getByTestId("code-interpreter-output")).toHaveTextContent("print('hello')");
});
it("should render generated image from chat completions via message.image", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "Here is your image",
image: { url: "https://example.com/generated.png", detail: "auto" },
}}
/>,
);
const images = screen.getAllByAltText("Generated image");
expect(images.some((img) => img.getAttribute("src") === "https://example.com/generated.png")).toBe(true);
});
});

View file

@ -0,0 +1,214 @@
import { RobotOutlined, UserOutlined } from "@ant-design/icons";
import React from "react";
import ReactMarkdown from "react-markdown";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { coy } from "react-syntax-highlighter/dist/esm/styles/prism";
import { CodeInterpreterResult } from "../llm_calls/code_interpreter_handler";
import A2AMetrics from "./A2AMetrics";
import AudioRenderer from "./AudioRenderer";
import ChatImageRenderer from "./ChatImageRenderer";
import CodeInterpreterOutput from "./CodeInterpreterOutput";
import { EndpointType } from "./mode_endpoint_mapping";
import MCPEventsDisplay from "./MCPEventsDisplay";
import type { MCPEvent } from "../../mcp_tools/types";
import ReasoningContent from "./ReasoningContent";
import ResponseMetrics from "./ResponseMetrics";
import ResponsesImageRenderer from "./ResponsesImageRenderer";
import { SearchResultsDisplay } from "./SearchResultsDisplay";
import { MessageType } from "./types";
interface ChatMessageBubbleProps {
message: MessageType;
/** Whether this is the last message in the chat history. */
isLastMessage: boolean;
endpointType: EndpointType;
/** MCP events to display on the last assistant message. */
mcpEvents: MCPEvent[];
/** Code interpreter result to display on the last assistant message. */
codeInterpreterResult: CodeInterpreterResult | null;
/** API key used to fetch code interpreter file downloads. */
accessToken: string;
}
function ChatMessageBubble({
message,
isLastMessage,
endpointType,
mcpEvents,
codeInterpreterResult,
accessToken,
}: ChatMessageBubbleProps) {
const isUser = message.role === "user";
return (
<div className={`mb-4 ${isUser ? "text-right" : "text-left"}`}>
<div
className="inline-block max-w-[80%] rounded-lg shadow-sm p-3.5 px-4"
style={{
backgroundColor: isUser ? "#f0f8ff" : "#ffffff",
border: isUser ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
textAlign: "left",
}}
>
{/* Header: role icon + name + model badge */}
<div className="flex items-center gap-2 mb-1.5">
<div
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
style={{
backgroundColor: isUser ? "#e6f0fa" : "#f5f5f5",
}}
>
{isUser ? (
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
) : (
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
)}
</div>
<strong className="text-sm capitalize">{message.role}</strong>
{message.role === "assistant" && message.model && (
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
{message.model}
</span>
)}
</div>
{/* Reasoning content (chain-of-thought) */}
{message.reasoningContent && <ReasoningContent reasoningContent={message.reasoningContent} />}
{/* MCP events at the start of the last assistant message */}
{message.role === "assistant" &&
isLastMessage &&
mcpEvents.length > 0 &&
(endpointType === EndpointType.RESPONSES || endpointType === EndpointType.CHAT) && (
<div className="mb-3">
<MCPEventsDisplay events={mcpEvents} />
</div>
)}
{/* Search results */}
{message.role === "assistant" && message.searchResults && (
<SearchResultsDisplay searchResults={message.searchResults} />
)}
{/* Code Interpreter output for the last assistant message */}
{message.role === "assistant" &&
isLastMessage &&
codeInterpreterResult &&
endpointType === EndpointType.RESPONSES && (
<CodeInterpreterOutput
code={codeInterpreterResult.code}
containerId={codeInterpreterResult.containerId}
annotations={codeInterpreterResult.annotations}
accessToken={accessToken}
/>
)}
{/* Message body */}
<div
className="whitespace-pre-wrap break-words max-w-full message-content"
style={{
wordWrap: "break-word",
overflowWrap: "break-word",
wordBreak: "break-word",
hyphens: "auto",
}}
>
{message.isImage ? (
<img
src={typeof message.content === "string" ? message.content : ""}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
) : message.isAudio ? (
<AudioRenderer message={message} />
) : (
<>
{/* Attached image for user messages based on endpoint */}
{endpointType === EndpointType.RESPONSES && <ResponsesImageRenderer message={message} />}
{endpointType === EndpointType.CHAT && <ChatImageRenderer message={message} />}
<ReactMarkdown
components={{
code({
node,
inline,
className,
children,
...props
}: React.ComponentPropsWithoutRef<"code"> & {
inline?: boolean;
node?: unknown;
}) {
const match = /language-(\w+)/.exec(className || "");
return !inline && match ? (
<SyntaxHighlighter
style={coy as any}
language={match[1]}
PreTag="div"
className="rounded-md my-2"
wrapLines={true}
wrapLongLines={true}
{...props}
>
{String(children).replace(/\n$/, "")}
</SyntaxHighlighter>
) : (
<code
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
style={{ wordBreak: "break-word" }}
{...props}
>
{children}
</code>
);
},
pre: ({ node, ...props }) => (
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
),
}}
>
{typeof message.content === "string" ? message.content : ""}
</ReactMarkdown>
{/* Generated image from chat completions */}
{message.image && (
<div className="mt-3">
<img
src={message.image.url}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
</div>
)}
</>
)}
{/* Response metrics */}
{message.role === "assistant" &&
(message.timeToFirstToken || message.totalLatency || message.usage) &&
!message.a2aMetadata && (
<ResponseMetrics
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
usage={message.usage}
toolName={message.toolName}
/>
)}
{/* A2A Metrics */}
{message.role === "assistant" && message.a2aMetadata && (
<A2AMetrics
a2aMetadata={message.a2aMetadata}
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
/>
)}
</div>
</div>
</div>
);
}
export default ChatMessageBubble;

View file

@ -63,6 +63,7 @@ import EndpointSelector from "./EndpointSelector";
import FilePreviewCard from "./FilePreviewCard";
import MCPEventsDisplay from "./MCPEventsDisplay";
import type { MCPEvent } from "../../mcp_tools/types";
import ChatMessageBubble from "./ChatMessageBubble";
import { EndpointType, getEndpointType } from "./mode_endpoint_mapping";
import ReasoningContent from "./ReasoningContent";
import ResponseMetrics, { TokenUsage } from "./ResponseMetrics";
@ -1932,168 +1933,14 @@ const ChatUI: React.FC<ChatUIProps> = ({
{chatHistory.map((message, index) => (
<div key={index}>
<div className={`mb-4 ${message.role === "user" ? "text-right" : "text-left"}`}>
<div
className="inline-block max-w-[80%] rounded-lg shadow-sm p-3.5 px-4"
style={{
backgroundColor: message.role === "user" ? "#f0f8ff" : "#ffffff",
border: message.role === "user" ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
textAlign: "left",
}}
>
<div className="flex items-center gap-2 mb-1.5">
<div
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
style={{
backgroundColor: message.role === "user" ? "#e6f0fa" : "#f5f5f5",
}}
>
{message.role === "user" ? (
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
) : (
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
)}
</div>
<strong className="text-sm capitalize">{message.role}</strong>
{message.role === "assistant" && message.model && (
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
{message.model}
</span>
)}
</div>
{message.reasoningContent && <ReasoningContent reasoningContent={message.reasoningContent} />}
{/* Show MCP events at the start of assistant messages */}
{message.role === "assistant" &&
index === chatHistory.length - 1 &&
mcpEvents.length > 0 &&
(endpointType === EndpointType.RESPONSES || endpointType === EndpointType.CHAT) && (
<div className="mb-3">
<MCPEventsDisplay events={mcpEvents} />
</div>
)}
{/* Show search results at the start of assistant messages */}
{message.role === "assistant" && message.searchResults && (
<SearchResultsDisplay searchResults={message.searchResults} />
)}
{/* Show Code Interpreter output for the last assistant message */}
{message.role === "assistant" &&
index === chatHistory.length - 1 &&
codeInterpreter.result &&
endpointType === EndpointType.RESPONSES && (
<CodeInterpreterOutput
code={codeInterpreter.result.code}
containerId={codeInterpreter.result.containerId}
annotations={codeInterpreter.result.annotations}
accessToken={apiKeySource === "session" ? accessToken || "" : apiKey}
/>
)}
<div
className="whitespace-pre-wrap break-words max-w-full message-content"
style={{
wordWrap: "break-word",
overflowWrap: "break-word",
wordBreak: "break-word",
hyphens: "auto",
}}
>
{message.isImage ? (
<img
src={typeof message.content === "string" ? message.content : ""}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
) : message.isAudio ? (
<AudioRenderer message={message} />
) : (
<>
{/* Show attached image for user messages based on current endpoint */}
{endpointType === EndpointType.RESPONSES && <ResponsesImageRenderer message={message} />}
{endpointType === EndpointType.CHAT && <ChatImageRenderer message={message} />}
<ReactMarkdown
components={{
code({
node,
inline,
className,
children,
...props
}: React.ComponentPropsWithoutRef<"code"> & {
inline?: boolean;
node?: any;
}) {
const match = /language-(\w+)/.exec(className || "");
return !inline && match ? (
<SyntaxHighlighter
style={coy as any}
language={match[1]}
PreTag="div"
className="rounded-md my-2"
wrapLines={true}
wrapLongLines={true}
{...props}
>
{String(children).replace(/\n$/, "")}
</SyntaxHighlighter>
) : (
<code
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
style={{ wordBreak: "break-word" }}
{...props}
>
{children}
</code>
);
},
pre: ({ node, ...props }) => (
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
),
}}
>
{typeof message.content === "string" ? message.content : ""}
</ReactMarkdown>
{/* Show generated image from chat completions */}
{message.image && (
<div className="mt-3">
<img
src={message.image.url}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
</div>
)}
</>
)}
{message.role === "assistant" &&
(message.timeToFirstToken || message.totalLatency || message.usage) &&
!message.a2aMetadata && (
<ResponseMetrics
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
usage={message.usage}
toolName={message.toolName}
/>
)}
{/* A2A Metrics - show for A2A agent responses */}
{message.role === "assistant" && message.a2aMetadata && (
<A2AMetrics
a2aMetadata={message.a2aMetadata}
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
/>
)}
</div>
</div>
</div>
<ChatMessageBubble
message={message}
isLastMessage={index === chatHistory.length - 1}
endpointType={endpointType as EndpointType}
mcpEvents={mcpEvents}
codeInterpreterResult={codeInterpreter.result}
accessToken={apiKeySource === "session" ? accessToken || "" : apiKey}
/>
</div>
))}

View file

@ -151,6 +151,39 @@ describe("GuardrailViewer", () => {
expect(screen.queryByText(/Raw Bedrock Guardrail Response/)).not.toBeInTheDocument();
});
it("renders without crashing when guardrail_mode is null", () => {
const data = makeGuardrailInformation({ guardrail_mode: null });
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
// Null mode should display as dash
expect(screen.getByText("—")).toBeInTheDocument();
});
it("renders without crashing when guardrail_mode is an object", () => {
const data = makeGuardrailInformation({
guardrail_mode: { default: "pre_call", tags: {} },
});
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
expect(screen.getByText("PRE-CALL")).toBeInTheDocument();
});
it("renders without crashing when guardrail_mode is an array and shows in both timeline buckets", () => {
const data = makeGuardrailInformation({
guardrail_mode: ["pre_call", "post_call"],
});
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
// Mode badge shows first element formatted
expect(screen.getByText("PRE-CALL")).toBeInTheDocument();
// Entry should appear in both pre-call and post-call timeline sections
expect(screen.getByText(/Pre-call guardrail:/)).toBeInTheDocument();
expect(screen.getByText(/Post-call guardrail:/)).toBeInTheDocument();
});
it("integration: renders with real Bedrock details without mocks", async () => {
const user = userEvent.setup();
const data = makeGuardrailInformation({

View file

@ -40,7 +40,7 @@ interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_mode: string | string[] | Record<string, unknown> | null;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse | any;
@ -77,9 +77,50 @@ const PROVIDERS_WITH_CUSTOM_RENDERERS = new Set([
"litellm_content_filter",
]);
const formatMode = (mode: unknown): string => {
if (mode == null || mode === "") return "—";
const s = typeof mode === "string" ? mode : String(mode);
/**
* Extracts a plain string from guardrail_mode for display purposes.
* Returns the first mode when multiple are present.
*/
const resolveMode = (mode: GuardrailInformation["guardrail_mode"]): string | null => {
if (mode == null) return null;
if (typeof mode === "string") return mode;
if (Array.isArray(mode)) {
const first = mode[0];
return typeof first === "string" ? first : null;
}
if (typeof mode === "object" && "default" in mode) {
const def = mode.default;
if (typeof def === "string") return def;
if (Array.isArray(def)) {
const first = def[0];
return typeof first === "string" ? first : null;
}
}
return null;
};
/**
* Checks whether guardrail_mode includes the given target stage.
* Handles arrays (multi-stage guardrails) by checking all elements.
*/
const modeMatches = (
mode: GuardrailInformation["guardrail_mode"],
target: string,
): boolean => {
if (mode == null) return false;
if (typeof mode === "string") return mode === target;
if (Array.isArray(mode)) return mode.includes(target);
if (typeof mode === "object" && "default" in mode) {
const def = mode.default;
if (typeof def === "string") return def === target;
if (Array.isArray(def)) return def.some((x) => typeof x === "string" && x === target);
}
return false;
};
const formatMode = (mode: GuardrailInformation["guardrail_mode"]): string => {
const s = resolveMode(mode);
if (s == null || s === "") return "—";
return s.replace(/_/g, "-").toUpperCase();
};
@ -301,10 +342,13 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
// Request received
items.push({ type: "request", label: "Request received", offsetMs: 0 });
// Pre-call guardrails
const preCalls = sorted.filter((e) => e.guardrail_mode === "pre_call");
const postCalls = sorted.filter((e) => e.guardrail_mode === "post_call" || e.guardrail_mode === "logging_only");
const duringCalls = sorted.filter((e) => e.guardrail_mode === "during_call");
// Pre-call guardrails — use modeMatches so array modes (e.g. ["pre_call", "post_call"])
// place the entry in every matching bucket.
const preCalls = sorted.filter((e) => modeMatches(e.guardrail_mode, "pre_call"));
const postCalls = sorted.filter(
(e) => modeMatches(e.guardrail_mode, "post_call") || modeMatches(e.guardrail_mode, "logging_only"),
);
const duringCalls = sorted.filter((e) => modeMatches(e.guardrail_mode, "during_call"));
for (const e of preCalls) {
const offsetMs = Math.round((e.end_time - baseTime) * 1000);

View file

@ -23,7 +23,7 @@ export interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_mode: string | string[] | Record<string, unknown> | null;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse;