mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
commit
784f9431ad
62 changed files with 4844 additions and 1207 deletions
5
.github/workflows/test-linting.yml
vendored
5
.github/workflows/test-linting.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
---
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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="",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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 ? (
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
@ -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>
|
||||
))}
|
||||
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue