diff --git a/docs/my-website/docs/pass_through/anthropic_completion.md b/docs/my-website/docs/pass_through/anthropic_completion.md
index e644b7d348f..e0c7c7c5496 100644
--- a/docs/my-website/docs/pass_through/anthropic_completion.md
+++ b/docs/my-website/docs/pass_through/anthropic_completion.md
@@ -1,7 +1,7 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
-# Anthropic SDK
+# Anthropic Passthrough
Pass-through endpoints for Anthropic - call provider-specific endpoint, in native format (no translation).
diff --git a/docs/my-website/docs/providers/milvus_vector_stores.md b/docs/my-website/docs/providers/milvus_vector_stores.md
index c1a3042051b..84f16fbc74a 100644
--- a/docs/my-website/docs/providers/milvus_vector_stores.md
+++ b/docs/my-website/docs/providers/milvus_vector_stores.md
@@ -178,7 +178,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/vector_stores/my-collection-name/search' \
| Guardrails | ❌ Not Yet Supported | Guardrails are not currently supported for vector stores |
| Cost Tracking | ✅ Supported | Cost is $0 for Milvus searches |
| Unified API | ✅ Supported | Call via OpenAI compatible `/v1/vector_stores/search` endpoint |
-| Passthrough | ❌ Not yet supported | |
+| Passthrough | ✅ Supported | Use native Milvus API format |
## Response Format
@@ -208,6 +208,313 @@ The response follows the standard LiteLLM vector store format:
}
```
+## Passthrough API (Native Milvus Format)
+
+Use this to allow developers to **create** and **search** vector stores using the native Milvus API format, without giving them the Milvus credentials.
+
+This is for the proxy only.
+
+### Admin Flow
+
+#### 1. Add the vector store to LiteLLM
+
+```yaml
+model_list:
+ - model_name: embedding-model
+ litellm_params:
+ model: azure/text-embedding-3-large
+ api_base: https://your-endpoint.cognitiveservices.azure.com/
+ api_key: os.environ/AZURE_API_KEY
+ api_version: "2025-09-01"
+
+vector_store_registry:
+ - vector_store_name: "milvus-store"
+ litellm_params:
+ vector_store_id: "can-be-anything" # vector store id can be anything for the purpose of passthrough api
+ custom_llm_provider: "milvus"
+ api_key: os.environ/MILVUS_API_KEY
+ api_base: https://your-milvus-instance.milvus.io
+
+general_settings:
+ database_url: "postgresql://user:password@host:port/database"
+ master_key: "sk-1234"
+```
+
+Add your vector store credentials to LiteLLM.
+
+#### 2. Start the proxy
+
+```bash
+litellm --config /path/to/config.yaml
+
+# RUNNING on http://0.0.0.0:4000
+```
+
+#### 3. Create a virtual index
+
+```bash
+curl -L -X POST 'http://0.0.0.0:4000/v1/indexes' \
+-H 'Content-Type: application/json' \
+-H 'Authorization: Bearer sk-1234' \
+-d '{
+ "index_name": "dall-e-6",
+ "litellm_params": {
+ "vector_store_index": "real-collection-name",
+ "vector_store_name": "milvus-store"
+ }
+}'
+```
+
+This is a virtual index, which the developer can use to create and search vector stores.
+
+#### 4. Create a key with the vector store permissions
+
+```bash
+curl -L -X POST 'http://0.0.0.0:4000/key/generate' \
+-H 'Content-Type: application/json' \
+-H 'Authorization: Bearer sk-1234' \
+-d '{
+ "allowed_vector_store_indexes": [{"index_name": "dall-e-6", "index_permissions": ["write", "read"]}],
+ "models": ["embedding-model"]
+}'
+```
+
+Give the key access to the virtual index and the embedding model.
+
+**Expected response**
+
+```json
+{
+ "key": "sk-my-virtual-key"
+}
+```
+
+### Developer Flow
+
+#### 1. Create a collection with schema
+
+Note: Use the `/milvus` endpoint for the passthrough api that uses the `milvus` provider in your config.
+
+```python
+from milvus_rest_client import MilvusRESTClient, DataType
+import random
+import time
+
+# Configuration
+uri = "http://0.0.0.0:4000/milvus" # IMPORTANT: Use the '/milvus' endpoint for passthrough
+token = "sk-my-virtual-key"
+collection_name = "dall-e-6" # Virtual index name
+
+# Initialize client
+milvus_client = MilvusRESTClient(uri=uri, token=token)
+print(f"Connected to DB: {uri} successfully")
+
+# Check if the collection exists and drop if it does
+check_collection = milvus_client.has_collection(collection_name)
+if check_collection:
+ milvus_client.drop_collection(collection_name)
+ print(f"Dropped the existing collection {collection_name} successfully")
+
+# Define schema
+dim = 64 # Vector dimension
+
+print("Start to create the collection schema")
+schema = milvus_client.create_schema()
+schema.add_field(
+ "book_id", DataType.INT64, is_primary=True, description="customized primary id"
+)
+schema.add_field("word_count", DataType.INT64, description="word count")
+schema.add_field(
+ "book_intro", DataType.FLOAT_VECTOR, dim=dim, description="book introduction"
+)
+
+# Prepare index parameters
+print("Start to prepare index parameters with default AUTOINDEX")
+index_params = milvus_client.prepare_index_params()
+index_params.add_index("book_intro", metric_type="L2")
+
+# Create collection
+print(f"Start to create example collection: {collection_name}")
+milvus_client.create_collection(
+ collection_name, schema=schema, index_params=index_params
+)
+collection_property = milvus_client.describe_collection(collection_name)
+print("Collection details: %s" % collection_property)
+```
+
+#### 2. Insert data into the collection
+
+```python
+# Insert data with customized ids
+nb = 1000
+insert_rounds = 2
+start = 0 # first primary key id
+total_rt = 0 # total response time for insert
+
+print(
+ f"Start to insert {nb*insert_rounds} entities into example collection: {collection_name}"
+)
+for i in range(insert_rounds):
+ vector = [random.random() for _ in range(dim)]
+ rows = [
+ {"book_id": i, "word_count": random.randint(1, 100), "book_intro": vector}
+ for i in range(start, start + nb)
+ ]
+ t0 = time.time()
+ milvus_client.insert(collection_name, rows)
+ ins_rt = time.time() - t0
+ start += nb
+ total_rt += ins_rt
+print(f"Insert completed in {round(total_rt, 4)} seconds")
+
+# Flush the collection
+print("Start to flush")
+start_flush = time.time()
+milvus_client.flush(collection_name)
+end_flush = time.time()
+print(f"Flush completed in {round(end_flush - start_flush, 4)} seconds")
+```
+
+#### 3. Search the collection
+
+```python
+# Search configuration
+nq = 3 # Number of query vectors
+search_params = {"metric_type": "L2", "params": {"level": 2}}
+limit = 2 # Number of results to return
+
+# Perform searches
+for i in range(5):
+ search_vectors = [[random.random() for _ in range(dim)] for _ in range(nq)]
+ t0 = time.time()
+ results = milvus_client.search(
+ collection_name,
+ data=search_vectors,
+ limit=limit,
+ search_params=search_params,
+ anns_field="book_intro",
+ )
+ t1 = time.time()
+ print(f"Search {i} results: {results}")
+ print(f"Search {i} latency: {round(t1-t0, 4)} seconds")
+```
+
+#### Complete Example
+
+Here's a full working example:
+
+```python
+from milvus_rest_client import MilvusRESTClient, DataType
+import random
+import time
+
+# ----------------------------
+# 🔐 CONFIGURATION
+# ----------------------------
+uri = "http://0.0.0.0:4000/milvus" # IMPORTANT: Use the '/milvus' endpoint
+token = "sk-my-virtual-key"
+collection_name = "dall-e-6" # Your virtual index name
+
+# ----------------------------
+# 📋 STEP 1 — Initialize Client
+# ----------------------------
+milvus_client = MilvusRESTClient(uri=uri, token=token)
+print(f"✅ Connected to DB: {uri} successfully")
+
+# ----------------------------
+# 🗑️ STEP 2 — Drop Existing Collection (if needed)
+# ----------------------------
+check_collection = milvus_client.has_collection(collection_name)
+if check_collection:
+ milvus_client.drop_collection(collection_name)
+ print(f"🗑️ Dropped the existing collection {collection_name} successfully")
+
+# ----------------------------
+# 📐 STEP 3 — Create Collection Schema
+# ----------------------------
+dim = 64 # Vector dimension
+
+print("📐 Creating the collection schema")
+schema = milvus_client.create_schema()
+schema.add_field(
+ "book_id", DataType.INT64, is_primary=True, description="customized primary id"
+)
+schema.add_field("word_count", DataType.INT64, description="word count")
+schema.add_field(
+ "book_intro", DataType.FLOAT_VECTOR, dim=dim, description="book introduction"
+)
+
+# ----------------------------
+# 🔍 STEP 4 — Create Index
+# ----------------------------
+print("🔍 Preparing index parameters with default AUTOINDEX")
+index_params = milvus_client.prepare_index_params()
+index_params.add_index("book_intro", metric_type="L2")
+
+# ----------------------------
+# 🏗️ STEP 5 — Create Collection
+# ----------------------------
+print(f"🏗️ Creating collection: {collection_name}")
+milvus_client.create_collection(
+ collection_name, schema=schema, index_params=index_params
+)
+collection_property = milvus_client.describe_collection(collection_name)
+print(f"✅ Collection created: {collection_property}")
+
+# ----------------------------
+# 📤 STEP 6 — Insert Data
+# ----------------------------
+nb = 1000
+insert_rounds = 2
+start = 0
+total_rt = 0
+
+print(f"📤 Inserting {nb*insert_rounds} entities into collection")
+for i in range(insert_rounds):
+ vector = [random.random() for _ in range(dim)]
+ rows = [
+ {"book_id": i, "word_count": random.randint(1, 100), "book_intro": vector}
+ for i in range(start, start + nb)
+ ]
+ t0 = time.time()
+ milvus_client.insert(collection_name, rows)
+ ins_rt = time.time() - t0
+ start += nb
+ total_rt += ins_rt
+print(f"✅ Insert completed in {round(total_rt, 4)} seconds")
+
+# ----------------------------
+# 💾 STEP 7 — Flush Collection
+# ----------------------------
+print("💾 Flushing collection")
+start_flush = time.time()
+milvus_client.flush(collection_name)
+end_flush = time.time()
+print(f"✅ Flush completed in {round(end_flush - start_flush, 4)} seconds")
+
+# ----------------------------
+# 🔍 STEP 8 — Search
+# ----------------------------
+nq = 3
+search_params = {"metric_type": "L2", "params": {"level": 2}}
+limit = 2
+
+print(f"🔍 Performing {5} search operations")
+for i in range(5):
+ search_vectors = [[random.random() for _ in range(dim)] for _ in range(nq)]
+ t0 = time.time()
+ results = milvus_client.search(
+ collection_name,
+ data=search_vectors,
+ limit=limit,
+ search_params=search_params,
+ anns_field="book_intro",
+ )
+ t1 = time.time()
+ print(f"✅ Search {i} results: {results}")
+ print(f" Search {i} latency: {round(t1-t0, 4)} seconds")
+```
+
## How It Works
When you search:
diff --git a/docs/my-website/docs/proxy/logging.md b/docs/my-website/docs/proxy/logging.md
index baf3d42234d..20fc6103883 100644
--- a/docs/my-website/docs/proxy/logging.md
+++ b/docs/my-website/docs/proxy/logging.md
@@ -1335,6 +1335,7 @@ litellm_settings:
s3_aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY # AWS Secret Access Key for S3
s3_path: my-test-path # [OPTIONAL] set path in bucket you want to write logs to
s3_endpoint_url: https://s3.amazonaws.com # [OPTIONAL] S3 endpoint URL, if you want to use Backblaze/cloudflare s3 buckets
+ s3_strip_base64_files: false # [OPTIONAL] remove base64 files before storing in s3
```
**Step 3**: Start the proxy, make a test request
diff --git a/docs/my-website/docs/proxy/logging_spec.md b/docs/my-website/docs/proxy/logging_spec.md
index 6364b8c4444..e0281c9c9bf 100644
--- a/docs/my-website/docs/proxy/logging_spec.md
+++ b/docs/my-website/docs/proxy/logging_spec.md
@@ -91,7 +91,7 @@ Inherits from `StandardLoggingUserAPIKeyMetadata` and adds:
| `applied_guardrails` | `Optional[List[str]]` | List of applied guardrail names |
| `usage_object` | `Optional[dict]` | Raw usage object from the LLM provider |
| `cold_storage_object_key` | `Optional[str]` | S3/GCS object key for cold storage retrieval |
-| `guardrail_information` | `Optional[StandardLoggingGuardrailInformation]` | Guardrail information |
+| `guardrail_information` | `Optional[list[StandardLoggingGuardrailInformation]]` | Guardrail information |
## StandardLoggingVectorStoreRequest
@@ -170,7 +170,7 @@ A literal type with two possible values:
| `guardrail_mode` | `Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]]` | Guardrail mode |
| `guardrail_request` | `Optional[dict]` | Guardrail request |
| `guardrail_response` | `Optional[Union[dict, str, List[dict]]]` | Guardrail response |
-| `guardrail_status` | `Literal["success", "failure", "blocked"]` | Guardrail execution status: `success` = no violations detected, `blocked` = content blocked/modified due to policy violations, `failure` = technical error or API failure |
+| `guardrail_status` | `Literal["success", "guardrail_intervened", "guardrail_failed_to_respond"]` | Guardrail execution status: `success` = no violations detected, `blocked` = content blocked/modified due to policy violations, `failure` = technical error or API failure |
| `start_time` | `Optional[float]` | Start time of the guardrail |
| `end_time` | `Optional[float]` | End time of the guardrail |
| `duration` | `Optional[float]` | Duration of the guardrail in seconds |
diff --git a/litellm/integrations/_types/open_inference.py b/litellm/integrations/_types/open_inference.py
index af2ff2347c8..0fde1ff7525 100644
--- a/litellm/integrations/_types/open_inference.py
+++ b/litellm/integrations/_types/open_inference.py
@@ -201,6 +201,10 @@ class MessageAttributes:
"""
The id of the tool call.
"""
+ MESSAGE_REASONING_SUMMARY = "message.reasoning_summary"
+ """
+ The reasoning summary from the model's chain-of-thought process.
+ """
class MessageContentAttributes:
diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py
index e93ef128b4a..1359f279ee9 100644
--- a/litellm/integrations/arize/_utils.py
+++ b/litellm/integrations/arize/_utils.py
@@ -1,41 +1,97 @@
import json
-from typing import TYPE_CHECKING, Any, Optional, Union
+from abc import ABC
+from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, Union
+
+from typing_extensions import override
from litellm._logging import verbose_logger
+from litellm.integrations.opentelemetry_utils.base_otel_llm_obs_attributes import (
+ BaseLLMObsOTELAttributes,
+ safe_set_attribute,
+)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.types.utils import StandardLoggingPayload
if TYPE_CHECKING:
- from opentelemetry.trace import Span as _Span
-
- Span = Union[_Span, Any]
-else:
- Span = Any
+ from opentelemetry.trace import Span
-def cast_as_primitive_value_type(value) -> Union[str, bool, int, float]:
- """
- Converts a value to an OTEL-supported primitive for Arize/Phoenix observability.
- """
- if value is None:
- return ""
- if isinstance(value, (str, bool, int, float)):
- return value
- try:
- return str(value)
- except Exception:
- return ""
+class ArizeOTELAttributes(BaseLLMObsOTELAttributes):
+
+ @staticmethod
+ @override
+ def set_messages(span: "Span", kwargs: Dict[str, Any]):
+ from litellm.integrations._types.open_inference import (
+ MessageAttributes,
+ SpanAttributes,
+ )
+
+ messages = kwargs.get("messages")
+
+ # for /chat/completions
+ # https://docs.arize.com/arize/large-language-models/tracing/semantic-conventions
+ if messages:
+ last_message = messages[-1]
+ safe_set_attribute(
+ span,
+ SpanAttributes.INPUT_VALUE,
+ last_message.get("content", ""),
+ )
+
+ # LLM_INPUT_MESSAGES shows up under `input_messages` tab on the span page.
+ for idx, msg in enumerate(messages):
+ prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.{idx}"
+ # Set the role per message.
+ safe_set_attribute(
+ span, f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", msg.get("role")
+ )
+ # Set the content per message.
+ safe_set_attribute(
+ span,
+ f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}",
+ msg.get("content", ""),
+ )
+
+ @staticmethod
+ @override
+ def set_response_output_messages(span: "Span", response_obj):
+ """
+ Sets output message attributes on the span from the LLM response.
+
+ Args:
+ span: The OpenTelemetry span to set attributes on
+ response_obj: The response object containing choices with messages
+ """
+ from litellm.integrations._types.open_inference import (
+ MessageAttributes,
+ SpanAttributes,
+ )
+
+ for idx, choice in enumerate(response_obj.get("choices", [])):
+ response_message = choice.get("message", {})
+ safe_set_attribute(
+ span,
+ SpanAttributes.OUTPUT_VALUE,
+ response_message.get("content", ""),
+ )
+
+ # This shows up under `output_messages` tab on the span page.
+ prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{idx}"
+ safe_set_attribute(
+ span,
+ f"{prefix}.{MessageAttributes.MESSAGE_ROLE}",
+ response_message.get("role"),
+ )
+ safe_set_attribute(
+ span,
+ f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}",
+ response_message.get("content", ""),
+ )
-def safe_set_attribute(span: Span, key: str, value: Any):
- """
- Sets a span attribute safely with OTEL-compliant primitive typing for Arize/Phoenix.
- """
- primitive_value = cast_as_primitive_value_type(value)
- span.set_attribute(key, primitive_value)
-
-
-def set_attributes(span: Span, kwargs, response_obj): # noqa: PLR0915
+def set_attributes(
+ span: "Span", kwargs, response_obj, attributes: Type[BaseLLMObsOTELAttributes]
+): # noqa: PLR0915
"""
Populates span with OpenInference-compliant LLM attributes for Arize and Phoenix tracing.
"""
@@ -45,6 +101,7 @@ def set_attributes(span: Span, kwargs, response_obj): # noqa: PLR0915
SpanAttributes,
ToolCallAttributes,
)
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
try:
optional_params = kwargs.get("optional_params", {})
@@ -151,31 +208,7 @@ def set_attributes(span: Span, kwargs, response_obj): # noqa: PLR0915
SpanAttributes.OPENINFERENCE_SPAN_KIND,
OpenInferenceSpanKindValues.LLM.value,
)
- messages = kwargs.get("messages")
-
- # for /chat/completions
- # https://docs.arize.com/arize/large-language-models/tracing/semantic-conventions
- if messages:
- last_message = messages[-1]
- safe_set_attribute(
- span,
- SpanAttributes.INPUT_VALUE,
- last_message.get("content", ""),
- )
-
- # LLM_INPUT_MESSAGES shows up under `input_messages` tab on the span page.
- for idx, msg in enumerate(messages):
- prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.{idx}"
- # Set the role per message.
- safe_set_attribute(
- span, f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", msg.get("role")
- )
- # Set the content per message.
- safe_set_attribute(
- span,
- f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}",
- msg.get("content", ""),
- )
+ attributes.set_messages(span, kwargs)
# Capture tools (function definitions) used in the LLM call.
tools = optional_params.get("tools")
@@ -235,6 +268,7 @@ def set_attributes(span: Span, kwargs, response_obj): # noqa: PLR0915
# Captures response tokens, message, and content.
if hasattr(response_obj, "get"):
+ # Handle chat completions API (choices field)
for idx, choice in enumerate(response_obj.get("choices", [])):
response_message = choice.get("message", {})
safe_set_attribute(
@@ -256,6 +290,51 @@ def set_attributes(span: Span, kwargs, response_obj): # noqa: PLR0915
response_message.get("content", ""),
)
+ # Handle responses API (output field)
+ output_items = response_obj.get("output", [])
+ if output_items:
+ for i, item in enumerate(output_items):
+ prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{i}"
+
+ if hasattr(item, "type"):
+ item_type = item.type
+
+ # Extract reasoning summary
+ if item_type == "reasoning" and hasattr(item, "summary"):
+ for summary in item.summary:
+ if hasattr(summary, "text"):
+ safe_set_attribute(
+ span,
+ f"{prefix}.{MessageAttributes.MESSAGE_REASONING_SUMMARY}",
+ summary.text,
+ )
+
+ # Extract message content
+ elif item_type == "message" and hasattr(item, "content"):
+ message_content = ""
+
+ content_list = item.content
+ if content_list and len(content_list) > 0:
+ first_content = content_list[0]
+ message_content = getattr(first_content, "text", "")
+ message_role = getattr(item, "role", "assistant")
+
+ safe_set_attribute(
+ span,
+ SpanAttributes.OUTPUT_VALUE,
+ message_content,
+ )
+ safe_set_attribute(
+ span,
+ f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}",
+ message_content,
+ )
+ safe_set_attribute(
+ span,
+ f"{prefix}.{MessageAttributes.MESSAGE_ROLE}",
+ message_role,
+ )
+
# Token usage info.
usage = response_obj and response_obj.get("usage")
if usage:
@@ -266,18 +345,33 @@ def set_attributes(span: Span, kwargs, response_obj): # noqa: PLR0915
)
# The number of tokens used in the LLM response (completion).
- safe_set_attribute(
- span,
- SpanAttributes.LLM_TOKEN_COUNT_COMPLETION,
- usage.get("completion_tokens"),
- )
+ # Responses API uses "output_tokens", chat completions uses "completion_tokens"
+ completion_tokens = usage.get("completion_tokens") or usage.get("output_tokens")
+ if completion_tokens:
+ safe_set_attribute(
+ span,
+ SpanAttributes.LLM_TOKEN_COUNT_COMPLETION,
+ completion_tokens,
+ )
# The number of tokens used in the LLM prompt.
- safe_set_attribute(
- span,
- SpanAttributes.LLM_TOKEN_COUNT_PROMPT,
- usage.get("prompt_tokens"),
- )
+ # Responses API uses "input_tokens", chat completions uses "prompt_tokens"
+ prompt_tokens = usage.get("prompt_tokens") or usage.get("input_tokens")
+ if prompt_tokens:
+ safe_set_attribute(
+ span,
+ SpanAttributes.LLM_TOKEN_COUNT_PROMPT,
+ prompt_tokens,
+ )
+
+ # The number of reasoning tokens in the output, if available.
+ reasoning_tokens = usage.get("output_tokens_details", {}).get("reasoning_tokens")
+ if reasoning_tokens:
+ safe_set_attribute(
+ span,
+ SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING,
+ reasoning_tokens,
+ )
except Exception as e:
verbose_logger.error(
diff --git a/litellm/integrations/arize/arize.py b/litellm/integrations/arize/arize.py
index 06e05f1271d..9d587dcfa0e 100644
--- a/litellm/integrations/arize/arize.py
+++ b/litellm/integrations/arize/arize.py
@@ -9,6 +9,7 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any, Optional, Union
from litellm.integrations.arize import _utils
+from litellm.integrations.arize._utils import ArizeOTELAttributes
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.types.integrations.arize import ArizeConfig
from litellm.types.services import ServiceLoggerPayload
@@ -33,7 +34,7 @@ class ArizeLogger(OpenTelemetry):
@staticmethod
def set_arize_attributes(span: Span, kwargs, response_obj):
- _utils.set_attributes(span, kwargs, response_obj)
+ _utils.set_attributes(span, kwargs, response_obj, ArizeOTELAttributes)
return
@staticmethod
@@ -107,39 +108,38 @@ class ArizeLogger(OpenTelemetry):
async def async_health_check(self):
"""
Performs a health check for Arize integration.
-
+
Returns:
dict: Health check result with status and message
"""
try:
config = self.get_arize_config()
-
+
if not config.space_key:
return {
"status": "unhealthy",
- "error_message": "ARIZE_SPACE_KEY environment variable not set"
+ "error_message": "ARIZE_SPACE_KEY environment variable not set",
}
-
+
if not config.api_key:
return {
- "status": "unhealthy",
- "error_message": "ARIZE_API_KEY environment variable not set"
+ "status": "unhealthy",
+ "error_message": "ARIZE_API_KEY environment variable not set",
}
-
+
return {
"status": "healthy",
- "message": "Arize credentials are configured properly"
+ "message": "Arize credentials are configured properly",
}
-
+
except Exception as e:
return {
"status": "unhealthy",
- "error_message": f"Arize health check failed: {str(e)}"
+ "error_message": f"Arize health check failed: {str(e)}",
}
def construct_dynamic_otel_headers(
- self,
- standard_callback_dynamic_params: StandardCallbackDynamicParams
+ self, standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> Optional[dict]:
"""
Construct dynamic Arize headers from standard callback dynamic params
@@ -163,7 +163,7 @@ class ArizeLogger(OpenTelemetry):
dynamic_headers["arize-space-id"] = standard_callback_dynamic_params.get(
"arize_space_key"
)
-
+
#########################################################
# `api_key` handling
#########################################################
@@ -171,5 +171,5 @@ class ArizeLogger(OpenTelemetry):
dynamic_headers["api_key"] = standard_callback_dynamic_params.get(
"arize_api_key"
)
-
+
return dynamic_headers
diff --git a/litellm/integrations/arize/arize_phoenix.py b/litellm/integrations/arize/arize_phoenix.py
index 044486fcd27..60566ee55c0 100644
--- a/litellm/integrations/arize/arize_phoenix.py
+++ b/litellm/integrations/arize/arize_phoenix.py
@@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Union
from litellm._logging import verbose_logger
from litellm.integrations.arize import _utils
+from litellm.integrations.arize._utils import ArizeOTELAttributes
from litellm.types.integrations.arize_phoenix import ArizePhoenixConfig
if TYPE_CHECKING:
@@ -28,7 +29,7 @@ ARIZE_HOSTED_PHOENIX_ENDPOINT = "https://app.phoenix.arize.com/v1/traces"
class ArizePhoenixLogger:
@staticmethod
def set_arize_phoenix_attributes(span: Span, kwargs, response_obj):
- _utils.set_attributes(span, kwargs, response_obj)
+ _utils.set_attributes(span, kwargs, response_obj, ArizeOTELAttributes)
return
@staticmethod
@@ -70,7 +71,9 @@ class ArizePhoenixLogger:
otlp_auth_headers = f"api_key={api_key}"
elif api_key is not None:
# api_key/auth is optional for self hosted phoenix
- otlp_auth_headers = f"Authorization={urllib.parse.quote(f'Bearer {api_key}')}"
+ otlp_auth_headers = (
+ f"Authorization={urllib.parse.quote(f'Bearer {api_key}')}"
+ )
return ArizePhoenixConfig(
otlp_auth_headers=otlp_auth_headers, protocol=protocol, endpoint=endpoint
diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py
index c3e1a31c3ef..b50d05ed2ec 100644
--- a/litellm/integrations/custom_guardrail.py
+++ b/litellm/integrations/custom_guardrail.py
@@ -59,7 +59,6 @@ class CustomGuardrail(CustomLogger):
self.mask_response_content: bool = mask_response_content
if supported_event_hooks:
-
## validate event_hook is in supported_event_hooks
self._validate_event_hook(event_hook, supported_event_hooks)
super().__init__(**kwargs)
@@ -80,7 +79,6 @@ class CustomGuardrail(CustomLogger):
],
supported_event_hooks: List[GuardrailEventHooks],
) -> None:
-
def _validate_event_hook_list_is_in_supported_event_hooks(
event_hook: Union[List[GuardrailEventHooks], List[str]],
supported_event_hooks: List[GuardrailEventHooks],
@@ -130,15 +128,12 @@ class CustomGuardrail(CustomLogger):
self,
requested_guardrails: Union[List[str], List[Dict[str, DynamicGuardrailParams]]],
) -> bool:
-
for _guardrail in requested_guardrails:
if isinstance(_guardrail, dict):
if self.guardrail_name in _guardrail:
-
return True
elif isinstance(_guardrail, str):
if self.guardrail_name == _guardrail:
-
return True
return False
@@ -146,7 +141,6 @@ class CustomGuardrail(CustomLogger):
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
) -> Optional[dict]:
-
from litellm.proxy._types import UserAPIKeyAuth
# should run guardrail
@@ -385,14 +379,24 @@ class CustomGuardrail(CustomLogger):
duration=duration,
masked_entity_count=masked_entity_count,
)
+
+ def _append_guardrail_info(container: dict) -> None:
+ key = "standard_logging_guardrail_information"
+ existing = container.get(key)
+ if existing is None:
+ container[key] = [slg]
+ elif isinstance(existing, list):
+ existing.append(slg)
+ else:
+ # should not happen
+ container[key] = [existing, slg]
+
if "metadata" in request_data:
if request_data["metadata"] is None:
request_data["metadata"] = {}
- request_data["metadata"]["standard_logging_guardrail_information"] = slg
+ _append_guardrail_info(request_data["metadata"])
elif "litellm_metadata" in request_data:
- request_data["litellm_metadata"][
- "standard_logging_guardrail_information"
- ] = slg
+ _append_guardrail_info(request_data["litellm_metadata"])
else:
verbose_logger.warning(
"unable to log guardrail information. No metadata found in request_data"
@@ -497,37 +501,46 @@ class CustomGuardrail(CustomLogger):
"""
for key, value in vars(litellm_params).items():
setattr(self, key, value)
-
- def get_guardrails_messages_for_call_type(self, call_type: CallTypes, data: Optional[dict] = None) -> Optional[List[AllMessageValues]]:
+
+ def get_guardrails_messages_for_call_type(
+ self, call_type: CallTypes, data: Optional[dict] = None
+ ) -> Optional[List[AllMessageValues]]:
"""
Returns the messages for the given call type and data
"""
if call_type is None or data is None:
return None
-
+
#########################################################
- # /chat/completions
- # /messages
+ # /chat/completions
+ # /messages
# Both endpoints store the messages in the "messages" key
#########################################################
- if call_type == CallTypes.completion.value or call_type == CallTypes.acompletion.value or call_type == CallTypes.anthropic_messages.value:
+ if (
+ call_type == CallTypes.completion.value
+ or call_type == CallTypes.acompletion.value
+ or call_type == CallTypes.anthropic_messages.value
+ ):
return data.get("messages")
-
+
#########################################################
- # /responses
+ # /responses
# User/System messages are stored in the "input" key, use litellm transformation to get the messages
#########################################################
- if call_type == CallTypes.responses.value or call_type == CallTypes.aresponses.value:
+ if (
+ call_type == CallTypes.responses.value
+ or call_type == CallTypes.aresponses.value
+ ):
from typing import cast
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
-
+
input_data = data.get("input")
if input_data is None:
return None
-
+
messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=input_data,
responses_api_request=data,
diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py
index fc3cf4b9ff2..b44762d0af8 100644
--- a/litellm/integrations/datadog/datadog_llm_obs.py
+++ b/litellm/integrations/datadog/datadog_llm_obs.py
@@ -498,7 +498,9 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
"guardrail_information": standard_logging_payload.get(
"guardrail_information", None
),
- "is_streamed_request": self._get_stream_value_from_payload(standard_logging_payload),
+ "is_streamed_request": self._get_stream_value_from_payload(
+ standard_logging_payload
+ ),
}
#########################################################
@@ -548,21 +550,24 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
# Guardrail overhead latency
guardrail_info: Optional[
- StandardLoggingGuardrailInformation
+ list[StandardLoggingGuardrailInformation]
] = standard_logging_payload.get("guardrail_information")
if guardrail_info is not None:
- _guardrail_duration_seconds: Optional[float] = guardrail_info.get(
- "duration"
- )
- if _guardrail_duration_seconds is not None:
+ total_duration = 0.0
+ for info in guardrail_info:
+ _guardrail_duration_seconds: Optional[float] = info.get("duration")
+ if _guardrail_duration_seconds is not None:
+ total_duration += float(_guardrail_duration_seconds)
+
+ if total_duration > 0:
# Convert from seconds to milliseconds for consistency
- latency_metrics["guardrail_overhead_time_ms"] = (
- _guardrail_duration_seconds * 1000
- )
+ latency_metrics["guardrail_overhead_time_ms"] = total_duration * 1000
return latency_metrics
- def _get_stream_value_from_payload(self, standard_logging_payload: StandardLoggingPayload) -> bool:
+ def _get_stream_value_from_payload(
+ self, standard_logging_payload: StandardLoggingPayload
+ ) -> bool:
"""
Extract the stream value from standard logging payload.
diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py
index 7f807bb8b0c..a067d285245 100644
--- a/litellm/integrations/langfuse/langfuse.py
+++ b/litellm/integrations/langfuse/langfuse.py
@@ -688,11 +688,17 @@ class LangFuseLogger:
"completion_tokens": _usage_obj.completion_tokens,
"total_cost": cost if self._supports_costs() else None,
}
- usage_details = LangfuseUsageDetails(input=_usage_obj.prompt_tokens,
- output=_usage_obj.completion_tokens,
- total=_usage_obj.total_tokens,
- cache_creation_input_tokens=_usage_obj.get('cache_creation_input_tokens', 0),
- cache_read_input_tokens=_usage_obj.get('cache_read_input_tokens', 0))
+ usage_details = LangfuseUsageDetails(
+ input=_usage_obj.prompt_tokens,
+ output=_usage_obj.completion_tokens,
+ total=_usage_obj.total_tokens,
+ cache_creation_input_tokens=_usage_obj.get(
+ "cache_creation_input_tokens", 0
+ ),
+ cache_read_input_tokens=_usage_obj.get(
+ "cache_read_input_tokens", 0
+ ),
+ )
generation_name = clean_metadata.pop("generation_name", None)
if generation_name is None:
@@ -790,7 +796,7 @@ class LangFuseLogger:
"""
Get the responses API content for Langfuse logging
"""
- if hasattr(response_obj, 'output') and response_obj.output:
+ if hasattr(response_obj, "output") and response_obj.output:
# ResponsesAPIResponse.output is a list of strings
return response_obj.output
else:
@@ -880,29 +886,44 @@ class LangFuseLogger:
guardrail_information = standard_logging_object.get(
"guardrail_information", None
)
- if guardrail_information is None:
+ if not guardrail_information:
verbose_logger.debug(
- "Not logging guardrail information as span because guardrail_information is None"
+ "Not logging guardrail information as span because guardrail_information is empty"
)
return
- span = trace.span(
- name="guardrail",
- input=guardrail_information.get("guardrail_request", None),
- output=guardrail_information.get("guardrail_response", None),
- metadata={
- "guardrail_name": guardrail_information.get("guardrail_name", None),
- "guardrail_mode": guardrail_information.get("guardrail_mode", None),
- "guardrail_masked_entity_count": guardrail_information.get(
- "masked_entity_count", None
- ),
- },
- start_time=guardrail_information.get("start_time", None), # type: ignore
- end_time=guardrail_information.get("end_time", None), # type: ignore
- )
+ if not isinstance(guardrail_information, list):
+ verbose_logger.debug(
+ "Not logging guardrail information as span because guardrail_information is not a list: %s",
+ type(guardrail_information),
+ )
+ return
- verbose_logger.debug(f"Logged guardrail information as span: {span}")
- span.end()
+ for guardrail_entry in guardrail_information:
+ if not isinstance(guardrail_entry, dict):
+ verbose_logger.debug(
+ "Skipping guardrail entry with unexpected type: %s",
+ type(guardrail_entry),
+ )
+ continue
+
+ span = trace.span(
+ name="guardrail",
+ input=guardrail_entry.get("guardrail_request", None),
+ output=guardrail_entry.get("guardrail_response", None),
+ metadata={
+ "guardrail_name": guardrail_entry.get("guardrail_name", None),
+ "guardrail_mode": guardrail_entry.get("guardrail_mode", None),
+ "guardrail_masked_entity_count": guardrail_entry.get(
+ "masked_entity_count", None
+ ),
+ },
+ start_time=guardrail_entry.get("start_time", None), # type: ignore
+ end_time=guardrail_entry.get("end_time", None), # type: ignore
+ )
+
+ verbose_logger.debug(f"Logged guardrail information as span: {span}")
+ span.end()
def _add_prompt_to_generation_params(
diff --git a/litellm/integrations/langfuse/langfuse_otel.py b/litellm/integrations/langfuse/langfuse_otel.py
index 43d16b5e4cb..fc010928101 100644
--- a/litellm/integrations/langfuse/langfuse_otel.py
+++ b/litellm/integrations/langfuse/langfuse_otel.py
@@ -5,6 +5,9 @@ from typing import TYPE_CHECKING, Any, Optional, Union
from litellm._logging import verbose_logger
from litellm.integrations.arize import _utils
+from litellm.integrations.langfuse.langfuse_otel_attributes import (
+ LangfuseLLMObsOTELAttributes,
+)
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.types.integrations.langfuse_otel import (
LangfuseOtelConfig,
@@ -33,27 +36,24 @@ LANGFUSE_CLOUD_EU_ENDPOINT = "https://cloud.langfuse.com/api/public/otel"
LANGFUSE_CLOUD_US_ENDPOINT = "https://us.cloud.langfuse.com/api/public/otel"
-
class LangfuseOtelLogger(OpenTelemetry):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
-
@staticmethod
def set_langfuse_otel_attributes(span: Span, kwargs, response_obj):
"""
Sets OpenTelemetry span attributes for Langfuse observability.
Uses the same attribute setting logic as Arize Phoenix for consistency.
"""
- _utils.set_attributes(span, kwargs, response_obj)
+
+ _utils.set_attributes(span, kwargs, response_obj, LangfuseLLMObsOTELAttributes)
#########################################################
# Set Langfuse specific attributes
#########################################################
LangfuseOtelLogger._set_langfuse_specific_attributes(
- span=span,
- kwargs=kwargs,
- response_obj=response_obj
+ span=span, kwargs=kwargs, response_obj=response_obj
)
return
@@ -158,6 +158,7 @@ class LangfuseOtelLogger(OpenTelemetry):
# Set observation output (response with tool_calls if present)
if response_obj and hasattr(response_obj, "get"):
+ # Handle chat completions API (choices field)
choices = response_obj.get("choices", [])
if choices:
# Extract the first choice's message
@@ -175,7 +176,11 @@ class LangfuseOtelLogger(OpenTelemetry):
# Parse arguments from JSON string to object
try:
- arguments_obj = json.loads(arguments_str) if isinstance(arguments_str, str) else arguments_str
+ arguments_obj = (
+ json.loads(arguments_str)
+ if isinstance(arguments_str, str)
+ else arguments_str
+ )
except json.JSONDecodeError:
arguments_obj = {}
@@ -212,6 +217,44 @@ class LangfuseOtelLogger(OpenTelemetry):
safe_dumps(output_data),
)
+ # Handle responses API (output field)
+ output = response_obj.get("output", [])
+ if output:
+ output_data = []
+ for item in output:
+ if hasattr(item, "type"):
+ item_type = item.type
+
+ if item_type == "reasoning" and hasattr(item, "summary"):
+ for summary in item.summary:
+ if hasattr(summary, "text"):
+ output_data.append({
+ "role": "reasoning_summary",
+ "content": summary.text
+ })
+ elif item_type == "message":
+ output_data.append({
+ "role": getattr(item, "role", "assistant"),
+ "content": getattr(getattr(item, "content", [{}])[0], "text", "")
+ })
+ elif item_type == "function_call":
+ arguments_str = getattr(item, "arguments", "{}")
+ arguments_obj = json.loads(arguments_str) if isinstance(arguments_str, str) else arguments_str
+ langfuse_tool_call = {
+ "id": getattr(item, "id", ""),
+ "name": getattr(item, "name", ""),
+ "call_id": getattr(item, "call_id", ""),
+ "type": "function_call",
+ "arguments": arguments_obj,
+ }
+ output_data.append(langfuse_tool_call)
+ if output_data:
+ safe_set_attribute(
+ span,
+ LangfuseSpanAttributes.OBSERVATION_OUTPUT.value,
+ safe_dumps(output_data),
+ )
+
@staticmethod
def _get_langfuse_otel_host() -> Optional[str]:
"""
@@ -262,8 +305,7 @@ class LangfuseOtelLogger(OpenTelemetry):
verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}")
auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
- public_key=public_key,
- secret_key=secret_key
+ public_key=public_key, secret_key=secret_key
)
otlp_auth_headers = f"Authorization={auth_header}"
@@ -274,7 +316,7 @@ class LangfuseOtelLogger(OpenTelemetry):
return LangfuseOtelConfig(
otlp_auth_headers=otlp_auth_headers, protocol="otlp_http"
)
-
+
@staticmethod
def _get_langfuse_authorization_header(public_key: str, secret_key: str) -> str:
"""
@@ -282,11 +324,10 @@ class LangfuseOtelLogger(OpenTelemetry):
"""
auth_string = f"{public_key}:{secret_key}"
auth_header = base64.b64encode(auth_string.encode()).decode()
- return f'Basic {auth_header}'
-
+ return f"Basic {auth_header}"
+
def construct_dynamic_otel_headers(
- self,
- standard_callback_dynamic_params: StandardCallbackDynamicParams
+ self, standard_callback_dynamic_params: StandardCallbackDynamicParams
) -> Optional[dict]:
"""
Construct dynamic Langfuse headers from standard callback dynamic params
@@ -298,13 +339,17 @@ class LangfuseOtelLogger(OpenTelemetry):
"""
dynamic_headers = {}
- dynamic_langfuse_public_key = standard_callback_dynamic_params.get("langfuse_public_key")
- dynamic_langfuse_secret_key = standard_callback_dynamic_params.get("langfuse_secret_key")
+ dynamic_langfuse_public_key = standard_callback_dynamic_params.get(
+ "langfuse_public_key"
+ )
+ dynamic_langfuse_secret_key = standard_callback_dynamic_params.get(
+ "langfuse_secret_key"
+ )
if dynamic_langfuse_public_key and dynamic_langfuse_secret_key:
auth_header = LangfuseOtelLogger._get_langfuse_authorization_header(
public_key=dynamic_langfuse_public_key,
- secret_key=dynamic_langfuse_secret_key
+ secret_key=dynamic_langfuse_secret_key,
)
dynamic_headers["Authorization"] = auth_header
-
+
return dynamic_headers
diff --git a/litellm/integrations/langfuse/langfuse_otel_attributes.py b/litellm/integrations/langfuse/langfuse_otel_attributes.py
new file mode 100644
index 00000000000..f14412ad53b
--- /dev/null
+++ b/litellm/integrations/langfuse/langfuse_otel_attributes.py
@@ -0,0 +1,110 @@
+"""
+If the LLM Obs has any specific attributes to log request or response, we can add them here.
+
+Relevant Issue: https://github.com/BerriAI/litellm/issues/13764
+"""
+
+import json
+from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
+
+from numpy import isin
+from pydantic import BaseModel
+from typing_extensions import override
+
+import litellm
+from litellm.integrations.opentelemetry_utils.base_otel_llm_obs_attributes import (
+ BaseLLMObsOTELAttributes,
+ safe_set_attribute,
+)
+from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse
+from litellm.types.utils import (
+ EmbeddingResponse,
+ ImageResponse,
+ ModelResponse,
+ RerankResponse,
+ TextCompletionResponse,
+ TranscriptionResponse,
+)
+
+if TYPE_CHECKING:
+ from opentelemetry.trace import Span
+
+
+def get_output_content_by_type(
+ response_obj: Union[
+ None,
+ dict,
+ EmbeddingResponse,
+ ModelResponse,
+ TextCompletionResponse,
+ ImageResponse,
+ TranscriptionResponse,
+ RerankResponse,
+ HttpxBinaryResponseContent,
+ ResponsesAPIResponse,
+ list,
+ ],
+ kwargs: Optional[Dict[str, Any]] = None,
+) -> str:
+ """
+ Extract output content from response objects based on their type.
+
+ This utility function handles the type-specific logic for converting
+ various response objects into appropriate output formats for Langfuse logging.
+
+ Args:
+ response_obj: The response object returned by the function
+ kwargs: Optional keyword arguments containing call_type and other metadata
+
+ Returns:
+ The formatted output content suitable for Langfuse logging, or None
+ """
+ if response_obj is None:
+ return ""
+
+ kwargs = kwargs or {}
+ call_type = kwargs.get("call_type", None)
+
+ # Embedding responses - no output content
+ if call_type == "embedding" or isinstance(response_obj, EmbeddingResponse):
+ return "embedding-output"
+
+ # Binary/Speech responses
+ if isinstance(response_obj, HttpxBinaryResponseContent):
+ return "speech-output"
+
+ if isinstance(response_obj, BaseModel):
+ return response_obj.model_dump_json()
+
+ if response_obj and (
+ isinstance(response_obj, dict) or isinstance(response_obj, list)
+ ):
+ return json.dumps(response_obj)
+ else:
+ return ""
+
+
+class LangfuseLLMObsOTELAttributes(BaseLLMObsOTELAttributes):
+ @staticmethod
+ @override
+ def set_messages(span: "Span", kwargs: Dict[str, Any]):
+ prompt = {"messages": kwargs.get("messages")}
+ optional_params = kwargs.get("optional_params", {})
+ functions = optional_params.get("functions")
+ tools = optional_params.get("tools")
+ if functions is not None:
+ prompt["functions"] = functions
+ if tools is not None:
+ prompt["tools"] = tools
+
+ input = prompt
+ safe_set_attribute(span, "langfuse.observation.input", json.dumps(input))
+
+ @staticmethod
+ @override
+ def set_response_output_messages(span: "Span", response_obj):
+ safe_set_attribute(
+ span,
+ "langfuse.observation.output",
+ get_output_content_by_type(response_obj),
+ )
diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py
index 9315384ad96..9a17244a06d 100644
--- a/litellm/integrations/opentelemetry.py
+++ b/litellm/integrations/opentelemetry.py
@@ -141,7 +141,6 @@ class OpenTelemetry(CustomLogger):
meter_provider: Optional[Any] = None,
**kwargs,
):
-
if config is None:
config = OpenTelemetryConfig.from_env()
@@ -203,13 +202,14 @@ class OpenTelemetry(CustomLogger):
# Check if a TracerProvider is already set globally (e.g., by Langfuse SDK)
try:
from opentelemetry.trace import ProxyTracerProvider
+
existing_provider = trace.get_tracer_provider()
# If an actual provider exists (not the default proxy), use it
if not isinstance(existing_provider, ProxyTracerProvider):
verbose_logger.debug(
"OpenTelemetry: Using existing TracerProvider: %s",
- type(existing_provider).__name__
+ type(existing_provider).__name__,
)
tracer_provider = existing_provider
# Don't call set_tracer_provider to preserve existing context
@@ -223,7 +223,7 @@ class OpenTelemetry(CustomLogger):
# Fallback: create a new provider if something goes wrong
verbose_logger.debug(
"OpenTelemetry: Exception checking existing provider, creating new one: %s",
- str(e)
+ str(e),
)
tracer_provider = TracerProvider(resource=_get_litellm_resource())
tracer_provider.add_span_processor(self._get_span_processor())
@@ -232,7 +232,7 @@ class OpenTelemetry(CustomLogger):
# Tracer provider explicitly provided (e.g., for testing)
verbose_logger.debug(
"OpenTelemetry: Using provided TracerProvider: %s",
- type(tracer_provider).__name__
+ type(tracer_provider).__name__,
)
trace.set_tracer_provider(tracer_provider)
@@ -514,9 +514,9 @@ class OpenTelemetry(CustomLogger):
def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> Optional[dict]:
"""Extract dynamic headers from kwargs if available."""
- standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
- kwargs.get("standard_callback_dynamic_params")
- )
+ standard_callback_dynamic_params: Optional[
+ StandardCallbackDynamicParams
+ ] = kwargs.get("standard_callback_dynamic_params")
if not standard_callback_dynamic_params:
return None
@@ -775,52 +775,63 @@ class OpenTelemetry(CustomLogger):
if standard_logging_payload is None:
return
- guardrail_information = standard_logging_payload.get("guardrail_information")
- if guardrail_information is None:
+ guardrail_information_data = standard_logging_payload.get(
+ "guardrail_information"
+ )
+ if not guardrail_information_data:
return
- start_time_float = guardrail_information.get("start_time")
- end_time_float = guardrail_information.get("end_time")
- start_time_datetime = datetime.now()
- if start_time_float is not None:
- start_time_datetime = datetime.fromtimestamp(start_time_float)
- end_time_datetime = datetime.now()
- if end_time_float is not None:
- end_time_datetime = datetime.fromtimestamp(end_time_float)
+ guardrail_information_list = [
+ information
+ for information in guardrail_information_data
+ if isinstance(information, dict)
+ ]
+
+ if not guardrail_information_list:
+ return
otel_tracer: Tracer = self.get_tracer_to_use_for_request(kwargs)
- guardrail_span = otel_tracer.start_span(
- name="guardrail",
- start_time=self._to_ns(start_time_datetime),
- context=context,
- )
+ for guardrail_information in guardrail_information_list:
+ start_time_float = guardrail_information.get("start_time")
+ end_time_float = guardrail_information.get("end_time")
+ start_time_datetime = datetime.now()
+ if start_time_float is not None:
+ start_time_datetime = datetime.fromtimestamp(start_time_float)
+ end_time_datetime = datetime.now()
+ if end_time_float is not None:
+ end_time_datetime = datetime.fromtimestamp(end_time_float)
- self.safe_set_attribute(
- span=guardrail_span,
- key="guardrail_name",
- value=guardrail_information.get("guardrail_name"),
- )
-
- self.safe_set_attribute(
- span=guardrail_span,
- key="guardrail_mode",
- value=guardrail_information.get("guardrail_mode"),
- )
-
- # Set masked_entity_count directly without conversion
- masked_entity_count = guardrail_information.get("masked_entity_count")
- if masked_entity_count is not None:
- guardrail_span.set_attribute(
- "masked_entity_count", safe_dumps(masked_entity_count)
+ guardrail_span = otel_tracer.start_span(
+ name="guardrail",
+ start_time=self._to_ns(start_time_datetime),
+ context=context,
)
- self.safe_set_attribute(
- span=guardrail_span,
- key="guardrail_response",
- value=guardrail_information.get("guardrail_response"),
- )
+ self.safe_set_attribute(
+ span=guardrail_span,
+ key="guardrail_name",
+ value=guardrail_information.get("guardrail_name"),
+ )
- guardrail_span.end(end_time=self._to_ns(end_time_datetime))
+ self.safe_set_attribute(
+ span=guardrail_span,
+ key="guardrail_mode",
+ value=guardrail_information.get("guardrail_mode"),
+ )
+
+ masked_entity_count = guardrail_information.get("masked_entity_count")
+ if masked_entity_count is not None:
+ guardrail_span.set_attribute(
+ "masked_entity_count", safe_dumps(masked_entity_count)
+ )
+
+ self.safe_set_attribute(
+ span=guardrail_span,
+ key="guardrail_response",
+ value=guardrail_information.get("guardrail_response"),
+ )
+
+ guardrail_span.end(end_time=self._to_ns(end_time_datetime))
def _handle_failure(self, kwargs, response_obj, start_time, end_time):
from opentelemetry.trace import Status, StatusCode
@@ -841,10 +852,10 @@ class OpenTelemetry(CustomLogger):
)
span.set_status(Status(StatusCode.ERROR))
self.set_attributes(span, kwargs, response_obj)
-
+
# Record exception information using OTEL standard method
self._record_exception_on_span(span=span, kwargs=kwargs)
-
+
span.end(end_time=self._to_ns(end_time))
# Create span for guardrail information
@@ -856,7 +867,7 @@ class OpenTelemetry(CustomLogger):
def _record_exception_on_span(self, span: Span, kwargs: dict):
"""
Record exception information on the span using OTEL standard methods.
-
+
This extracts error information from StandardLoggingPayload and:
1. Uses span.record_exception() for the actual exception object (OTEL standard)
2. Sets structured error attributes from StandardLoggingPayloadErrorInformation
@@ -866,22 +877,22 @@ class OpenTelemetry(CustomLogger):
# Get the exception object if available
exception = kwargs.get("exception")
-
+
# Record the exception using OTEL's standard method
if exception is not None:
span.record_exception(exception)
-
+
# Get StandardLoggingPayload for structured error information
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object"
)
-
+
if standard_logging_payload is None:
return
-
+
# Extract error_information from StandardLoggingPayload
error_information = standard_logging_payload.get("error_information")
-
+
if error_information is None:
# Fallback to error_str if error_information is not available
error_str = standard_logging_payload.get("error_str")
@@ -892,7 +903,7 @@ class OpenTelemetry(CustomLogger):
value=error_str,
)
return
-
+
# Set structured error attributes from StandardLoggingPayloadErrorInformation
if error_information.get("error_code"):
self.safe_set_attribute(
@@ -900,35 +911,35 @@ class OpenTelemetry(CustomLogger):
key=ErrorAttributes.ERROR_CODE,
value=error_information["error_code"],
)
-
+
if error_information.get("error_class"):
self.safe_set_attribute(
span=span,
key=ErrorAttributes.ERROR_TYPE,
value=error_information["error_class"],
)
-
+
if error_information.get("error_message"):
self.safe_set_attribute(
span=span,
key=ErrorAttributes.ERROR_MESSAGE,
value=error_information["error_message"],
)
-
+
if error_information.get("llm_provider"):
self.safe_set_attribute(
span=span,
key=ErrorAttributes.ERROR_LLM_PROVIDER,
value=error_information["llm_provider"],
)
-
+
if error_information.get("traceback"):
self.safe_set_attribute(
span=span,
key=ErrorAttributes.ERROR_STACK_TRACE,
value=error_information["traceback"],
)
-
+
except Exception as e:
verbose_logger.exception(
"OpenTelemetry: Error recording exception on span: %s", str(e)
@@ -1363,12 +1374,16 @@ class OpenTelemetry(CustomLogger):
# Priority 1: Explicit parent span from metadata
if parent_otel_span is not None:
- verbose_logger.debug("OpenTelemetry: Using explicit parent span from metadata")
+ verbose_logger.debug(
+ "OpenTelemetry: Using explicit parent span from metadata"
+ )
return trace.set_span_in_context(parent_otel_span), parent_otel_span
# Priority 2: HTTP traceparent header
if traceparent is not None:
- verbose_logger.debug("OpenTelemetry: Using traceparent header for context propagation")
+ verbose_logger.debug(
+ "OpenTelemetry: Using traceparent header for context propagation"
+ )
carrier = {"traceparent": traceparent}
return TraceContextTextMapPropagator().extract(carrier=carrier), None
@@ -1381,16 +1396,20 @@ class OpenTelemetry(CustomLogger):
verbose_logger.debug(
"OpenTelemetry: Using active span from global context: %s (trace_id=%s, span_id=%s, is_recording=%s)",
current_span,
- format(span_context.trace_id, '032x'),
- format(span_context.span_id, '016x'),
- current_span.is_recording()
+ format(span_context.trace_id, "032x"),
+ format(span_context.span_id, "016x"),
+ current_span.is_recording(),
)
return context.get_current(), current_span
except Exception as e:
- verbose_logger.debug("OpenTelemetry: Error getting current span: %s", str(e))
+ verbose_logger.debug(
+ "OpenTelemetry: Error getting current span: %s", str(e)
+ )
# Priority 4: No parent context
- verbose_logger.debug("OpenTelemetry: No parent context found, creating root span")
+ verbose_logger.debug(
+ "OpenTelemetry: No parent context found, creating root span"
+ )
return None, None
def _get_span_processor(self, dynamic_headers: Optional[dict] = None):
diff --git a/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py b/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py
new file mode 100644
index 00000000000..8d7b2d681ab
--- /dev/null
+++ b/litellm/integrations/opentelemetry_utils/base_otel_llm_obs_attributes.py
@@ -0,0 +1,37 @@
+from abc import ABC
+from typing import TYPE_CHECKING, Any, Dict, List, Union
+
+if TYPE_CHECKING:
+ from opentelemetry.trace import Span
+
+
+class BaseLLMObsOTELAttributes(ABC):
+ @staticmethod
+ def set_messages(span: "Span", kwargs: Dict[str, Any]):
+ pass
+
+ @staticmethod
+ def set_response_output_messages(span: "Span", response_obj):
+ pass
+
+
+def cast_as_primitive_value_type(value) -> Union[str, bool, int, float]:
+ """
+ Converts a value to an OTEL-supported primitive for Arize/Phoenix observability.
+ """
+ if value is None:
+ return ""
+ if isinstance(value, (str, bool, int, float)):
+ return value
+ try:
+ return str(value)
+ except Exception:
+ return ""
+
+
+def safe_set_attribute(span: "Span", key: str, value: Any):
+ """
+ Sets a span attribute safely with OTEL-compliant primitive typing for Arize/Phoenix.
+ """
+ primitive_value = cast_as_primitive_value_type(value)
+ span.set_attribute(key, primitive_value)
diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py
index a65500c80dc..cc44450e737 100644
--- a/litellm/integrations/s3_v2.py
+++ b/litellm/integrations/s3_v2.py
@@ -49,6 +49,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_batch_size: Optional[int] = DEFAULT_S3_BATCH_SIZE,
s3_config=None,
s3_use_team_prefix: bool = False,
+ s3_strip_base64_files: bool = False,
**kwargs,
):
try:
@@ -80,6 +81,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_config=s3_config,
s3_path=s3_path,
s3_use_team_prefix=s3_use_team_prefix,
+ s3_strip_base64_files=s3_strip_base64_files
)
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
@@ -124,6 +126,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_config=None,
s3_path: Optional[str] = None,
s3_use_team_prefix: bool = False,
+ s3_strip_base64_files: bool = False,
):
"""
Initialize the s3 params for this logging callback
@@ -194,6 +197,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
or s3_use_team_prefix
)
+ self.s3_strip_base64_files = (
+ bool(litellm.s3_callback_params.get("s3_strip_base64_files", False))
+ or s3_strip_base64_files
+ )
+
return
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
@@ -364,6 +372,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
if standard_logging_payload is None:
return None
+ if self.s3_strip_base64_files:
+ import asyncio
+ standard_logging_payload = asyncio.run(self._strip_base64_from_messages(standard_logging_payload))
+
team_alias = standard_logging_payload["metadata"].get("user_api_key_team_alias")
team_alias_prefix = ""
diff --git a/litellm/integrations/sqs.py b/litellm/integrations/sqs.py
index 545aebbec6d..b353c3670f3 100644
--- a/litellm/integrations/sqs.py
+++ b/litellm/integrations/sqs.py
@@ -256,6 +256,8 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
standard_logging_payload = kwargs.get("standard_logging_object")
if standard_logging_payload is None:
raise ValueError("standard_logging_payload is None")
+ if self.sqs_strip_base64_files:
+ standard_logging_payload = await self._strip_base64_from_messages(standard_logging_payload)
self.log_queue.append(standard_logging_payload)
verbose_logger.debug(
diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py
index afb6da0bae2..fac57e038f0 100644
--- a/litellm/litellm_core_utils/litellm_logging.py
+++ b/litellm/litellm_core_utils/litellm_logging.py
@@ -68,7 +68,9 @@ from litellm.litellm_core_utils.redact_messages import (
redact_message_input_output_from_custom_logger,
redact_message_input_output_from_logging,
)
+from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.responses.utils import ResponseAPILoggingUtils
+from litellm.types.containers.main import ContainerObject
from litellm.types.llms.openai import (
AllMessageValues,
Batch,
@@ -76,6 +78,7 @@ from litellm.types.llms.openai import (
HttpxBinaryResponseContent,
OpenAIFileObject,
OpenAIModerationResponse,
+ ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIResponse,
)
@@ -116,9 +119,7 @@ from litellm.types.utils import (
Usage,
)
from litellm.types.videos.main import VideoObject
-from litellm.types.containers.main import ContainerObject
from litellm.utils import _get_base_model_from_metadata, executor, print_verbose
-from litellm.llms.base_llm.ocr.transformation import OCRResponse
from ..integrations.argilla import ArgillaLogger
from ..integrations.arize.arize_phoenix import ArizePhoenixLogger
@@ -307,9 +308,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.litellm_trace_id: str = litellm_trace_id or str(uuid.uuid4())
self.function_id = function_id
self.streaming_chunks: List[Any] = [] # for generating complete stream response
- self.sync_streaming_chunks: List[Any] = (
- []
- ) # for generating complete stream response
+ self.sync_streaming_chunks: List[
+ Any
+ ] = [] # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
# Initialize dynamic callbacks
@@ -685,9 +686,9 @@ class Logging(LiteLLMLoggingBaseClass):
if anthropic_cache_control_logger := AnthropicCacheControlHook.get_custom_logger_for_anthropic_cache_control_hook(
non_default_params
):
- self.model_call_details["prompt_integration"] = (
- anthropic_cache_control_logger.__class__.__name__
- )
+ self.model_call_details[
+ "prompt_integration"
+ ] = anthropic_cache_control_logger.__class__.__name__
return anthropic_cache_control_logger
#########################################################
@@ -699,9 +700,9 @@ class Logging(LiteLLMLoggingBaseClass):
internal_usage_cache=None,
llm_router=None,
)
- self.model_call_details["prompt_integration"] = (
- vector_store_custom_logger.__class__.__name__
- )
+ self.model_call_details[
+ "prompt_integration"
+ ] = vector_store_custom_logger.__class__.__name__
# Add to global callbacks so post-call hooks are invoked
if (
vector_store_custom_logger
@@ -761,9 +762,9 @@ class Logging(LiteLLMLoggingBaseClass):
model
): # if model name was changes pre-call, overwrite the initial model call name with the new one
self.model_call_details["model"] = model
- self.model_call_details["litellm_params"]["api_base"] = (
- self._get_masked_api_base(additional_args.get("api_base", ""))
- )
+ self.model_call_details["litellm_params"][
+ "api_base"
+ ] = self._get_masked_api_base(additional_args.get("api_base", ""))
def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915
# Log the exact input to the LLM API
@@ -792,10 +793,10 @@ class Logging(LiteLLMLoggingBaseClass):
try:
# [Non-blocking Extra Debug Information in metadata]
if turn_off_message_logging is True:
- _metadata["raw_request"] = (
- "redacted by litellm. \
+ _metadata[
+ "raw_request"
+ ] = "redacted by litellm. \
'litellm.turn_off_message_logging=True'"
- )
else:
curl_command = self._get_request_curl_command(
api_base=additional_args.get("api_base", ""),
@@ -806,32 +807,32 @@ class Logging(LiteLLMLoggingBaseClass):
_metadata["raw_request"] = str(curl_command)
# split up, so it's easier to parse in the UI
- self.model_call_details["raw_request_typed_dict"] = (
- RawRequestTypedDict(
- raw_request_api_base=str(
- additional_args.get("api_base") or ""
- ),
- raw_request_body=self._get_raw_request_body(
- additional_args.get("complete_input_dict", {})
- ),
- raw_request_headers=self._get_masked_headers(
- additional_args.get("headers", {}) or {},
- ignore_sensitive_headers=True,
- ),
- error=None,
- )
+ self.model_call_details[
+ "raw_request_typed_dict"
+ ] = RawRequestTypedDict(
+ raw_request_api_base=str(
+ additional_args.get("api_base") or ""
+ ),
+ raw_request_body=self._get_raw_request_body(
+ additional_args.get("complete_input_dict", {})
+ ),
+ raw_request_headers=self._get_masked_headers(
+ additional_args.get("headers", {}) or {},
+ ignore_sensitive_headers=True,
+ ),
+ error=None,
)
except Exception as e:
- self.model_call_details["raw_request_typed_dict"] = (
- RawRequestTypedDict(
- error=str(e),
- )
+ self.model_call_details[
+ "raw_request_typed_dict"
+ ] = RawRequestTypedDict(
+ error=str(e),
)
- _metadata["raw_request"] = (
- "Unable to Log \
+ _metadata[
+ "raw_request"
+ ] = "Unable to Log \
raw request: {}".format(
- str(e)
- )
+ str(e)
)
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
try:
@@ -1132,13 +1133,13 @@ class Logging(LiteLLMLoggingBaseClass):
for callback in callbacks:
try:
if isinstance(callback, CustomLogger):
- response: Optional[MCPPostCallResponseObject] = (
- await callback.async_post_mcp_tool_call_hook(
- kwargs=kwargs,
- response_obj=post_mcp_tool_call_response_obj,
- start_time=start_time,
- end_time=end_time,
- )
+ response: Optional[
+ MCPPostCallResponseObject
+ ] = await callback.async_post_mcp_tool_call_hook(
+ kwargs=kwargs,
+ response_obj=post_mcp_tool_call_response_obj,
+ start_time=start_time,
+ end_time=end_time,
)
######################################################################
# if any of the callbacks modify the response, use the modified response
@@ -1301,16 +1302,16 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
- self.model_call_details["response_cost_failure_debug_information"] = (
- debug_info
- )
+ self.model_call_details[
+ "response_cost_failure_debug_information"
+ ] = debug_info
return None
try:
-
response_cost = litellm.response_cost_calculator(
**response_cost_calculator_kwargs
)
+
verbose_logger.debug(f"response_cost: {response_cost}")
return response_cost
except Exception as e: # error calculating cost
@@ -1329,9 +1330,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
- self.model_call_details["response_cost_failure_debug_information"] = (
- debug_info
- )
+ self.model_call_details[
+ "response_cost_failure_debug_information"
+ ] = debug_info
return None
@@ -1475,9 +1476,9 @@ class Logging(LiteLLMLoggingBaseClass):
end_time = datetime.datetime.now()
if self.completion_start_time is None:
self.completion_start_time = end_time
- self.model_call_details["completion_start_time"] = (
- self.completion_start_time
- )
+ self.model_call_details[
+ "completion_start_time"
+ ] = self.completion_start_time
self.model_call_details["log_event_type"] = "successful_api_call"
self.model_call_details["end_time"] = end_time
self.model_call_details["cache_hit"] = cache_hit
@@ -1530,39 +1531,39 @@ class Logging(LiteLLMLoggingBaseClass):
"response_cost"
]
else:
- self.model_call_details["response_cost"] = (
- self._response_cost_calculator(result=logging_result)
- )
+ self.model_call_details[
+ "response_cost"
+ ] = self._response_cost_calculator(result=logging_result)
## STANDARDIZED LOGGING PAYLOAD
- self.model_call_details["standard_logging_object"] = (
- get_standard_logging_object_payload(
- kwargs=self.model_call_details,
- init_response_obj=logging_result,
- start_time=start_time,
- end_time=end_time,
- logging_obj=self,
- status="success",
- standard_built_in_tools_params=self.standard_built_in_tools_params,
- )
+ self.model_call_details[
+ "standard_logging_object"
+ ] = get_standard_logging_object_payload(
+ kwargs=self.model_call_details,
+ init_response_obj=logging_result,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=self,
+ status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
elif isinstance(result, dict) or isinstance(result, list):
## STANDARDIZED LOGGING PAYLOAD
- self.model_call_details["standard_logging_object"] = (
- get_standard_logging_object_payload(
- kwargs=self.model_call_details,
- init_response_obj=result,
- start_time=start_time,
- end_time=end_time,
- logging_obj=self,
- status="success",
- standard_built_in_tools_params=self.standard_built_in_tools_params,
- )
+ self.model_call_details[
+ "standard_logging_object"
+ ] = get_standard_logging_object_payload(
+ kwargs=self.model_call_details,
+ init_response_obj=result,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=self,
+ status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
elif standard_logging_object is not None:
- self.model_call_details["standard_logging_object"] = (
- standard_logging_object
- )
+ self.model_call_details[
+ "standard_logging_object"
+ ] = standard_logging_object
else: # streaming chunks + image gen.
self.model_call_details["response_cost"] = None
@@ -1718,23 +1719,23 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
"Logging Details LiteLLM-Success Call streaming complete"
)
- self.model_call_details["complete_streaming_response"] = (
- complete_streaming_response
- )
- self.model_call_details["response_cost"] = (
- self._response_cost_calculator(result=complete_streaming_response)
- )
+ self.model_call_details[
+ "complete_streaming_response"
+ ] = complete_streaming_response
+ self.model_call_details[
+ "response_cost"
+ ] = self._response_cost_calculator(result=complete_streaming_response)
## STANDARDIZED LOGGING PAYLOAD
- self.model_call_details["standard_logging_object"] = (
- get_standard_logging_object_payload(
- kwargs=self.model_call_details,
- init_response_obj=complete_streaming_response,
- start_time=start_time,
- end_time=end_time,
- logging_obj=self,
- status="success",
- standard_built_in_tools_params=self.standard_built_in_tools_params,
- )
+ self.model_call_details[
+ "standard_logging_object"
+ ] = get_standard_logging_object_payload(
+ kwargs=self.model_call_details,
+ init_response_obj=complete_streaming_response,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=self,
+ status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=self.dynamic_success_callbacks,
@@ -2062,10 +2063,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
- self.model_call_details["complete_response"] = (
- self.model_call_details.get(
- "complete_streaming_response", {}
- )
+ self.model_call_details[
+ "complete_response"
+ ] = self.model_call_details.get(
+ "complete_streaming_response", {}
)
result = self.model_call_details["complete_response"]
openMeterLogger.log_success_event(
@@ -2104,10 +2105,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
- self.model_call_details["complete_response"] = (
- self.model_call_details.get(
- "complete_streaming_response", {}
- )
+ self.model_call_details[
+ "complete_response"
+ ] = self.model_call_details.get(
+ "complete_streaming_response", {}
)
result = self.model_call_details["complete_response"]
@@ -2245,9 +2246,9 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
print_verbose("Async success callbacks: Got a complete streaming response")
- self.model_call_details["async_complete_streaming_response"] = (
- complete_streaming_response
- )
+ self.model_call_details[
+ "async_complete_streaming_response"
+ ] = complete_streaming_response
try:
if self.model_call_details.get("cache_hit", False) is True:
@@ -2258,10 +2259,10 @@ class Logging(LiteLLMLoggingBaseClass):
model_call_details=self.model_call_details
)
# base_model defaults to None if not set on model_info
- self.model_call_details["response_cost"] = (
- self._response_cost_calculator(
- result=complete_streaming_response
- )
+ self.model_call_details[
+ "response_cost"
+ ] = self._response_cost_calculator(
+ result=complete_streaming_response
)
verbose_logger.debug(
@@ -2274,16 +2275,16 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = None
## STANDARDIZED LOGGING PAYLOAD
- self.model_call_details["standard_logging_object"] = (
- get_standard_logging_object_payload(
- kwargs=self.model_call_details,
- init_response_obj=complete_streaming_response,
- start_time=start_time,
- end_time=end_time,
- logging_obj=self,
- status="success",
- standard_built_in_tools_params=self.standard_built_in_tools_params,
- )
+ self.model_call_details[
+ "standard_logging_object"
+ ] = get_standard_logging_object_payload(
+ kwargs=self.model_call_details,
+ init_response_obj=complete_streaming_response,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=self,
+ status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=self.dynamic_async_success_callbacks,
@@ -2496,18 +2497,18 @@ class Logging(LiteLLMLoggingBaseClass):
## STANDARDIZED LOGGING PAYLOAD
- self.model_call_details["standard_logging_object"] = (
- get_standard_logging_object_payload(
- kwargs=self.model_call_details,
- init_response_obj={},
- start_time=start_time,
- end_time=end_time,
- logging_obj=self,
- status="failure",
- error_str=str(exception),
- original_exception=exception,
- standard_built_in_tools_params=self.standard_built_in_tools_params,
- )
+ self.model_call_details[
+ "standard_logging_object"
+ ] = get_standard_logging_object_payload(
+ kwargs=self.model_call_details,
+ init_response_obj={},
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=self,
+ status="failure",
+ error_str=str(exception),
+ original_exception=exception,
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
return start_time, end_time
@@ -2997,6 +2998,17 @@ class Logging(LiteLLMLoggingBaseClass):
elif isinstance(result, TextCompletionResponse):
return result
elif isinstance(result, ResponseCompletedEvent):
+ ## return unified Usage object
+ if isinstance(result.response.usage, ResponseAPIUsage):
+ setattr(
+ result.response,
+ "usage",
+ (
+ ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
+ result.response.usage
+ )
+ ),
+ )
return result.response
else:
return None
@@ -3395,9 +3407,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
endpoint=arize_config.endpoint,
)
- os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
- f"space_id={arize_config.space_key},api_key={arize_config.api_key}"
- )
+ os.environ[
+ "OTEL_EXPORTER_OTLP_TRACES_HEADERS"
+ ] = f"space_id={arize_config.space_key},api_key={arize_config.api_key}"
for callback in _in_memory_loggers:
if (
isinstance(callback, ArizeLogger)
@@ -3421,9 +3433,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
# auth can be disabled on local deployments of arize phoenix
if arize_phoenix_config.otlp_auth_headers is not None:
- os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
- arize_phoenix_config.otlp_auth_headers
- )
+ os.environ[
+ "OTEL_EXPORTER_OTLP_TRACES_HEADERS"
+ ] = arize_phoenix_config.otlp_auth_headers
for callback in _in_memory_loggers:
if (
@@ -3555,9 +3567,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
exporter="otlp_http",
endpoint="https://langtrace.ai/api/trace",
)
- os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
- f"api_key={os.getenv('LANGTRACE_API_KEY')}"
- )
+ os.environ[
+ "OTEL_EXPORTER_OTLP_TRACES_HEADERS"
+ ] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetry)
@@ -4257,10 +4269,10 @@ class StandardLoggingPayloadSetup:
for key in StandardLoggingHiddenParams.__annotations__.keys():
if key in hidden_params:
if key == "additional_headers":
- clean_hidden_params["additional_headers"] = (
- StandardLoggingPayloadSetup.get_additional_headers(
- hidden_params[key]
- )
+ clean_hidden_params[
+ "additional_headers"
+ ] = StandardLoggingPayloadSetup.get_additional_headers(
+ hidden_params[key]
)
else:
clean_hidden_params[key] = hidden_params[key] # type: ignore
@@ -4484,7 +4496,7 @@ class StandardLoggingPayloadSetup:
def _get_status_fields(
status: StandardLoggingPayloadStatus,
- guardrail_information: Optional[dict],
+ guardrail_information: Optional[list[dict]],
error_str: Optional[str],
) -> "StandardLoggingPayloadStatusFields":
"""
@@ -4515,9 +4527,13 @@ def _get_status_fields(
# Map - guardrail_information.guardrail_status to guardrail_status
#########################################################
guardrail_status: GuardrailStatus = "not_run"
- if guardrail_information and isinstance(guardrail_information, dict):
- raw_status = guardrail_information.get("guardrail_status", "not_run")
- guardrail_status = GUARDRAIL_STATUS_MAP.get(raw_status, "not_run")
+ if guardrail_information and isinstance(guardrail_information, list):
+ for information in guardrail_information:
+ if isinstance(information, dict):
+ raw_status = information.get("guardrail_status", "not_run")
+ if raw_status != "not_run":
+ guardrail_status = GUARDRAIL_STATUS_MAP.get(raw_status, "not_run")
+ break
return StandardLoggingPayloadStatusFields(
llm_api_status=llm_api_status, guardrail_status=guardrail_status
@@ -4819,9 +4835,9 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]):
):
for k, v in metadata["user_api_key_metadata"].items():
if k == "logging": # prevent logging user logging keys
- cleaned_user_api_key_metadata[k] = (
- "scrubbed_by_litellm_for_sensitive_keys"
- )
+ cleaned_user_api_key_metadata[
+ k
+ ] = "scrubbed_by_litellm_for_sensitive_keys"
else:
cleaned_user_api_key_metadata[k] = v
diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py
index 64223a9ba4e..a5c3862cf30 100644
--- a/litellm/litellm_core_utils/streaming_handler.py
+++ b/litellm/litellm_core_utils/streaming_handler.py
@@ -734,6 +734,14 @@ class CustomStreamWrapper:
"function_call" in completion_obj
and completion_obj["function_call"] is not None
)
+ or (
+ "tool_calls" in model_response.choices[0].delta
+ and model_response.choices[0].delta["tool_calls"] is not None
+ )
+ or (
+ "function_call" in model_response.choices[0].delta
+ and model_response.choices[0].delta["function_call"] is not None
+ )
or (
"reasoning_content" in model_response.choices[0].delta
and model_response.choices[0].delta.reasoning_content is not None
diff --git a/litellm/llms/deepgram/audio_transcription/transformation.py b/litellm/llms/deepgram/audio_transcription/transformation.py
index faa68e97eae..fe63ad11bc7 100644
--- a/litellm/llms/deepgram/audio_transcription/transformation.py
+++ b/litellm/llms/deepgram/audio_transcription/transformation.py
@@ -90,8 +90,21 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
first_channel = response_json["results"]["channels"][0]
first_alternative = first_channel["alternatives"][0]
- # Extract the full transcript
- text = first_alternative["transcript"]
+ # Detect if diarization is active by checking if words have 'speaker' field
+ has_diarization = False
+ if "words" in first_alternative and len(first_alternative["words"]) > 0:
+ has_diarization = "speaker" in first_alternative["words"][0]
+
+ # Extract the transcript based on diarization mode
+ if not has_diarization:
+ # No diarization: use the standard transcript
+ text = first_alternative["transcript"]
+ elif "paragraphs" in first_alternative:
+ # Diarization with paragraphs: use the pre-formatted diarized transcript
+ text = first_alternative["paragraphs"]["transcript"]
+ else:
+ # Diarization without paragraphs: reconstruct from words
+ text = self._reconstruct_diarized_transcript(first_alternative["words"])
# Create TranscriptionResponse object
response = TranscriptionResponse(text=text)
@@ -122,6 +135,46 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
f"Error transforming Deepgram response: {str(e)}\nResponse: {raw_response.text}"
)
+ def _reconstruct_diarized_transcript(self, words: list) -> str:
+ """
+ Reconstructs a diarized transcript from words with speaker information.
+
+ Args:
+ words: List of word objects with speaker, word, and optionally punctuated_word
+
+ Returns:
+ Formatted transcript with speaker labels
+ """
+ if not words:
+ return ""
+
+ segments = []
+ current_speaker = None
+ current_words = []
+
+ for word_obj in words:
+ speaker = word_obj.get("speaker")
+ # Use punctuated_word if available, otherwise fall back to word
+ word_text = word_obj.get("punctuated_word", word_obj.get("word", ""))
+
+ if speaker != current_speaker:
+ # New speaker: save previous segment and start new one
+ if current_words:
+ segments.append(
+ f"Speaker {current_speaker}: {' '.join(current_words)}"
+ )
+ current_speaker = speaker
+ current_words = [word_text]
+ else:
+ # Same speaker: add word to current segment
+ current_words.append(word_text)
+
+ # Add the last segment
+ if current_words:
+ segments.append(f"\nSpeaker {current_speaker}: {' '.join(current_words)}\n")
+
+ return "\n".join(segments)
+
def get_complete_url(
self,
api_base: Optional[str],
diff --git a/litellm/llms/perplexity/chat/__init__.py b/litellm/llms/perplexity/chat/__init__.py
new file mode 100644
index 00000000000..f4f9edf38e5
--- /dev/null
+++ b/litellm/llms/perplexity/chat/__init__.py
@@ -0,0 +1 @@
+"""Perplexity chat completion transformations."""
diff --git a/litellm/llms/perplexity/chat/transformation.py b/litellm/llms/perplexity/chat/transformation.py
index 27e6415ff8b..831b009de7a 100644
--- a/litellm/llms/perplexity/chat/transformation.py
+++ b/litellm/llms/perplexity/chat/transformation.py
@@ -1,25 +1,32 @@
-"""
-Translate from OpenAI's `/v1/chat/completions` to Perplexity's `/v1/chat/completions`
-"""
+"""Translate from OpenAI's `/v1/chat/completions` to Perplexity's `/v1/chat/completions`."""
-from typing import Any, List, Optional, Tuple
+from __future__ import annotations
+
+from typing import TYPE_CHECKING, Any, List, Optional, Tuple
-import httpx
import litellm
from litellm._logging import verbose_logger
-from litellm.secret_managers.main import get_secret_str
-from litellm.types.llms.openai import AllMessageValues
-from litellm.types.utils import Usage, PromptTokensDetailsWrapper
-from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
-from litellm.types.utils import ModelResponse
-from litellm.types.llms.openai import ChatCompletionAnnotation
-from litellm.types.llms.openai import ChatCompletionAnnotationURLCitation
+from litellm.secret_managers.main import get_secret_str
+from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage
+
+if TYPE_CHECKING:
+ import httpx
+
+ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+ from litellm.types.llms.openai import (
+ AllMessageValues,
+ ChatCompletionAnnotation,
+ ChatCompletionAnnotationURLCitation,
+ )
class PerplexityChatConfig(OpenAIGPTConfig):
+ """Configuration for Perplexity chat completions."""
+
@property
- def custom_llm_provider(self) -> Optional[str]:
+ def custom_llm_provider(self) -> str | None:
+ """Return the custom LLM provider name."""
return "perplexity"
def _get_openai_compatible_provider_info(
@@ -33,6 +40,38 @@ class PerplexityChatConfig(OpenAIGPTConfig):
)
return api_base, dynamic_api_key
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: list,
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ """Validate Perplexity environment and set headers."""
+ # Get API key from environment if not provided
+ if api_key is None:
+ _, api_key = self._get_openai_compatible_provider_info(
+ api_base=api_base, api_key=api_key
+ )
+
+ # Validate API key is present
+ if api_key is None:
+ raise ValueError(
+ "The api_key client option must be set either by passing api_key to the client or by setting the PERPLEXITY_API_KEY environment variable"
+ )
+
+ # Set authorization header
+ headers["Authorization"] = f"Bearer {api_key}"
+
+ # Ensure Content-Type is set to application/json
+ if "content-type" not in headers and "Content-Type" not in headers:
+ headers["Content-Type"] = "application/json"
+
+ return headers
+
def get_supported_openai_params(self, model: str) -> list:
"""
Perplexity supports a subset of OpenAI params
@@ -72,7 +111,8 @@ class PerplexityChatConfig(OpenAIGPTConfig):
return base_openai_params
- def transform_response(
+
+ def transform_response( # noqa: PLR0913
self,
model: str,
raw_response: httpx.Response,
@@ -82,10 +122,11 @@ class PerplexityChatConfig(OpenAIGPTConfig):
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
- encoding: Any,
+ encoding: Any,
api_key: Optional[str] = None,
- json_mode: Optional[bool] = None,
+ json_mode: Optional[bool] = None,
) -> ModelResponse:
+ """Transform Perplexity response to standard format."""
# Call the parent transform_response first to handle the standard transformation
model_response = super().transform_response(
model=model,
@@ -104,28 +145,29 @@ class PerplexityChatConfig(OpenAIGPTConfig):
# Extract and enhance usage with Perplexity-specific fields
try:
raw_response_json = raw_response.json()
+ self.add_cost_to_usage(model_response, raw_response_json)
self._enhance_usage_with_perplexity_fields(
- model_response, raw_response_json
+ model_response, raw_response_json,
)
self._add_citations_as_annotations(model_response, raw_response_json)
- except Exception as e:
+ except (ValueError, TypeError, KeyError) as e:
verbose_logger.debug(f"Error extracting Perplexity-specific usage fields: {e}")
return model_response
- def _enhance_usage_with_perplexity_fields(
- self, model_response: ModelResponse, raw_response_json: dict
+ def _enhance_usage_with_perplexity_fields(
+ self, model_response: ModelResponse, raw_response_json: dict,
) -> None:
- """
- Extract citation tokens and search queries from Perplexity API response
- and add them to the usage object using standard LiteLLM fields.
+ """Extract citation tokens and search queries from Perplexity API response.
+
+ Add them to the usage object using standard LiteLLM fields.
"""
if not hasattr(model_response, "usage") or model_response.usage is None:
# Create a usage object if it doesn't exist (when usage was None)
model_response.usage = Usage( # type: ignore[attr-defined]
prompt_tokens=0,
completion_tokens=0,
- total_tokens=0
+ total_tokens=0,
)
usage = model_response.usage # type: ignore[attr-defined]
@@ -146,7 +188,7 @@ class PerplexityChatConfig(OpenAIGPTConfig):
# Extract search queries count from usage or response metadata
# Perplexity might include this in the usage object or as separate metadata
perplexity_usage = raw_response_json.get("usage", {})
-
+
# Try to extract search queries from usage field first, then root level
num_search_queries = perplexity_usage.get("num_search_queries")
if num_search_queries is None:
@@ -155,18 +197,18 @@ class PerplexityChatConfig(OpenAIGPTConfig):
num_search_queries = perplexity_usage.get("search_queries")
if num_search_queries is None:
num_search_queries = raw_response_json.get("search_queries")
-
+
# Create or update prompt_tokens_details to include web search requests and citation tokens
if citation_tokens > 0 or (
num_search_queries is not None and num_search_queries > 0
):
if usage.prompt_tokens_details is None:
usage.prompt_tokens_details = PromptTokensDetailsWrapper()
-
+
# Store citation tokens count for cost calculation
if citation_tokens > 0:
- setattr(usage, "citation_tokens", citation_tokens)
-
+ usage.citation_tokens = citation_tokens
+
# Store search queries count in the standard web_search_requests field
if num_search_queries is not None and num_search_queries > 0:
usage.prompt_tokens_details.web_search_requests = num_search_queries
@@ -248,4 +290,35 @@ class PerplexityChatConfig(OpenAIGPTConfig):
if citations:
setattr(model_response, "citations", citations)
if search_results:
- setattr(model_response, "search_results", search_results)
\ No newline at end of file
+ setattr(model_response, "search_results", search_results)
+
+ def add_cost_to_usage(self, model_response: ModelResponse, raw_response_json: dict) -> None:
+ """Add the cost to the usage object."""
+ try:
+ usage_data = raw_response_json.get("usage")
+ if usage_data:
+ # Try different possible cost field locations
+ response_cost = None
+
+ # Check if cost is directly in usage (flat structure)
+ if "total_cost" in usage_data:
+ response_cost = usage_data["total_cost"]
+ # Check if cost is nested (cost.total_cost structure)
+ elif "cost" in usage_data and isinstance(usage_data["cost"], dict):
+ response_cost = usage_data["cost"].get("total_cost")
+ # Check if cost is a simple value
+ elif "cost" in usage_data:
+ response_cost = usage_data["cost"]
+
+ if response_cost is not None:
+ # Store cost in hidden params for the cost calculator to use
+ if not hasattr(model_response, "_hidden_params"):
+ model_response._hidden_params = {}
+ if "additional_headers" not in model_response._hidden_params:
+ model_response._hidden_params["additional_headers"] = {}
+ model_response._hidden_params["additional_headers"][
+ "llm_provider-x-litellm-response-cost"
+ ] = float(response_cost)
+ except (ValueError, TypeError, KeyError) as e:
+ verbose_logger.debug(f"Error adding cost to usage: {e}")
+ # If we can't extract cost, continue without it - don't fail the response
diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py
index a750b5e985f..2c534577366 100644
--- a/litellm/llms/vertex_ai/common_utils.py
+++ b/litellm/llms/vertex_ai/common_utils.py
@@ -27,6 +27,7 @@ class VertexAIError(BaseLLMException):
class VertexAIModelRoute(str, Enum):
"""Enum for Vertex AI model routing"""
+
PARTNER_MODELS = "partner_models"
GEMINI = "gemini"
GEMMA = "gemma"
@@ -34,27 +35,29 @@ class VertexAIModelRoute(str, Enum):
NON_GEMINI = "non_gemini"
-def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None) -> VertexAIModelRoute:
+def get_vertex_ai_model_route(
+ model: str, litellm_params: Optional[dict] = None
+) -> VertexAIModelRoute:
"""
Determine which handler to use for a Vertex AI model based on the model name.
-
+
Args:
model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "openai/gpt-oss-120b")
litellm_params: Optional litellm parameters dict that may contain base_model for routing
-
+
Returns:
VertexAIModelRoute: The route enum indicating which handler should be used
-
+
Examples:
>>> get_vertex_ai_model_route("llama3-405b")
VertexAIModelRoute.PARTNER_MODELS
-
+
>>> get_vertex_ai_model_route("gemini-pro")
VertexAIModelRoute.GEMINI
-
+
>>> get_vertex_ai_model_route("gemma/gemma-3-12b-it")
VertexAIModelRoute.GEMMA
-
+
>>> get_vertex_ai_model_route("openai/gpt-oss-120b")
VertexAIModelRoute.MODEL_GARDEN
"""
@@ -66,23 +69,23 @@ def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None)
if litellm_params and litellm_params.get("base_model") is not None:
if "gemini" in litellm_params["base_model"]:
return VertexAIModelRoute.GEMINI
-
+
# Check for partner models (llama, mistral, claude, etc.)
if VertexAIPartnerModels.is_vertex_partner_model(model=model):
return VertexAIModelRoute.PARTNER_MODELS
-
+
# Check for gemma models
if "gemma/" in model:
return VertexAIModelRoute.GEMMA
-
+
# Check for model garden openai models
if "openai" in model:
return VertexAIModelRoute.MODEL_GARDEN
-
+
# Check for gemini models
if "gemini" in model:
return VertexAIModelRoute.GEMINI
-
+
# Default to non-gemini (legacy vertex models like chat-bison, text-bison, etc.)
return VertexAIModelRoute.NON_GEMINI
@@ -253,8 +256,10 @@ def _check_text_in_content(parts: List[PartType]) -> bool:
def _fix_enum_empty_strings(schema, depth=0):
"""Fix empty strings in enum values by replacing them with None. Gemini doesn't accept empty strings in enums."""
if depth > DEFAULT_MAX_RECURSE_DEPTH:
- raise ValueError(f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while processing schema.")
-
+ raise ValueError(
+ f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while processing schema."
+ )
+
if "enum" in schema and isinstance(schema["enum"], list):
schema["enum"] = [None if value == "" else value for value in schema["enum"]]
@@ -529,19 +534,18 @@ def _convert_vertex_datetime_to_openai_datetime(vertex_datetime: str) -> int:
def _convert_schema_types(schema, depth=0):
"""
Convert type arrays and lowercase types for Vertex AI compatibility.
-
- Transforms OpenAI-style schemas to Vertex AI format by converting type arrays
+
+ Transforms OpenAI-style schemas to Vertex AI format by converting type arrays
like ["string", "number"] to anyOf format and converting all types to uppercase.
"""
if depth > DEFAULT_MAX_RECURSE_DEPTH:
raise ValueError(
f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while processing schema. Please check the schema for excessive nesting."
)
-
+
if not isinstance(schema, dict):
return
-
# Handle type field
if "type" in schema:
type_val = schema["type"]
@@ -553,7 +557,7 @@ def _convert_schema_types(schema, depth=0):
schema["type"] = type_val[0]
elif isinstance(type_val, str):
schema["type"] = type_val
-
+
# Recursively process nested properties, items, and anyOf
for key in ["properties", "items", "anyOf"]:
if key in schema:
@@ -567,6 +571,7 @@ def _convert_schema_types(schema, depth=0):
for anyof_schema in value:
_convert_schema_types(anyof_schema, depth + 1)
+
def get_vertex_project_id_from_url(url: str) -> Optional[str]:
"""
Get the vertex project id from the url
@@ -665,17 +670,18 @@ def is_global_only_vertex_model(model: str) -> bool:
return False
return "global" in supported_regions
-class VertexAIModelInfo(BaseLLMModelInfo):
+
+class VertexAIModelInfo(BaseLLMModelInfo):
def get_token_counter(self) -> Optional[BaseTokenCounter]:
"""
Factory method to create a token counter for this provider.
-
+
Returns:
Optional TokenCounterInterface implementation for this provider,
or None if token counting is not supported.
"""
return VertexAITokenCounter()
-
+
def validate_environment(
self,
headers: dict,
@@ -687,7 +693,7 @@ class VertexAIModelInfo(BaseLLMModelInfo):
api_base: Optional[str] = None,
) -> dict:
raise NotImplementedError("Vertex AI models are not supported yet")
-
+
def get_models(
self, api_key: Optional[str] = None, api_base: Optional[str] = None
) -> List[str]:
@@ -706,8 +712,6 @@ class VertexAIModelInfo(BaseLLMModelInfo):
) -> Optional[str]:
raise NotImplementedError("Vertex AI models are not supported yet")
-
-
@staticmethod
def get_base_model(model: str) -> Optional[str]:
"""
@@ -721,13 +725,15 @@ class VertexAIModelInfo(BaseLLMModelInfo):
class VertexAITokenCounter(BaseTokenCounter):
"""Token counter implementation for Google AI Studio provider."""
+
def should_use_token_counting_api(
- self,
+ self,
custom_llm_provider: Optional[str] = None,
) -> bool:
from litellm.types.utils import LlmProviders
+
return custom_llm_provider == LlmProviders.VERTEX_AI.value
-
+
async def count_tokens(
self,
model_to_use: str,
@@ -738,25 +744,68 @@ class VertexAITokenCounter(BaseTokenCounter):
) -> Optional[TokenCountResponse]:
import copy
- from litellm.llms.vertex_ai.count_tokens.handler import VertexAITokenCounter
- deployment = deployment or {}
- count_tokens_params_request = copy.deepcopy(deployment.get("litellm_params", {}))
- count_tokens_params = {
- "model": model_to_use,
- "contents": contents,
- }
- count_tokens_params_request.update(count_tokens_params)
- result = await VertexAITokenCounter().acount_tokens(
- **count_tokens_params_request,
+ from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
+ VertexAIPartnerModels,
)
-
- if result is not None:
- return TokenCountResponse(
- total_tokens=result.get("totalTokens", 0),
- request_model=request_model,
- model_used=model_to_use,
- tokenizer_type=result.get("tokenizer_used", ""),
- original_response=result,
+
+ deployment = deployment or {}
+ count_tokens_params_request = copy.deepcopy(
+ deployment.get("litellm_params", {})
+ )
+
+ # Check if this is a partner model (Claude, Mistral, etc.)
+ if VertexAIPartnerModels.is_vertex_partner_model(model_to_use):
+ # Use partner models token counter
+ partner_models_handler = VertexAIPartnerModels()
+
+ # Extract vertex-specific params from litellm_params
+ vertex_project = count_tokens_params_request.get(
+ "vertex_project"
+ ) or count_tokens_params_request.get("vertex_ai_project")
+ vertex_location = count_tokens_params_request.get(
+ "vertex_location"
+ ) or count_tokens_params_request.get("vertex_ai_location")
+ vertex_credentials = count_tokens_params_request.get(
+ "vertex_credentials"
+ ) or count_tokens_params_request.get("vertex_ai_credentials")
+
+ result = await partner_models_handler.count_tokens(
+ model=model_to_use,
+ messages=messages or [],
+ litellm_params=count_tokens_params_request,
+ vertex_project=vertex_project,
+ vertex_location=vertex_location,
+ vertex_credentials=vertex_credentials,
)
-
- return None
\ No newline at end of file
+
+ if result is not None:
+ return TokenCountResponse(
+ total_tokens=result.get("input_tokens", 0),
+ request_model=request_model,
+ model_used=model_to_use,
+ tokenizer_type=result.get("tokenizer_used", ""),
+ original_response=result,
+ )
+ else:
+ # Use standard Vertex AI (Gemini) token counter
+ from litellm.llms.vertex_ai.count_tokens.handler import VertexAITokenCounter
+
+ count_tokens_params = {
+ "model": model_to_use,
+ "contents": contents,
+ }
+ count_tokens_params_request.update(count_tokens_params)
+ result = await VertexAITokenCounter().acount_tokens(
+ **count_tokens_params_request,
+ )
+
+ if result is not None:
+ return TokenCountResponse(
+ total_tokens=result.get("totalTokens", 0),
+ request_model=request_model,
+ model_used=model_to_use,
+ tokenizer_type=result.get("tokenizer_used", ""),
+ original_response=result,
+ )
+
+ return None
diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py
index 4f84586cfc1..08c91a6fad1 100644
--- a/litellm/llms/vertex_ai/gemini/transformation.py
+++ b/litellm/llms/vertex_ai/gemini/transformation.py
@@ -473,7 +473,7 @@ def _transform_request_body(
labels = {k: v for k, v in rm.items() if isinstance(v, str)}
filtered_params = {
- k: v for k, v in optional_params.items() if k in config_fields
+ k: v for k, v in optional_params.items() if _get_equivalent_key(k, set(config_fields))
}
generation_config: Optional[GenerationConfig] = GenerationConfig(
diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py
new file mode 100644
index 00000000000..008957c87a2
--- /dev/null
+++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py
@@ -0,0 +1 @@
+# Count tokens handler for Vertex AI Partner Models (Anthropic, Mistral, etc.)
diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py
new file mode 100644
index 00000000000..0f073d02269
--- /dev/null
+++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py
@@ -0,0 +1,157 @@
+"""
+Token counter for Vertex AI Partner Models (Anthropic Claude, Mistral, etc.)
+
+This handler provides token counting for partner models hosted on Vertex AI.
+Unlike Gemini models which use Google's token counting API, partner models use
+their respective publisher-specific count-tokens endpoints.
+"""
+from typing import Any, Dict, Optional
+
+from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
+from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
+from litellm.types.llms.vertex_ai import VertexPartnerProvider
+
+
+class VertexAIPartnerModelsTokenCounter(VertexBase):
+ """
+ Token counter for Vertex AI Partner Models.
+
+ Handles token counting for models like Claude (Anthropic), Mistral, etc.
+ that are available through Vertex AI's partner model program.
+ """
+
+ def _get_publisher_for_model(self, model: str) -> str:
+ """
+ Determine the publisher name for the given model.
+
+ Args:
+ model: The model name (e.g., "claude-3-5-sonnet-20241022")
+
+ Returns:
+ Publisher name to use in the Vertex AI endpoint URL
+
+ Raises:
+ ValueError: If the model is not a recognized partner model
+ """
+ if "claude" in model:
+ return "anthropic"
+ elif "mistral" in model or "codestral" in model:
+ return "mistralai"
+ elif "llama" in model or "meta/" in model:
+ return "meta"
+ else:
+ raise ValueError(f"Unknown partner model: {model}")
+
+ def _build_count_tokens_endpoint(
+ self,
+ model: str,
+ project_id: str,
+ vertex_location: str,
+ api_base: Optional[str] = None,
+ ) -> str:
+ """
+ Build the count-tokens endpoint URL for a partner model.
+
+ Args:
+ model: The model name
+ project_id: Google Cloud project ID
+ vertex_location: Vertex AI location (e.g., "us-east5")
+ api_base: Optional custom API base URL
+
+ Returns:
+ Full endpoint URL for the count-tokens API
+ """
+ publisher = self._get_publisher_for_model(model)
+
+ # Use custom api_base if provided, otherwise construct default
+ if api_base:
+ base_url = api_base
+ else:
+ base_url = f"https://{vertex_location}-aiplatform.googleapis.com"
+
+ # Construct the count-tokens endpoint
+ # Format: /v1/projects/{project}/locations/{location}/publishers/{publisher}/models/count-tokens:rawPredict
+ endpoint = (
+ f"{base_url}/v1/projects/{project_id}/locations/{vertex_location}/"
+ f"publishers/{publisher}/models/count-tokens:rawPredict"
+ )
+
+ return endpoint
+
+ async def handle_count_tokens_request(
+ self,
+ model: str,
+ request_data: Dict[str, Any],
+ litellm_params: Dict[str, Any],
+ ) -> Dict[str, Any]:
+ """
+ Handle token counting request for a Vertex AI partner model.
+
+ Args:
+ model: The model name
+ request_data: Request payload (Anthropic Messages API format)
+ litellm_params: LiteLLM parameters containing credentials, project, location
+
+ Returns:
+ Dict containing token count information
+
+ Raises:
+ ValueError: If required parameters are missing or invalid
+ """
+ # Validate request
+ if "messages" not in request_data:
+ raise ValueError("messages required for token counting")
+
+ # Extract Vertex AI credentials and settings
+ vertex_credentials = self.get_vertex_ai_credentials(litellm_params)
+ vertex_project = self.get_vertex_ai_project(litellm_params)
+ vertex_location = self.get_vertex_ai_location(litellm_params)
+
+ # Get access token and resolved project ID
+ access_token, project_id = await self._ensure_access_token_async(
+ credentials=vertex_credentials,
+ project_id=vertex_project,
+ custom_llm_provider="vertex_ai",
+ )
+
+ # Build the endpoint URL
+ endpoint_url = self._build_count_tokens_endpoint(
+ model=model,
+ project_id=project_id,
+ vertex_location=vertex_location or "us-central1",
+ api_base=litellm_params.get("api_base"),
+ )
+
+ # Prepare headers
+ headers = {"Authorization": f"Bearer {access_token}"}
+
+ # Get async HTTP client
+ from litellm import LlmProviders
+
+ async_client = get_async_httpx_client(llm_provider=LlmProviders.VERTEX_AI)
+
+ # Make the request
+ # Note: Partner models (especially Claude) accept Anthropic Messages API format directly
+ response = await async_client.post(
+ endpoint_url,
+ headers=headers,
+ json=request_data,
+ timeout=30.0,
+ )
+
+ # Check for errors
+ if response.status_code != 200:
+ error_text = response.text
+ raise ValueError(
+ f"Token counting request failed with status {response.status_code}: {error_text}"
+ )
+
+ # Parse response
+ result = response.json()
+
+ # Return token count
+ # Vertex AI Anthropic returns: {"input_tokens": 123}
+ return {
+ "input_tokens": result.get("input_tokens", 0),
+ "tokenizer_used": "vertex_ai_partner_models",
+ }
diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py
index ea29970f0aa..85b1a6bc0db 100644
--- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py
+++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py
@@ -28,6 +28,7 @@ class VertexAIError(Exception):
self.message
) # Call the base class constructor with the parameters it needs
+
class PartnerModelPrefixes(str, Enum):
META_PREFIX = "meta/"
DEEPSEEK_PREFIX = "deepseek-ai"
@@ -64,7 +65,7 @@ class VertexAIPartnerModels(VertexBase):
):
return True
return False
-
+
@staticmethod
def should_use_openai_handler(model: str):
OPENAI_LIKE_VERTEX_PROVIDERS = [
@@ -258,3 +259,77 @@ class VertexAIPartnerModels(VertexBase):
if hasattr(e, "status_code"):
raise e
raise VertexAIError(status_code=500, message=str(e))
+
+ async def count_tokens(
+ self,
+ model: str,
+ messages: list,
+ litellm_params: dict,
+ vertex_project=None,
+ vertex_location=None,
+ vertex_credentials=None,
+ ):
+ """
+ Count tokens for Vertex AI partner models (Anthropic Claude, Mistral, etc.)
+
+ Args:
+ model: The model name (e.g., "claude-3-5-sonnet-20241022")
+ messages: List of messages in Anthropic Messages API format
+ litellm_params: LiteLLM parameters dict
+ vertex_project: Optional Google Cloud project ID
+ vertex_location: Optional Vertex AI location
+ vertex_credentials: Optional Vertex AI credentials
+
+ Returns:
+ Dict containing token count information
+ """
+ try:
+ import vertexai
+ except Exception as e:
+ raise VertexAIError(
+ status_code=400,
+ message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""",
+ )
+
+ if not (
+ hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models")
+ ):
+ raise VertexAIError(
+ status_code=400,
+ message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""",
+ )
+
+ try:
+ from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler import (
+ VertexAIPartnerModelsTokenCounter,
+ )
+
+ # Prepare request data in Anthropic Messages API format
+ request_data = {
+ "model": model,
+ "messages": messages,
+ }
+
+ # Prepare litellm_params with credentials
+ _litellm_params = litellm_params.copy()
+ if vertex_project:
+ _litellm_params["vertex_project"] = vertex_project
+ if vertex_location:
+ _litellm_params["vertex_location"] = vertex_location
+ if vertex_credentials:
+ _litellm_params["vertex_credentials"] = vertex_credentials
+
+ # Call the token counter
+ token_counter = VertexAIPartnerModelsTokenCounter()
+ result = await token_counter.handle_count_tokens_request(
+ model=model,
+ request_data=request_data,
+ litellm_params=_litellm_params,
+ )
+
+ return result
+
+ except Exception as e:
+ if hasattr(e, "status_code"):
+ raise e
+ raise VertexAIError(status_code=500, message=str(e))
diff --git a/litellm/main.py b/litellm/main.py
index b7af4e8d39c..9cae34d1678 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -2033,11 +2033,36 @@ def completion( # type: ignore # noqa: PLR0915
logging.post_call(
input=messages, api_key=api_key, original_response=response
)
+ elif custom_llm_provider == "perplexity":
+ response = base_llm_http_handler.completion(
+ model=model,
+ messages=messages,
+ headers=headers,
+ model_response=model_response,
+ api_key=api_key,
+ api_base=api_base,
+ acompletion=acompletion,
+ logging_obj=logging,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ shared_session=shared_session,
+ timeout=timeout,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ encoding=encoding,
+ stream=stream,
+ provider_config=provider_config,
+ )
+
+ ## LOGGING - Call after response has been processed by transform_response
+ logging.post_call(
+ input=messages, api_key=api_key, original_response=response
+ )
+
elif (
model in litellm.open_ai_chat_completion_models
or custom_llm_provider == "custom_openai"
or custom_llm_provider == "deepinfra"
- or custom_llm_provider == "perplexity"
or custom_llm_provider == "nvidia_nim"
or custom_llm_provider == "cerebras"
or custom_llm_provider == "baseten"
diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/api-reference.html
rename to litellm/proxy/_experimental/out/api-reference/index.html
diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails.html
deleted file mode 100644
index 99ec72d9c9d..00000000000
--- a/litellm/proxy/_experimental/out/guardrails.html
+++ /dev/null
@@ -1 +0,0 @@
-
LiteLLM Dashboard
\ No newline at end of file
diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/logs.html
rename to litellm/proxy/_experimental/out/logs/index.html
diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model-hub.html
rename to litellm/proxy/_experimental/out/model-hub/index.html
diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/model_hub_table.html
rename to litellm/proxy/_experimental/out/model_hub_table/index.html
diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/models-and-endpoints.html
rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html
diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html
deleted file mode 100644
index a03ee54b8f7..00000000000
--- a/litellm/proxy/_experimental/out/onboarding.html
+++ /dev/null
@@ -1 +0,0 @@
-LiteLLM Dashboard
\ No newline at end of file
diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/organizations.html
rename to litellm/proxy/_experimental/out/organizations/index.html
diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/teams.html
rename to litellm/proxy/_experimental/out/teams/index.html
diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/test-key.html
rename to litellm/proxy/_experimental/out/test-key/index.html
diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/usage.html
rename to litellm/proxy/_experimental/out/usage/index.html
diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/users.html
rename to litellm/proxy/_experimental/out/users/index.html
diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html
similarity index 100%
rename from litellm/proxy/_experimental/out/virtual-keys.html
rename to litellm/proxy/_experimental/out/virtual-keys/index.html
diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml
index 1f57cd9f159..9ef9812dd49 100644
--- a/litellm/proxy/_new_secret_config.yaml
+++ b/litellm/proxy/_new_secret_config.yaml
@@ -12,4 +12,10 @@ vector_store_registry:
vector_store_id: "litellm-docs_1761094140318"
custom_llm_provider: "vertex_ai/search_api"
vertex_project: "test-vector-store-db"
- vertex_location: "global"
\ No newline at end of file
+ vertex_location: "global"
+ - vector_store_name: "milvus-litellm-website-knowledgebase"
+ litellm_params:
+ vector_store_id: "can-be-anything"
+ custom_llm_provider: "milvus"
+ api_base: os.environ/MILVUS_API_BASE
+ api_key: os.environ/MILVUS_API_KEY
\ No newline at end of file
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 20a78d839a3..0a6240b0f9f 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -346,6 +346,7 @@ class LiteLLMRoutes(enum.Enum):
"/eu.assemblyai",
"/vllm",
"/mistral",
+ "/milvus",
]
#########################################################
@@ -764,9 +765,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
allowed_cache_controls: Optional[list] = []
config: Optional[dict] = {}
permissions: Optional[dict] = {}
- model_max_budget: Optional[dict] = (
- {}
- ) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
+ model_max_budget: Optional[
+ dict
+ ] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
@@ -1192,12 +1193,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
- allowed_model_region: Optional[AllowedModelRegion] = (
- None # require all user requests to use models in this specific region
- )
- default_model: Optional[str] = (
- None # if no equivalent model in allowed region - default all requests to this model
- )
+ allowed_model_region: Optional[
+ AllowedModelRegion
+ ] = None # require all user requests to use models in this specific region
+ default_model: Optional[
+ str
+ ] = None # if no equivalent model in allowed region - default all requests to this model
@model_validator(mode="before")
@classmethod
@@ -1219,12 +1220,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
- allowed_model_region: Optional[AllowedModelRegion] = (
- None # require all user requests to use models in this specific region
- )
- default_model: Optional[str] = (
- None # if no equivalent model in allowed region - default all requests to this model
- )
+ allowed_model_region: Optional[
+ AllowedModelRegion
+ ] = None # require all user requests to use models in this specific region
+ default_model: Optional[
+ str
+ ] = None # if no equivalent model in allowed region - default all requests to this model
class DeleteCustomerRequest(LiteLLMPydanticObjectBase):
@@ -1308,15 +1309,15 @@ class NewTeamRequest(TeamBase):
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
model_tpm_limit: Optional[Dict[str, int]] = None
- team_member_budget: Optional[float] = (
- None # allow user to set a budget for all team members
- )
- team_member_rpm_limit: Optional[int] = (
- None # allow user to set RPM limit for all team members
- )
- team_member_tpm_limit: Optional[int] = (
- None # allow user to set TPM limit for all team members
- )
+ team_member_budget: Optional[
+ float
+ ] = None # allow user to set a budget for all team members
+ team_member_rpm_limit: Optional[
+ int
+ ] = None # allow user to set RPM limit for all team members
+ team_member_tpm_limit: Optional[
+ int
+ ] = None # allow user to set TPM limit for all team members
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
@@ -1400,9 +1401,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
- callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
- "success_and_failure"
- )
+ callback_type: Optional[
+ Literal["success", "failure", "success_and_failure"]
+ ] = "success_and_failure"
callback_vars: Dict[str, str]
@model_validator(mode="before")
@@ -1687,9 +1688,9 @@ class ConfigList(LiteLLMPydanticObjectBase):
stored_in_db: Optional[bool]
field_default_value: Any
premium_field: bool = False
- nested_fields: Optional[List[FieldDetail]] = (
- None # For nested dictionary or Pydantic fields
- )
+ nested_fields: Optional[
+ List[FieldDetail]
+ ] = None # For nested dictionary or Pydantic fields
class UserHeaderMapping(LiteLLMPydanticObjectBase):
@@ -2069,9 +2070,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase):
budget_id: Optional[str] = None
created_at: datetime
updated_at: datetime
- user: Optional[Any] = (
- None # You might want to replace 'Any' with a more specific type if available
- )
+ user: Optional[
+ Any
+ ] = None # You might want to replace 'Any' with a more specific type if available
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
model_config = ConfigDict(protected_namespaces=())
@@ -2520,7 +2521,7 @@ class SpendLogsMetadata(TypedDict):
applied_guardrails: Optional[List[str]]
mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall]
vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]]
- guardrail_information: Optional[StandardLoggingGuardrailInformation]
+ guardrail_information: Optional[list[StandardLoggingGuardrailInformation]]
status: StandardLoggingPayloadStatus
proxy_server_request: Optional[str]
batch_models: Optional[List[str]]
@@ -3004,9 +3005,9 @@ class TeamModelDeleteRequest(BaseModel):
# Organization Member Requests
class OrganizationMemberAddRequest(OrgMemberAddRequest):
organization_id: str
- max_budget_in_organization: Optional[float] = (
- None # Users max budget within the organization
- )
+ max_budget_in_organization: Optional[
+ float
+ ] = None # Users max budget within the organization
class OrganizationMemberDeleteRequest(MemberDeleteRequest):
@@ -3219,9 +3220,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase):
Maps provider names to their budget configs.
"""
- providers: Dict[str, ProviderBudgetResponseObject] = (
- {}
- ) # Dictionary mapping provider names to their budget configurations
+ providers: Dict[
+ str, ProviderBudgetResponseObject
+ ] = {} # Dictionary mapping provider names to their budget configurations
class ProxyStateVariables(TypedDict):
@@ -3355,9 +3356,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
enforce_rbac: bool = False
roles_jwt_field: Optional[str] = None # v2 on role mappings
role_mappings: Optional[List[RoleMapping]] = None
- object_id_jwt_field: Optional[str] = (
- None # can be either user / team, inferred from the role mapping
- )
+ object_id_jwt_field: Optional[
+ str
+ ] = None # can be either user / team, inferred from the role mapping
scope_mappings: Optional[List[ScopeMapping]] = None
enforce_scope_based_access: bool = False
enforce_team_based_model_access: bool = False
diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py
index dec399c7f74..a4c1581a71b 100644
--- a/litellm/proxy/management_endpoints/ui_sso.py
+++ b/litellm/proxy/management_endpoints/ui_sso.py
@@ -24,6 +24,7 @@ from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.caching import DualCache
from litellm.constants import MAX_SPENDLOG_ROWS_TO_QUERY
+from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
@@ -273,12 +274,14 @@ def generic_response_convertor(
all_teams.extend(team_ids)
return CustomOpenID(
- id=response.get(generic_user_id_attribute_name),
- display_name=response.get(generic_user_display_name_attribute_name),
- email=response.get(generic_user_email_attribute_name),
- first_name=response.get(generic_user_first_name_attribute_name),
- last_name=response.get(generic_user_last_name_attribute_name),
- provider=response.get(generic_provider_attribute_name),
+ id=get_nested_value(response, generic_user_id_attribute_name),
+ display_name=get_nested_value(
+ response, generic_user_display_name_attribute_name
+ ),
+ email=get_nested_value(response, generic_user_email_attribute_name),
+ first_name=get_nested_value(response, generic_user_first_name_attribute_name),
+ last_name=get_nested_value(response, generic_user_last_name_attribute_name),
+ provider=get_nested_value(response, generic_provider_attribute_name),
team_ids=all_teams,
user_role=None,
)
@@ -1081,7 +1084,7 @@ class SSOAuthenticationHandler:
raise ValueError(
"Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
)
-
+
@staticmethod
async def get_generic_sso_redirect_response(
generic_sso: Any,
@@ -1094,6 +1097,7 @@ class SSOAuthenticationHandler:
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from litellm.proxy.proxy_server import user_api_key_cache
+
with generic_sso:
# TODO: state should be a random string and added to the user session with cookie
# or a cryptographicly signed state that we can verify stateless
@@ -1133,22 +1137,24 @@ class SSOAuthenticationHandler:
if pkce_params:
parsed_url = urlparse(str(redirect_response.headers["location"]))
query_params = parse_qs(parsed_url.query)
-
+
# Add PKCE parameters
for key, value in pkce_params.items():
query_params[key] = [value]
-
+
# Reconstruct the URL with PKCE parameters
new_query = urlencode(query_params, doseq=True)
- new_url = urlunparse((
- parsed_url.scheme,
- parsed_url.netloc,
- parsed_url.path,
- parsed_url.params,
- new_query,
- parsed_url.fragment
- ))
-
+ new_url = urlunparse(
+ (
+ parsed_url.scheme,
+ parsed_url.netloc,
+ parsed_url.path,
+ parsed_url.params,
+ new_query,
+ parsed_url.fragment,
+ )
+ )
+
# Update the redirect response
redirect_response.headers["location"] = new_url
verbose_proxy_logger.debug(
@@ -1175,7 +1181,7 @@ class SSOAuthenticationHandler:
generic_authorization_endpoint: Authorization endpoint URL
Returns:
- Tuple[dict, Optional[str]]:
+ Tuple[dict, Optional[str]]:
- Redirect parameters for SSO login (may include PKCE params)
- code_verifier (if PKCE is enabled, None otherwise)
"""
@@ -1202,7 +1208,9 @@ class SSOAuthenticationHandler:
# Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security
use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true"
if use_pkce:
- code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params()
+ code_verifier, code_challenge = (
+ SSOAuthenticationHandler.generate_pkce_params()
+ )
redirect_params["code_challenge"] = code_challenge
redirect_params["code_challenge_method"] = "S256"
verbose_proxy_logger.debug(
@@ -1693,11 +1701,10 @@ class SSOAuthenticationHandler:
redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
return redirect_response
-
@staticmethod
def prepare_token_exchange_parameters(
- request: Request,
+ request: Request,
generic_include_client_id: bool,
) -> dict:
"""
@@ -1712,50 +1719,54 @@ class SSOAuthenticationHandler:
"""
# Prepare token exchange parameters
token_params = {"include_client_id": generic_include_client_id}
-
+
# Retrieve PKCE code_verifier if PKCE was used in authorization
query_params = dict(request.query_params)
state = query_params.get("state")
if state:
from litellm.proxy.proxy_server import user_api_key_cache
-
+
cache_key = f"pkce_verifier:{state}"
code_verifier = user_api_key_cache.get_cache(key=cache_key)
-
+
if code_verifier:
# Add code_verifier to token exchange parameters
token_params["code_verifier"] = code_verifier
verbose_proxy_logger.debug(
"PKCE code_verifier retrieved and will be included in token exchange"
)
-
+
# Clean up the cache entry (single-use verifier)
user_api_key_cache.delete_cache(key=cache_key)
return token_params
-
@staticmethod
def generate_pkce_params() -> Tuple[str, str]:
"""
Generate PKCE (Proof Key for Code Exchange) parameters for OAuth 2.0.
-
+
Returns:
Tuple[str, str]: (code_verifier, code_challenge)
- code_verifier: Random 43-128 character string (we use 43 for efficiency)
- code_challenge: Base64-URL-encoded SHA256 hash of the code_verifier
-
+
Reference: https://datatracker.ietf.org/doc/html/rfc7636
"""
# Generate a cryptographically random code_verifier (43 characters)
# Using 32 random bytes which becomes 43 characters when base64-url-encoded
- code_verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).decode('utf-8').rstrip('=')
-
- # Generate code_challenge using S256 method (SHA256)
- code_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest()
- code_challenge = base64.urlsafe_b64encode(code_challenge_bytes).decode('utf-8').rstrip('=')
-
- return code_verifier, code_challenge
+ code_verifier = (
+ base64.urlsafe_b64encode(secrets.token_bytes(32))
+ .decode("utf-8")
+ .rstrip("=")
+ )
+ # Generate code_challenge using S256 method (SHA256)
+ code_challenge_bytes = hashlib.sha256(code_verifier.encode("utf-8")).digest()
+ code_challenge = (
+ base64.urlsafe_b64encode(code_challenge_bytes).decode("utf-8").rstrip("=")
+ )
+
+ return code_verifier, code_challenge
class MicrosoftSSOHandler:
diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py
index cd8f6c38000..b7a82bb8d60 100644
--- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py
+++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py
@@ -24,6 +24,7 @@ from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
+ _safe_set_request_parsed_body,
get_form_data,
get_request_body,
)
@@ -39,6 +40,7 @@ from litellm.proxy.vector_store_endpoints.utils import (
is_allowed_to_call_vector_store_endpoint,
)
from litellm.secret_managers.main import get_secret_str
+from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
from .passthrough_endpoint_router import PassthroughEndpointRouter
@@ -161,12 +163,12 @@ async def llm_passthrough_factory_proxy_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers=auth_headers,
+ is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- stream=is_streaming_request, # type: ignore
)
return received_value
@@ -230,13 +232,12 @@ async def gemini_proxy_route(
endpoint=endpoint,
target=str(updated_url),
custom_llm_provider="gemini",
+ is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- query_params=merged_params, # type: ignore
- stream=is_streaming_request, # type: ignore
)
return received_value
@@ -283,12 +284,12 @@ async def cohere_proxy_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"Authorization": "Bearer {}".format(cohere_api_key)},
+ is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- stream=is_streaming_request, # type: ignore
)
return received_value
@@ -412,12 +413,141 @@ async def mistral_proxy_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"Authorization": "Bearer {}".format(mistral_api_key)},
+ is_streaming_request=is_streaming_request,
+ ) # dynamically construct pass-through endpoint based on incoming path
+ received_value = await endpoint_func(
+ request,
+ fastapi_response,
+ user_api_key_dict,
+ )
+
+ return received_value
+
+
+@router.api_route(
+ "/milvus/{endpoint:path}",
+ methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
+ tags=["Milvus Pass-through", "pass-through"],
+)
+async def milvus_proxy_route(
+ endpoint: str,
+ request: Request,
+ fastapi_response: Response,
+ user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
+):
+ """
+ Enable using Milvus `/vectors` endpoint as a pass-through endpoint.
+ """
+
+ provider_config = ProviderConfigManager.get_provider_vector_stores_config(
+ provider=LlmProviders.MILVUS
+ )
+ if not provider_config:
+ raise HTTPException(
+ status_code=500,
+ detail="Unable to find Milvus vector store config.",
+ )
+
+ # check if managed vector store index is used
+ request_body = await get_request_body(request)
+
+ # check collectionName
+ collection_name = cast(Optional[str], request_body.get("collectionName"))
+ extra_headers = {}
+ base_target_url: Optional[str] = None
+ if not collection_name:
+ raise HTTPException(
+ status_code=400,
+ detail=f"Collection name is required. Got {request_body}",
+ )
+
+ if not litellm.vector_store_index_registry or not litellm.vector_store_registry:
+ raise HTTPException(
+ status_code=500,
+ detail="Unable to find Milvus vector store index registry or vector store registry.",
+ )
+
+ # check if vector store index
+ is_vector_store_index = litellm.vector_store_index_registry.is_vector_store_index(
+ vector_store_index_name=collection_name
+ )
+
+ if not is_vector_store_index:
+ raise HTTPException(
+ status_code=400,
+ detail=f"Collection {collection_name} is not a litellm managed vector store index. Only litellm managed vector store indexes are supported.",
+ )
+
+ is_allowed_to_call_vector_store_endpoint(
+ index_name=collection_name,
+ provider=LlmProviders.MILVUS,
+ request=request,
+ user_api_key_dict=user_api_key_dict,
+ )
+ # get the vector store name from index registry
+
+ index_object = (
+ (
+ litellm.vector_store_index_registry.get_vector_store_index_by_name(
+ vector_store_index_name=collection_name
+ )
+ )
+ if litellm.vector_store_index_registry is not None
+ else None
+ )
+ if index_object is None:
+ raise Exception(f"Vector store index not found for {collection_name}")
+
+ vector_store_name = index_object.litellm_params.vector_store_name
+ vector_store_index = index_object.litellm_params.vector_store_index
+
+ request_body["collectionName"] = vector_store_index
+
+ # Update the request object with the modified collection name
+ _safe_set_request_parsed_body(request, request_body)
+
+ vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry_by_name(
+ vector_store_name=vector_store_name
+ )
+ if vector_store is None:
+ raise Exception(f"Vector store not found for {vector_store_name}")
+ litellm_params = vector_store.get("litellm_params") or {}
+ auth_credentials = provider_config.get_auth_credentials(
+ litellm_params=litellm_params
+ )
+
+ extra_headers = auth_credentials.get("headers") or {}
+
+ litellm_params = vector_store.get("litellm_params") or {}
+
+ base_target_url = provider_config.get_complete_url(
+ api_base=litellm_params.get("api_base"), litellm_params=litellm_params
+ )
+
+ if base_target_url is None:
+ raise Exception(
+ f"api_base not found in vector store configuration for {vector_store_name}"
+ )
+
+ encoded_endpoint = httpx.URL(endpoint).path
+
+ # Ensure endpoint starts with '/' for proper URL construction
+ if not encoded_endpoint.startswith("/"):
+ encoded_endpoint = "/" + encoded_endpoint
+
+ # Construct the full target URL using httpx
+ base_url = httpx.URL(base_target_url)
+ updated_url = base_url.copy_with(path=encoded_endpoint)
+ ## CREATE PASS-THROUGH
+ endpoint_func = create_pass_through_route(
+ endpoint=endpoint,
+ target=str(updated_url),
+ custom_headers=extra_headers,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- stream=is_streaming_request, # type: ignore
)
return received_value
@@ -475,12 +605,12 @@ async def anthropic_proxy_route(
target=str(updated_url),
custom_headers={"x-api-key": "{}".format(anthropic_api_key)},
_forward_headers=True,
+ is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- stream=is_streaming_request, # type: ignore
)
return received_value
@@ -898,12 +1028,12 @@ async def bedrock_proxy_route(
endpoint=endpoint,
target=str(prepped.url),
custom_headers=prepped.headers, # type: ignore
+ is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- stream=is_streaming_request, # type: ignore
custom_body=data, # type: ignore
query_params={}, # type: ignore
)
@@ -981,12 +1111,12 @@ async def assemblyai_proxy_route(
endpoint=endpoint,
target=str(updated_url),
custom_headers={"Authorization": "{}".format(assemblyai_api_key)},
+ is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
- stream=is_streaming_request, # type: ignore
)
return received_value
@@ -1687,13 +1817,12 @@ class BaseOpenAIPassThroughHandler:
custom_headers=BaseOpenAIPassThroughHandler._assemble_headers(
api_key=api_key, request=request, extra_headers=extra_headers
),
+ is_streaming_request=is_streaming_request, # type: ignore
) # dynamically construct pass-through endpoint based on incoming path
received_value = await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
- stream=is_streaming_request, # type: ignore
- query_params=dict(request.query_params), # type: ignore
)
return received_value
diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py
index 2e43191e2db..02adcac0eca 100644
--- a/litellm/proxy/spend_tracking/spend_tracking_utils.py
+++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py
@@ -51,7 +51,7 @@ def _get_spend_logs_metadata(
vector_store_request_metadata: Optional[
List[StandardLoggingVectorStoreRequest]
] = None,
- guardrail_information: Optional[StandardLoggingGuardrailInformation] = None,
+ guardrail_information: Optional[list[StandardLoggingGuardrailInformation]] = None,
usage_object: Optional[dict] = None,
model_map_information: Optional[StandardLoggingModelInformation] = None,
cold_storage_object_key: Optional[str] = None,
@@ -95,9 +95,9 @@ def _get_spend_logs_metadata(
clean_metadata["applied_guardrails"] = applied_guardrails
clean_metadata["batch_models"] = batch_models
clean_metadata["mcp_tool_call_metadata"] = mcp_tool_call_metadata
- clean_metadata["vector_store_request_metadata"] = (
- _get_vector_store_request_for_spend_logs_payload(vector_store_request_metadata)
- )
+ clean_metadata[
+ "vector_store_request_metadata"
+ ] = _get_vector_store_request_for_spend_logs_payload(vector_store_request_metadata)
clean_metadata["guardrail_information"] = guardrail_information
clean_metadata["usage_object"] = usage_object
clean_metadata["model_map_information"] = model_map_information
diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py
index 1d88ad0edf4..81c205095c1 100644
--- a/litellm/proxy/vector_store_endpoints/utils.py
+++ b/litellm/proxy/vector_store_endpoints/utils.py
@@ -2,7 +2,7 @@ from typing import Any, Dict, Literal, Optional
from fastapi import HTTPException, Request
-from litellm.proxy._types import UserAPIKeyAuth
+from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
@@ -73,6 +73,11 @@ def is_allowed_to_call_vector_store_endpoint(
1. Creating a vector store index
2. Reading a vector store index (Search / List / Get)
"""
+ if (
+ user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
+ or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
+ ):
+ return True
# check what allowed permissions are for the key
key_metadata = user_api_key_dict.metadata
team_metadata = user_api_key_dict.team_metadata
diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py
index 00410fe7732..7dab61b151f 100644
--- a/litellm/types/llms/openai.py
+++ b/litellm/types/llms/openai.py
@@ -674,7 +674,9 @@ class OpenAIChatCompletionAssistantMessage(TypedDict, total=False):
class ChatCompletionAssistantMessage(OpenAIChatCompletionAssistantMessage, total=False):
cache_control: ChatCompletionCachedContent
- thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]]
+ thinking_blocks: Optional[
+ List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]
+ ]
class ChatCompletionToolMessage(TypedDict):
@@ -1410,6 +1412,7 @@ class ImageGenerationPartialImageEvent(BaseLiteLLMOpenAIResponseObject):
class ErrorEventError(BaseLiteLLMOpenAIResponseObject):
"""Nested error object within ErrorEvent"""
+
type: str # e.g., 'invalid_request_error'
code: str # e.g., 'context_length_exceeded'
message: str
@@ -1853,10 +1856,10 @@ class OpenAIMcpServerTool(TypedDict, total=False):
class CreateVideoRequest(TypedDict, total=False):
"""
CreateVideoRequest for OpenAI video generation API
-
+
Required Params:
prompt: str - Text prompt that describes the video to generate
-
+
Optional Params:
input_reference: Optional[str] - Optional image reference that guides generation
model: Optional[str] - The video generation model to use (defaults to sora-2)
@@ -1867,6 +1870,7 @@ class CreateVideoRequest(TypedDict, total=False):
extra_body: Optional[Dict[str, str]] - Additional body parameters
timeout: Optional[float] - Request timeout
"""
+
prompt: Required[str]
input_reference: Optional[str]
model: Optional[str]
@@ -1880,42 +1884,43 @@ class CreateVideoRequest(TypedDict, total=False):
class OpenAIVideoObject(BaseModel):
"""OpenAI Video Object representing a video generation job."""
+
id: str
"""Unique identifier for the video job."""
-
+
object: Literal["video"]
"""The object type, which is always 'video'."""
-
+
status: str
"""Current lifecycle status of the video job."""
-
+
created_at: int
"""Unix timestamp (seconds) for when the job was created."""
-
+
completed_at: Optional[int] = None
"""Unix timestamp (seconds) for when the job completed, if finished."""
-
+
expires_at: Optional[int] = None
"""Unix timestamp (seconds) for when the downloadable assets expire, if set."""
-
+
error: Optional[Dict[str, Any]] = None
"""Error payload that explains why generation failed, if applicable."""
-
+
progress: Optional[int] = None
"""Approximate completion percentage for the generation task."""
-
+
remixed_from_video_id: Optional[str] = None
"""Identifier of the source video if this video is a remix."""
-
+
seconds: Optional[str] = None
"""Duration of the generated clip in seconds."""
-
+
size: Optional[str] = None
"""The resolution of the generated video."""
-
+
model: Optional[str] = None
"""The video generation model that produced the job."""
-
+
_hidden_params: Dict[str, Any] = {}
def __contains__(self, key):
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index f91e7d68a8d..e62dfbb20f0 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -9,7 +9,6 @@ from typing import (
Literal,
Mapping,
Optional,
- Tuple,
Union,
)
@@ -1293,7 +1292,7 @@ class ModelResponse(ModelResponseBase):
choices: List[Union[Choices, StreamingChoices]]
"""The list of completion choices the model generated for the input prompt."""
- def __init__(
+ def __init__( # noqa: PLR0915
self,
id=None,
choices=None,
@@ -2201,7 +2200,7 @@ class StandardLoggingPayload(TypedDict):
error_information: Optional[StandardLoggingPayloadErrorInformation]
model_parameters: dict
hidden_params: StandardLoggingHiddenParams
- guardrail_information: Optional[StandardLoggingGuardrailInformation]
+ guardrail_information: Optional[list[StandardLoggingGuardrailInformation]]
standard_built_in_tools_params: Optional[StandardBuiltInToolsParams]
diff --git a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py
index 139df0d6f09..a0f24376614 100644
--- a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py
+++ b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py
@@ -151,5 +151,44 @@ async def test__transform_request_body_image_config():
rb: RequestBody = transformation._transform_request_body(**transform_request_params)
- assert "imageConfig" in rb
- assert rb["imageConfig"] == {"aspectRatio": "16:9"}
\ No newline at end of file
+ assert "generationConfig" in rb
+ assert "imageConfig" in rb["generationConfig"]
+ assert rb["generationConfig"]["imageConfig"] == {"aspectRatio": "16:9"}
+
+
+@pytest.mark.asyncio
+async def test__transform_request_body_image_config_snake_case():
+ """
+ Test that Vertex AI Gemini supports the image_config parameter (snake_case) for gemini-2.5-flash-image model.
+ This should be transformed to imageConfig with aspectRatio.
+ """
+ model = "gemini-2.5-flash-image"
+ messages = [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": "Create a picture of a nano banana dish in a fancy restaurant with a Gemini theme"
+ }
+ ]
+ }
+ ]
+ optional_params = {
+ "image_config": {"aspect_ratio": "16:9"}
+ }
+ litellm_params = {}
+ transform_request_params = {
+ "messages": messages,
+ "model": model,
+ "optional_params": optional_params,
+ "custom_llm_provider": "gemini",
+ "litellm_params": litellm_params,
+ "cached_content": None,
+ }
+
+ rb: RequestBody = transformation._transform_request_body(**transform_request_params)
+
+ assert "generationConfig" in rb
+ assert "image_config" in rb["generationConfig"]
+ assert rb["generationConfig"]["image_config"] == {"aspect_ratio": "16:9"}
\ No newline at end of file
diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py
index 21d2f82c984..24a0321d0c7 100644
--- a/tests/llm_responses_api_testing/test_openai_responses_api.py
+++ b/tests/llm_responses_api_testing/test_openai_responses_api.py
@@ -24,11 +24,13 @@ from litellm.types.llms.openai import (
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from base_responses_api import BaseResponsesAPITest, validate_responses_api_response
+
class TestOpenAIResponsesAPITest(BaseResponsesAPITest):
def get_base_completion_call_args(self):
return {
"model": "openai/gpt-4o",
}
+
def get_base_completion_reasoning_call_args(self):
return {
"model": "openai/gpt-5-mini",
@@ -195,12 +197,12 @@ async def test_openai_responses_api_returns_headers(sync_mode):
"""
Test that OpenAI responses API returns OpenAI headers in _hidden_params.
This ensures the proxy can forward these headers to clients.
-
+
Related issue: LiteLLM responses API should return OpenAI headers like chat completions does
"""
litellm._turn_on_debug()
litellm.set_verbose = True
-
+
if sync_mode:
response = litellm.responses(
model="gpt-4o",
@@ -213,23 +215,28 @@ async def test_openai_responses_api_returns_headers(sync_mode):
input="Say hello",
max_output_tokens=20,
)
-
+
# Verify response is valid
assert response is not None
assert isinstance(response, ResponsesAPIResponse)
-
+
# Verify _hidden_params exists
- assert hasattr(response, "_hidden_params"), "Response should have _hidden_params attribute"
+ assert hasattr(
+ response, "_hidden_params"
+ ), "Response should have _hidden_params attribute"
assert response._hidden_params is not None, "_hidden_params should not be None"
-
+
# Verify additional_headers exists in _hidden_params
- assert "additional_headers" in response._hidden_params, \
- "_hidden_params should contain 'additional_headers' key"
-
+ assert (
+ "additional_headers" in response._hidden_params
+ ), "_hidden_params should contain 'additional_headers' key"
+
additional_headers = response._hidden_params["additional_headers"]
- assert isinstance(additional_headers, dict), "additional_headers should be a dictionary"
+ assert isinstance(
+ additional_headers, dict
+ ), "additional_headers should be a dictionary"
assert len(additional_headers) > 0, "additional_headers should not be empty"
-
+
# Check for expected OpenAI rate limit headers
# These can be either direct (x-ratelimit-*) or prefixed (llm_provider-x-ratelimit-*)
rate_limit_headers = [
@@ -238,22 +245,26 @@ async def test_openai_responses_api_returns_headers(sync_mode):
"x-ratelimit-remaining-requests",
"x-ratelimit-limit-requests",
]
-
+
found_headers = []
for header_name in rate_limit_headers:
if header_name in additional_headers:
found_headers.append(header_name)
elif f"llm_provider-{header_name}" in additional_headers:
found_headers.append(f"llm_provider-{header_name}")
-
- assert len(found_headers) > 0, \
- f"Should find at least one OpenAI rate limit header. Headers found: {list(additional_headers.keys())}"
-
+
+ assert (
+ len(found_headers) > 0
+ ), f"Should find at least one OpenAI rate limit header. Headers found: {list(additional_headers.keys())}"
+
# Verify headers key also exists (raw headers)
- assert "headers" in response._hidden_params, \
- "_hidden_params should contain 'headers' key with raw response headers"
-
- print(f"✓ Successfully validated OpenAI headers in {'sync' if sync_mode else 'async'} mode")
+ assert (
+ "headers" in response._hidden_params
+ ), "_hidden_params should contain 'headers' key with raw response headers"
+
+ print(
+ f"✓ Successfully validated OpenAI headers in {'sync' if sync_mode else 'async'} mode"
+ )
print(f" Found {len(additional_headers)} headers total")
print(f" Rate limit headers found: {found_headers}")
@@ -659,8 +670,6 @@ async def test_openai_responses_litellm_router_no_metadata():
request_body = mock_post.call_args.kwargs["json"]
print("Request body:", json.dumps(request_body, indent=4))
-
-
# Assert metadata is not in the request
assert (
"metadata" not in request_body
@@ -1132,6 +1141,7 @@ def test_basic_computer_use_preview_tool_call():
"user": None,
"metadata": {},
}
+
class MockResponse:
def __init__(self, json_data, status_code):
self._json_data = json_data
@@ -1151,21 +1161,23 @@ def test_basic_computer_use_preview_tool_call():
# Call the responses API with computer_use_preview tool
response = litellm.responses(
model="openai/computer-use-preview",
- tools=[{
- "type": "computer_use_preview",
- "display_width": 1024,
- "display_height": 768,
- "environment": "linux" # other possible values: "mac", "windows", "ubuntu"
- }],
+ tools=[
+ {
+ "type": "computer_use_preview",
+ "display_width": 1024,
+ "display_height": 768,
+ "environment": "linux", # other possible values: "mac", "windows", "ubuntu"
+ }
+ ],
input="Check the latest OpenAI news on bing.com.",
reasoning={"summary": "concise"},
- truncation="auto"
+ truncation="auto",
)
# Verify the request was made correctly
mock_post.assert_called_once()
request_body = mock_post.call_args.kwargs["json"]
-
+
# Validate the request structure
assert request_body["model"] == "computer-use-preview"
assert len(request_body["tools"]) == 1
@@ -1173,15 +1185,14 @@ def test_basic_computer_use_preview_tool_call():
assert request_body["tools"][0]["display_width"] == 1024
assert request_body["tools"][0]["display_height"] == 768
assert request_body["tools"][0]["environment"] == "linux"
-
+
# Check that reasoning was passed correctly
assert request_body["reasoning"]["summary"] == "concise"
assert request_body["truncation"] == "auto"
-
+
# Validate the input format
assert isinstance(request_body["input"], str)
assert request_body["input"] == "Check the latest OpenAI news on bing.com."
-
def test_mcp_tools_with_responses_api():
@@ -1193,19 +1204,15 @@ def test_mcp_tools_with_responses_api():
"server_url": "https://mcp.zapier.com/api/mcp/mcp",
"headers": {
"Authorization": f"Bearer {os.getenv('ZAPIER_CI_CD_MCP_TOKEN')}"
- }
+ },
}
]
MODEL = "openai/gpt-4.1"
USER_QUERY = "how does tiktoken work?"
#########################################################
- # Step 1: OpenAI will use MCP LIST, and return a list of MCP calls for our approval
+ # Step 1: OpenAI will use MCP LIST, and return a list of MCP calls for our approval
try:
- response = litellm.responses(
- model=MODEL,
- tools=MCP_TOOLS,
- input=USER_QUERY
- )
+ response = litellm.responses(model=MODEL, tools=MCP_TOOLS, input=USER_QUERY)
print(response)
response = cast(ResponsesAPIResponse, response)
@@ -1225,20 +1232,26 @@ def test_mcp_tools_with_responses_api():
{
"type": "mcp_approval_response",
"approve": True,
- "approval_request_id": mcp_approval_id
+ "approval_request_id": mcp_approval_id,
}
],
previous_response_id=response.id,
)
print(response_with_mcp_call)
except litellm.APIError as e:
- if "424" in str(e) or "Failed Dependency" in str(e) or "external_connector_error" in str(e):
+ if (
+ "424" in str(e)
+ or "Failed Dependency" in str(e)
+ or "external_connector_error" in str(e)
+ ):
pytest.skip(f"Skipping test due to external MCP server error: {e}")
else:
raise e
except litellm.InternalServerError as e:
if "500" in str(e) or "server_error" in str(e):
- pytest.skip(f"Skipping test due to OpenAI server error (likely MCP server unavailable): {e}")
+ pytest.skip(
+ f"Skipping test due to OpenAI server error (likely MCP server unavailable): {e}"
+ )
else:
raise e
@@ -1248,29 +1261,28 @@ async def test_openai_responses_api_field_types():
"""Test that specific fields in the response have the correct types"""
litellm._turn_on_debug()
litellm.set_verbose = True
-
+
# Test with store=True
response = await litellm.aresponses(
model="gpt-4o",
input="hi",
)
-
+
# Verify created_at is an integer
assert isinstance(response.created_at, int), "created_at should be an integer"
-
+
# Verify store field is present and matches input
assert hasattr(response, "store"), "store field should be present"
assert response.store is True, "store field should match input value"
-
+
# Test without store parameter
- response_without_store = await litellm.aresponses(
- model="gpt-4o",
- input="hi"
- )
-
+ response_without_store = await litellm.aresponses(model="gpt-4o", input="hi")
+
# Verify created_at is still an integer
- assert isinstance(response_without_store.created_at, int), "created_at should be an integer"
-
+ assert isinstance(
+ response_without_store.created_at, int
+ ), "created_at should be an integer"
+
# Verify store field is present but None when not specified
assert hasattr(response_without_store, "store"), "store field should be present"
@@ -1279,7 +1291,7 @@ async def test_openai_responses_api_field_types():
async def test_store_field_transformation():
"""Test store field transformation with mocked API responses"""
config = OpenAIResponsesAPIConfig()
-
+
# Initialize logging object with required parameters
logging_obj = LiteLLMLoggingObj(
model="gpt-4o",
@@ -1288,7 +1300,7 @@ async def test_store_field_transformation():
call_type="aresponses",
start_time=time.time(),
litellm_call_id="test-call-id",
- function_id="test-function-id"
+ function_id="test-function-id",
)
# Base response data with all required fields
@@ -1297,7 +1309,17 @@ async def test_store_field_transformation():
"created_at": 1751443898,
"model": "gpt-4o",
"object": "response",
- "output": [{"type": "message", "id": "msg_1", "status": "completed", "role": "assistant", "content": [{"type": "output_text", "text": "Hello", "annotations": []}]}],
+ "output": [
+ {
+ "type": "message",
+ "id": "msg_1",
+ "status": "completed",
+ "role": "assistant",
+ "content": [
+ {"type": "output_text", "text": "Hello", "annotations": []}
+ ],
+ }
+ ],
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
@@ -1314,70 +1336,70 @@ async def test_store_field_transformation():
"text": None,
"truncation": "auto",
"usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30},
- "user": "test_user"
+ "user": "test_user",
}
# Test case 1: API returns store=True
mock_response_store_true = httpx.Response(
- status_code=200,
- content=json.dumps({**base_response, "store": True}).encode()
+ status_code=200, content=json.dumps({**base_response, "store": True}).encode()
)
# Test case 2: API returns store=False
mock_response_store_false = httpx.Response(
- status_code=200,
- content=json.dumps({**base_response, "store": False}).encode()
+ status_code=200, content=json.dumps({**base_response, "store": False}).encode()
)
# Test case 3: API returns store=null
mock_response_store_null = httpx.Response(
- status_code=200,
- content=json.dumps({**base_response, "store": None}).encode()
+ status_code=200, content=json.dumps({**base_response, "store": None}).encode()
)
# Test case 4: API omits store field
mock_response_no_store = httpx.Response(
- status_code=200,
- content=json.dumps(base_response).encode()
+ status_code=200, content=json.dumps(base_response).encode()
)
# Test when store=True in request
logging_obj.optional_params = {"store": True}
response = config.transform_response_api_response(
- model="gpt-4o",
- raw_response=mock_response_store_true,
- logging_obj=logging_obj
+ model="gpt-4o", raw_response=mock_response_store_true, logging_obj=logging_obj
)
- assert response.store is True, "store should be True when specified in request and API returns True"
+ assert (
+ response.store is True
+ ), "store should be True when specified in request and API returns True"
# Test when store=False in request
logging_obj.optional_params = {"store": False}
response = config.transform_response_api_response(
- model="gpt-4o",
- raw_response=mock_response_store_false,
- logging_obj=logging_obj
+ model="gpt-4o", raw_response=mock_response_store_false, logging_obj=logging_obj
)
- assert response.store is False, "store should be False when specified in request and API returns False"
+ assert (
+ response.store is False
+ ), "store should be False when specified in request and API returns False"
# Test when store not in request but API returns null
response = config.transform_response_api_response(
- model="gpt-4o",
- raw_response=mock_response_store_null,
- logging_obj=logging_obj
+ model="gpt-4o", raw_response=mock_response_store_null, logging_obj=logging_obj
)
- assert response.store is None, "store should be None when not specified in request and API returns null"
+ assert (
+ response.store is None
+ ), "store should be None when not specified in request and API returns null"
# Test when store not in request and API omits store field
response = config.transform_response_api_response(
- model="gpt-4o",
- raw_response=mock_response_no_store,
- logging_obj=logging_obj
+ model="gpt-4o", raw_response=mock_response_no_store, logging_obj=logging_obj
)
- assert response.store is None, "store should be None when not specified in request and API omits store"
+ assert (
+ response.store is None
+ ), "store should be None when not specified in request and API omits store"
# Verify created_at is always converted to integer
- assert isinstance(response.created_at, int), "created_at should always be converted to integer"
- assert response.created_at == 1751443898, "created_at should maintain the same value after conversion"
+ assert isinstance(
+ response.created_at, int
+ ), "created_at should always be converted to integer"
+ assert (
+ response.created_at == 1751443898
+ ), "created_at should maintain the same value after conversion"
@pytest.mark.asyncio
@@ -1455,10 +1477,14 @@ async def test_aresponses_service_tier_and_safety_identifier():
mock_post.assert_called_once()
request_body = mock_post.call_args.kwargs["json"]
print("request_body=", json.dumps(request_body, indent=4, default=str))
-
+
# Validate that both parameters are present in the request body
- assert request_body["service_tier"] == "flex", "service_tier should be 'flex' in request body"
- assert request_body["safety_identifier"] == "123", "safety_identifier should be '123' in request body"
+ assert (
+ request_body["service_tier"] == "flex"
+ ), "service_tier should be 'flex' in request body"
+ assert (
+ request_body["safety_identifier"] == "123"
+ ), "safety_identifier should be '123' in request body"
assert request_body["model"] == "gpt-4o"
assert request_body["input"] == "Test with service tier and safety identifier"
@@ -1469,11 +1495,11 @@ async def test_aresponses_service_tier_and_safety_identifier():
@pytest.mark.asyncio
async def test_openai_gpt5_reasoning_effort_parameter():
"""Test that reasoning_effort parameter is properly sent in the HTTP request for GPT-5 models."""
-
+
# Mock response for GPT-5 responses API (correct format)
mock_response = {
"id": "resp_01ABC123",
- "object": "response",
+ "object": "response",
"created_at": 1729621667,
"status": "completed",
"model": "gpt-5-mini",
@@ -1481,10 +1507,14 @@ async def test_openai_gpt5_reasoning_effort_parameter():
{
"type": "message",
"id": "msg_123",
- "status": "completed",
+ "status": "completed",
"role": "assistant",
"content": [
- {"type": "output_text", "text": "The capital of France is Paris.", "annotations": []}
+ {
+ "type": "output_text",
+ "text": "The capital of France is Paris.",
+ "annotations": [],
+ }
],
}
],
@@ -1544,8 +1574,12 @@ async def test_openai_gpt5_reasoning_effort_parameter():
print("request_body=", json.dumps(request_body, indent=4, default=str))
print("reasoning=", request_body["reasoning"])
# Validate that reasoning_effort is present in the request body
- assert "reasoning" in request_body, "reasoning should be present in request body"
- assert request_body["reasoning"]["effort"] == "minimal", "reasoning_effort should be 'minimal' in request body"
+ assert (
+ "reasoning" in request_body
+ ), "reasoning should be present in request body"
+ assert (
+ request_body["reasoning"]["effort"] == "minimal"
+ ), "reasoning_effort should be 'minimal' in request body"
assert request_body["model"] == "gpt-5-mini"
assert request_body["input"] == "What is the capital of France?"
@@ -1553,9 +1587,6 @@ async def test_openai_gpt5_reasoning_effort_parameter():
print("Response:", json.dumps(response, indent=4, default=str))
-
-
-
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", [True, False])
async def test_basic_openai_responses_with_websearch(stream):
@@ -1565,12 +1596,7 @@ async def test_basic_openai_responses_with_websearch(stream):
model=request_model,
stream=stream,
input="hi",
- tools=[
- {
- "type": "web_search",
- "search_context_size": "low"
- }
- ]
+ tools=[{"type": "web_search", "search_context_size": "low"}],
)
if stream:
async for chunk in response:
@@ -1596,10 +1622,49 @@ async def test_openai_responses_api_token_limit_error():
# This will raise ValidationError instead of showing the real error
response = await litellm.aresponses(
- model="gpt-5-mini",
- input=oversized_text,
- stream=True
+ model="gpt-5-mini", input=oversized_text, stream=True
)
async for event in response:
- print(event) # Never reaches here - ValidationError is raised
\ No newline at end of file
+ print(event) # Never reaches here - ValidationError is raised
+
+
+async def test_openai_streaming_logging():
+ """Test that hard_limit parameter is properly sent in the HTTP request for GPT-5 models."""
+ litellm._turn_on_debug()
+ from litellm.integrations.custom_logger import CustomLogger
+ from litellm.types.utils import Usage
+
+ class TestCustomLogger(CustomLogger):
+ validate_usage = False
+
+ def __init__(self):
+ self.standard_logging_object: Optional[StandardLoggingPayload] = None
+
+ async def async_log_success_event(
+ self, kwargs, response_obj, start_time, end_time
+ ):
+ print(f"response_obj: {response_obj.usage}")
+ assert isinstance(
+ response_obj.usage, Usage
+ ), f"Expected response_obj.usage to be of type Usage, but got {type(response_obj.usage)}"
+ print("\n\nVALIDATED USAGE\n\n")
+ self.validate_usage = True
+
+ tcl = TestCustomLogger()
+ litellm.callbacks = [tcl]
+ request_model = "gpt-5-mini"
+ response = await litellm.aresponses(
+ model=request_model,
+ input="What is the capital of France?",
+ stream=True,
+ )
+ print("response=", json.dumps(response, indent=4, default=str))
+
+ async for event in response:
+ if event.type == "response.completed":
+ final_response = event
+ print("litellm response=", json.dumps(event, indent=4, default=str))
+
+ await asyncio.sleep(2)
+ assert tcl.validate_usage, "Usage should be validated"
diff --git a/tests/llm_translation/test_openai.py b/tests/llm_translation/test_openai.py
index 55d2d3953c6..d1fa9df916f 100644
--- a/tests/llm_translation/test_openai.py
+++ b/tests/llm_translation/test_openai.py
@@ -1308,3 +1308,107 @@ def test_openai_gpt_5_codex_reasoning():
print("response: ", response)
for chunk in response:
print("chunk: ", chunk)
+
+
+# Tests moved from test_streaming_n_with_tools.py
+# Regression test for: https://github.com/BerriAI/litellm/issues/8977
+@pytest.mark.parametrize("model", ["gpt-4o", "gpt-4-turbo"])
+@pytest.mark.asyncio
+async def test_streaming_tool_calls_with_n_greater_than_1(model):
+ """
+ Test that the index field in a choice object is correctly populated
+ when using streaming mode with n>1 and tool calls.
+
+ Regression test for: https://github.com/BerriAI/litellm/issues/8977
+ """
+ tools = [
+ {
+ "type": "function",
+ "function": {
+ "strict": True,
+ "name": "get_current_weather",
+ "description": "Get the current weather in a given location",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "location": {
+ "type": "string",
+ "description": "The city and state, e.g. San Francisco, CA",
+ },
+ "unit": {
+ "type": "string",
+ "enum": ["celsius", "fahrenheit"],
+ },
+ },
+ "required": ["location", "unit"],
+ "additionalProperties": False,
+ },
+ },
+ }
+ ]
+
+ response = litellm.completion(
+ model=model,
+ messages=[
+ {
+ "role": "user",
+ "content": "What is the weather in San Francisco?",
+ },
+ ],
+ tools=tools,
+ stream=True,
+ n=3,
+ )
+
+ # Collect all chunks and their indices
+ indices_seen = []
+ for chunk in response:
+ assert len(chunk.choices) == 1, "Each streaming chunk should have exactly 1 choice"
+ assert hasattr(chunk.choices[0], "index"), "Choice should have an index attribute"
+ index = chunk.choices[0].index
+ indices_seen.append(index)
+
+ # Verify that we got chunks with different indices (0, 1, 2 for n=3)
+ unique_indices = set(indices_seen)
+ assert unique_indices == {0, 1, 2}, f"Should have indices 0, 1, 2 for n=3, got {unique_indices}"
+
+ print(f"✓ Test passed: streaming with n=3 and tool calls correctly populates index field")
+ print(f" Indices seen: {indices_seen}")
+ print(f" Unique indices: {unique_indices}")
+
+
+@pytest.mark.parametrize("model", ["gpt-4o"])
+@pytest.mark.asyncio
+async def test_streaming_content_with_n_greater_than_1(model):
+ """
+ Test that the index field is correctly populated for regular content streaming
+ (not tool calls) with n>1.
+ """
+ response = litellm.completion(
+ model=model,
+ messages=[
+ {
+ "role": "user",
+ "content": "Say hello in one word",
+ },
+ ],
+ stream=True,
+ n=2,
+ max_tokens=10,
+ )
+
+ # Collect all chunks and their indices
+ indices_seen = []
+ for chunk in response:
+ assert len(chunk.choices) == 1, "Each streaming chunk should have exactly 1 choice"
+ assert hasattr(chunk.choices[0], "index"), "Choice should have an index attribute"
+ index = chunk.choices[0].index
+ indices_seen.append(index)
+
+ # Verify that we got chunks with different indices (0, 1 for n=2)
+ unique_indices = set(indices_seen)
+ assert unique_indices == {0, 1}, f"Should have indices 0, 1 for n=2, got {unique_indices}"
+
+ print(f"✓ Test passed: streaming with n=2 and regular content correctly populates index field")
+ print(f" Indices seen: {indices_seen}")
+ print(f" Unique indices: {unique_indices}")
diff --git a/tests/llm_translation/test_perplexity_reasoning.py b/tests/llm_translation/test_perplexity_reasoning.py
index 70b665ea339..7db99014e9d 100644
--- a/tests/llm_translation/test_perplexity_reasoning.py
+++ b/tests/llm_translation/test_perplexity_reasoning.py
@@ -62,13 +62,12 @@ class TestPerplexityReasoning:
"""
Test that reasoning_effort is correctly passed in actual completion call (mocked)
"""
- from openai import OpenAI
- from openai.types.chat.chat_completion import ChatCompletion
+ import httpx
litellm.set_verbose = True
# Mock successful response with reasoning content
- response_object = {
+ response_json = {
"id": "cmpl-test",
"object": "chat.completion",
"created": 1677652288,
@@ -94,35 +93,37 @@ class TestPerplexityReasoning:
},
}
- pydantic_obj = ChatCompletion(**response_object)
+ def mock_post(*args, **kwargs):
+ # Create a mock response
+ mock_response = MagicMock(spec=httpx.Response)
+ mock_response.status_code = 200
+ mock_response.headers = {"content-type": "application/json"}
+ mock_response.json.return_value = response_json
+ mock_response.text = json.dumps(response_json)
+
+ # Store the request data for verification
+ mock_post.last_request_data = kwargs.get("data")
+ if isinstance(mock_post.last_request_data, (str, bytes)):
+ mock_post.last_request_data = json.loads(mock_post.last_request_data)
+
+ return mock_response
- def _return_pydantic_obj(*args, **kwargs):
- new_response = MagicMock()
- new_response.headers = {"content-type": "application/json"}
- new_response.parse.return_value = pydantic_obj
- return new_response
-
- openai_client = OpenAI(api_key="fake-api-key")
-
- with patch.object(
- openai_client.chat.completions.with_raw_response, "create", side_effect=_return_pydantic_obj
- ) as mock_client:
+ # Mock at the HTTP handler level
+ with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post", side_effect=mock_post) as mock_http:
response = completion(
model=model,
messages=[{"role": "user", "content": "Hello, please think about this carefully."}],
reasoning_effort="high",
- client=openai_client,
+ api_key="fake-api-key",
)
# Verify the call was made
- assert mock_client.called
-
- # Get the request data from the mock call
- call_args = mock_client.call_args
- request_data = call_args.kwargs
+ assert mock_http.called
# Verify reasoning_effort was included in the request
+ request_data = mock_post.last_request_data
+ assert request_data is not None
assert "reasoning_effort" in request_data
assert request_data["reasoning_effort"] == "high"
diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py
index 6f333c4eddf..d8ebb059b9a 100644
--- a/tests/local_testing/test_completion.py
+++ b/tests/local_testing/test_completion.py
@@ -1244,6 +1244,9 @@ def test_completion_fireworks_ai_dynamic_params(api_key, api_base):
# @pytest.mark.skip(reason="this test is flaky")
def test_completion_perplexity_api():
try:
+ import httpx
+ import json
+
response_object = {
"id": "a8f37485-026e-45da-81a9-cf0184896840",
"model": "llama-3-sonar-small-32k-online",
@@ -1270,25 +1273,17 @@ def test_completion_perplexity_api():
],
}
- from openai import OpenAI
- from openai.types.chat.chat_completion import ChatCompletion
+ def mock_post(*args, **kwargs):
+ # Create a mock response
+ mock_response = MagicMock(spec=httpx.Response)
+ mock_response.status_code = 200
+ mock_response.headers = {"content-type": "application/json"}
+ mock_response.json.return_value = response_object
+ mock_response.text = json.dumps(response_object)
+ return mock_response
- pydantic_obj = ChatCompletion(**response_object)
-
- def _return_pydantic_obj(*args, **kwargs):
- new_response = MagicMock()
- new_response.headers = {"hello": "world"}
-
- new_response.parse.return_value = pydantic_obj
- return new_response
-
- openai_client = OpenAI()
-
- with patch.object(
- openai_client.chat.completions.with_raw_response,
- "create",
- side_effect=_return_pydantic_obj,
- ) as mock_client:
+ # Mock at the HTTP handler level
+ with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post", side_effect=mock_post):
# litellm.set_verbose= True
messages = [
{"role": "system", "content": "You're a good bot"},
@@ -1302,10 +1297,9 @@ def test_completion_perplexity_api():
},
]
response = completion(
- model="mistral-7b-instruct",
+ model="perplexity/llama-3-sonar-small-32k-online",
messages=messages,
- api_base="https://api.perplexity.ai",
- client=openai_client,
+ api_key="fake-api-key",
)
print(response)
assert hasattr(response, "citations")
@@ -2470,8 +2464,6 @@ def test_completion_azure_key_completion_arg():
pytest.fail(f"Error occurred: {e}")
-
-
async def test_re_use_azure_async_client():
try:
print("azure gpt-3.5 ASYNC with clie nttest\n\n")
@@ -4272,7 +4264,6 @@ def test_deepseek_reasoning_content_completion():
pytest.skip("Model is timing out")
-
def test_qwen_text_completion():
# litellm._turn_on_debug()
resp = litellm.completion(
@@ -4407,3 +4398,37 @@ def test_completion_gpt_4o_empty_str():
messages=[{"role": "user", "content": ""}],
)
assert resp.choices[0].message.content is not None
+
+
+def test_edit_note():
+ litellm.callbacks = ["langfuse_otel"]
+ response = completion(
+ model="gpt-4o",
+ messages=[
+ {
+ "role": "system",
+ "content": "Your only job is to call the edit_note tool with the content specified in the user's message.",
+ },
+ {
+ "role": "user",
+ "content": "Edit the note with the content: 'This is a test note.'",
+ },
+ ],
+ tools=[
+ {
+ "type": "function",
+ "function": {
+ "name": "edit_note",
+ "description": "Edit the note with the content specified in the user's message.",
+ "parameters": {
+ "type": "object",
+ "properties": {
+ "content": {"type": "string"},
+ },
+ },
+ },
+ },
+ ],
+ )
+
+ return response
diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py
index 534569f15af..41132a65f2c 100644
--- a/tests/proxy_unit_tests/test_proxy_token_counter.py
+++ b/tests/proxy_unit_tests/test_proxy_token_counter.py
@@ -743,3 +743,79 @@ async def test_bedrock_count_tokens_endpoint():
await mock_count_tokens_handler(
request_data, {}, "anthropic.claude-3-sonnet-20240229-v1:0"
)
+
+
+@pytest.mark.asyncio
+async def test_vertex_ai_anthropic_token_counting():
+ """
+ Unit test for Vertex AI Anthropic token counting with mocked API calls.
+
+ This tests the token counting implementation for Vertex AI partner models
+ without making actual API calls. Mocks at the handler level to test the full flow.
+ """
+ from unittest.mock import AsyncMock, patch, MagicMock
+
+ # Mock the Vertex AI partner models token counter response
+ mock_token_response = {
+ "input_tokens": 15,
+ "tokenizer_used": "vertex_ai_partner_models",
+ }
+
+ llm_router = Router(
+ model_list=[
+ {
+ "model_name": "vertex_ai/claude-3-5-sonnet-20241022",
+ "litellm_params": {
+ "model": "vertex_ai/claude-3-5-sonnet-20241022",
+ "vertex_project": "test-project",
+ "vertex_location": "us-east5",
+ },
+ }
+ ]
+ )
+
+ setattr(litellm.proxy.proxy_server, "llm_router", llm_router)
+
+ # Mock the lower level handler method
+ with patch(
+ "litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler.VertexAIPartnerModelsTokenCounter.handle_count_tokens_request"
+ ) as mock_handle_count_tokens:
+ mock_handle_count_tokens.return_value = mock_token_response
+
+ # Test with messages format and call_endpoint=True
+ response = await token_counter(
+ request=TokenCountRequest(
+ model="vertex_ai/claude-3-5-sonnet-20241022",
+ messages=[
+ {
+ "role": "user",
+ "content": "Hello Claude on Vertex AI! How are you?",
+ }
+ ],
+ ),
+ call_endpoint=True,
+ )
+
+ # Validate that handle_count_tokens_request was called
+ assert mock_handle_count_tokens.called
+
+ # Verify the call arguments
+ call_args = mock_handle_count_tokens.call_args
+ assert call_args is not None
+ assert call_args.kwargs["model"] == "claude-3-5-sonnet-20241022"
+ assert "messages" in call_args.kwargs["request_data"]
+ assert (
+ call_args.kwargs["request_data"]["messages"][0]["content"]
+ == "Hello Claude on Vertex AI! How are you?"
+ )
+
+ # Validate response structure
+ assert response.model_used == "claude-3-5-sonnet-20241022"
+ assert response.request_model == "vertex_ai/claude-3-5-sonnet-20241022"
+ assert response.total_tokens == 15
+ assert response.tokenizer_type == "vertex_ai_partner_models"
+
+ # Validate original response contains input_tokens
+ assert response.original_response is not None
+ assert "input_tokens" in response.original_response
+ assert response.original_response["input_tokens"] == 15
diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py
index 5f67ce08c1a..6eb6e67a088 100644
--- a/tests/test_litellm/integrations/arize/test_arize_utils.py
+++ b/tests/test_litellm/integrations/arize/test_arize_utils.py
@@ -181,6 +181,107 @@ def test_arize_set_attributes():
span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 40)
+def test_arize_set_attributes_responses_api():
+ """
+ Test setting attributes for Responses API with mixed output (reasoning + message).
+ Verifies that multiple output types are correctly handled.
+ """
+ from unittest.mock import MagicMock
+ from litellm.types.llms.openai import ResponsesAPIResponse, ResponseAPIUsage, OutputTokensDetails
+ from openai.types.responses import ResponseReasoningItem, ResponseOutputMessage, ResponseOutputText
+ from openai.types.responses.response_reasoning_item import Summary
+
+ span = MagicMock() # Mocked tracing span to test attribute setting
+
+ # Construct kwargs to simulate a real LLM request scenario
+ kwargs = {
+ "model": "o3-mini",
+ "messages": [{"role": "user", "content": "What is the answer?"}],
+ "standard_logging_object": {
+ "model_parameters": {"user": "test_user", "stream": True},
+ "metadata": {"key_1": "value_1", "key_2": None},
+ "call_type": "responses",
+ },
+ "optional_params": {
+ "max_tokens": "100",
+ "temperature": "1",
+ "top_p": "5",
+ "stream": True,
+ "user": "test_user",
+ },
+ "litellm_params": {"custom_llm_provider": "openai"},
+ }
+
+ # Simulate Responses API response with mixed output
+ response_obj = ResponsesAPIResponse(
+ id="response-123",
+ created_at=1625247600,
+ output=[
+ ResponseReasoningItem(
+ id="reasoning-001",
+ type="reasoning",
+ summary=[
+ Summary(
+ text="First, I need to analyze...",
+ type="summary_text"
+ )
+ ]
+ ),
+ ResponseOutputMessage(
+ id="msg-001",
+ type="message",
+ role="assistant",
+ status="completed",
+ content=[
+ ResponseOutputText(
+ annotations=[],
+ text="The answer is 42",
+ type="output_text",
+ )
+ ]
+ )
+ ],
+ usage=ResponseAPIUsage(
+ input_tokens=120,
+ output_tokens=250,
+ total_tokens=370,
+ output_tokens_details=OutputTokensDetails(
+ reasoning_tokens=180
+ )
+ )
+ )
+
+ ArizeLogger.set_arize_attributes(span, kwargs, response_obj)
+
+ # Verify reasoning summary was set (index 0)
+ span.set_attribute.assert_any_call(
+ f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_REASONING_SUMMARY}",
+ "First, I need to analyze..."
+ )
+
+ # Verify message content was set (index 1)
+ span.set_attribute.assert_any_call(
+ SpanAttributes.OUTPUT_VALUE,
+ "The answer is 42"
+ )
+ span.set_attribute.assert_any_call(
+ f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.1.{MessageAttributes.MESSAGE_CONTENT}",
+ "The answer is 42"
+ )
+ span.set_attribute.assert_any_call(
+ f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.1.{MessageAttributes.MESSAGE_ROLE}",
+ "assistant"
+ )
+
+ # Verify token counts including reasoning tokens
+ span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370)
+ span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250)
+ span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120)
+ span.set_attribute.assert_any_call(
+ SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180
+ )
+
+
class TestArizeLogger(CustomLogger):
"""
Custom logger implementation to capture standard_callback_dynamic_params.
diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py b/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py
index d7db5ef00a0..f6b8af503fb 100644
--- a/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py
+++ b/tests/test_litellm/integrations/datadog/test_datadog_llm_observability.py
@@ -528,7 +528,7 @@ def create_standard_logging_payload_with_latency_metrics() -> StandardLoggingPay
error_information=None,
model_parameters={"stream": True},
hidden_params=hidden_params,
- guardrail_information=guardrail_info,
+ guardrail_information=[ guardrail_info ],
trace_id="test-trace-id-latency",
custom_llm_provider="openai",
)
@@ -607,11 +607,11 @@ def test_latency_metrics_edge_cases(mock_env_vars):
# Test case 3: Missing guardrail duration should not crash
standard_payload = create_standard_logging_payload_with_cache()
- standard_payload["guardrail_information"] = StandardLoggingGuardrailInformation(
+ standard_payload["guardrail_information"] = [StandardLoggingGuardrailInformation(
guardrail_name="test",
guardrail_status="success",
# duration is missing
- )
+ )]
metadata = logger._get_dd_llm_obs_payload_metadata(standard_payload)
assert "guardrail_overhead_time_ms" not in metadata
@@ -644,20 +644,20 @@ def test_guardrail_information_in_metadata(mock_env_vars):
# Verify the guardrail information structure
guardrail_info = metadata["guardrail_information"]
- assert guardrail_info["guardrail_name"] == "test_guardrail"
- assert guardrail_info["guardrail_status"] == "success"
- assert guardrail_info["duration"] == 0.5
+ assert guardrail_info[0]["guardrail_name"] == "test_guardrail"
+ assert guardrail_info[0]["guardrail_status"] == "success"
+ assert guardrail_info[0]["duration"] == 0.5
# Verify input/output fields are present
- assert "guardrail_request" in guardrail_info
- assert "guardrail_response" in guardrail_info
+ assert "guardrail_request" in guardrail_info[0]
+ assert "guardrail_response" in guardrail_info[0]
# Validate the input/output content
- assert guardrail_info["guardrail_request"]["input"] == "test input message"
- assert guardrail_info["guardrail_request"]["user_id"] == "test_user"
- assert guardrail_info["guardrail_response"]["output"] == "filtered output"
- assert guardrail_info["guardrail_response"]["flagged"] is False
- assert guardrail_info["guardrail_response"]["score"] == 0.1
+ assert guardrail_info[0]["guardrail_request"]["input"] == "test input message"
+ assert guardrail_info[0]["guardrail_request"]["user_id"] == "test_user"
+ assert guardrail_info[0]["guardrail_response"]["output"] == "filtered output"
+ assert guardrail_info[0]["guardrail_response"]["flagged"] is False
+ assert guardrail_info[0]["guardrail_response"]["score"] == 0.1
def create_standard_logging_payload_with_tool_calls() -> StandardLoggingPayload:
diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py
index 1106c0b3678..601c18077a5 100644
--- a/tests/test_litellm/integrations/test_custom_guardrail.py
+++ b/tests/test_litellm/integrations/test_custom_guardrail.py
@@ -237,3 +237,74 @@ class TestApplyGuardrailCheck:
assert hasattr(
child_with_override, "apply_guardrail"
), "All instances should have apply_guardrail via inheritance"
+
+
+class TestGuardrailLoggingAggregation:
+ def _make_guardrail(self):
+ from litellm.types.guardrails import GuardrailEventHooks
+
+ return CustomGuardrail(
+ guardrail_name="test_guardrail",
+ event_hook=GuardrailEventHooks.pre_call,
+ )
+
+ def _invoke_add_log(self, request_data: dict) -> None:
+ guardrail = self._make_guardrail()
+ guardrail.add_standard_logging_guardrail_information_to_request_data(
+ guardrail_json_response={"result": "ok"},
+ request_data=request_data,
+ guardrail_status="success",
+ start_time=1.0,
+ end_time=2.0,
+ duration=1.0,
+ masked_entity_count={"EMAIL": 1},
+ guardrail_provider="presidio",
+ )
+
+ def test_appends_to_existing_metadata_list(self):
+ request_data = {
+ "metadata": {
+ "standard_logging_guardrail_information": [
+ {"guardrail_name": "existing_guardrail"}
+ ]
+ }
+ }
+
+ self._invoke_add_log(request_data)
+
+ info = request_data["metadata"]["standard_logging_guardrail_information"]
+ assert isinstance(info, list)
+ assert len(info) == 2
+ assert info[0]["guardrail_name"] == "existing_guardrail"
+ assert info[1]["guardrail_name"] == "test_guardrail"
+
+ def test_converts_existing_metadata_dict_to_list(self):
+ request_data = {
+ "metadata": {
+ "standard_logging_guardrail_information": {"guardrail_name": "legacy"}
+ }
+ }
+
+ self._invoke_add_log(request_data)
+
+ info = request_data["metadata"]["standard_logging_guardrail_information"]
+ assert isinstance(info, list)
+ assert len(info) == 2
+ assert info[0]["guardrail_name"] == "legacy"
+ assert info[1]["guardrail_name"] == "test_guardrail"
+
+ def test_appends_to_litellm_metadata(self):
+ request_data = {
+ "litellm_metadata": {
+ "standard_logging_guardrail_information": [
+ {"guardrail_name": "litellm_existing"}
+ ]
+ }
+ }
+
+ self._invoke_add_log(request_data)
+
+ info = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
+ assert isinstance(info, list)
+ assert len(info) == 2
+ assert info[1]["guardrail_name"] == "test_guardrail"
diff --git a/tests/test_litellm/integrations/test_langfuse_otel.py b/tests/test_litellm/integrations/test_langfuse_otel.py
index a2fbbf74cec..f8c662979ad 100644
--- a/tests/test_litellm/integrations/test_langfuse_otel.py
+++ b/tests/test_litellm/integrations/test_langfuse_otel.py
@@ -1,5 +1,6 @@
import json
import os
+from datetime import datetime
from unittest.mock import MagicMock, patch
import pytest
@@ -7,7 +8,6 @@ import pytest
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
from litellm.types.integrations.langfuse_otel import LangfuseOtelConfig
from litellm.types.llms.openai import ResponsesAPIResponse
-from datetime import datetime
class TestLangfuseOtelIntegration:
@@ -75,6 +75,10 @@ class TestLangfuseOtelIntegration:
def test_set_langfuse_otel_attributes(self):
"""Test that set_langfuse_otel_attributes calls the Arize utils function."""
+ from litellm.integrations.langfuse.langfuse_otel_attributes import (
+ LangfuseLLMObsOTELAttributes,
+ )
+
mock_span = MagicMock()
mock_kwargs = {"test": "kwargs"}
mock_response = {"test": "response"}
@@ -82,7 +86,7 @@ class TestLangfuseOtelIntegration:
with patch('litellm.integrations.arize._utils.set_attributes') as mock_set_attributes:
LangfuseOtelLogger.set_langfuse_otel_attributes(mock_span, mock_kwargs, mock_response)
- mock_set_attributes.assert_called_once_with(mock_span, mock_kwargs, mock_response)
+ mock_set_attributes.assert_called_once_with(mock_span, mock_kwargs, mock_response, LangfuseLLMObsOTELAttributes)
def test_set_langfuse_environment_attribute(self):
"""Test that Langfuse environment is set correctly when environment variable is present."""
@@ -191,8 +195,8 @@ class TestLangfuseOtelIntegration:
def test_set_langfuse_specific_attributes_with_content(self):
"""Test that _set_langfuse_specific_attributes correctly sets observation.output with regular content response."""
- from litellm.types.utils import Choices, ModelResponse
from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes
+ from litellm.types.utils import Choices, ModelResponse
# Create response with content
response_obj = ModelResponse(
@@ -240,8 +244,13 @@ class TestLangfuseOtelIntegration:
def test_set_langfuse_specific_attributes_with_tool_calls(self):
"""Test that _set_langfuse_specific_attributes correctly sets observation.output with tool calls in Langfuse format."""
- from litellm.types.utils import Choices, Function, ChatCompletionMessageToolCall, ModelResponse
from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes
+ from litellm.types.utils import (
+ ChatCompletionMessageToolCall,
+ Choices,
+ Function,
+ ModelResponse,
+ )
# Create response with tool calls
response_obj = ModelResponse(
@@ -388,13 +397,17 @@ class TestLangfuseOtelResponsesAPI:
mock_span = MagicMock()
+ from litellm.integrations.langfuse.langfuse_otel_attributes import (
+ LangfuseLLMObsOTELAttributes,
+ )
+
with patch('litellm.integrations.arize._utils.set_attributes') as mock_set_attributes:
with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute:
logger = LangfuseOtelLogger()
logger.set_langfuse_otel_attributes(mock_span, kwargs, mock_response)
# Verify that set_attributes was called for general attributes
- mock_set_attributes.assert_called_once_with(mock_span, kwargs, mock_response)
+ mock_set_attributes.assert_called_once_with(mock_span, kwargs, mock_response, LangfuseLLMObsOTELAttributes)
# Verify that Langfuse-specific attributes were set
mock_safe_set_attribute.assert_any_call(
@@ -474,6 +487,132 @@ class TestLangfuseOtelResponsesAPI:
for expected_call in expected_calls:
mock_safe_set_attribute.assert_any_call(*expected_call)
+ def test_responses_api_with_output(self):
+ """Test Langfuse OTEL logger with Responses API output (reasoning + message)."""
+ from openai.types.responses import ResponseReasoningItem, ResponseOutputMessage, ResponseOutputText
+ from openai.types.responses.response_reasoning_item import Summary
+ from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes
+
+ # Create Responses API response with reasoning and message
+ response_obj = ResponsesAPIResponse(
+ id="response-456",
+ created_at=1625247600,
+ output=[
+ ResponseReasoningItem(
+ id="reasoning-001",
+ type="reasoning",
+ summary=[
+ Summary(
+ text="Let me analyze this problem step by step...",
+ type="summary_text"
+ )
+ ]
+ ),
+ ResponseOutputMessage(
+ id="msg-001",
+ type="message",
+ role="assistant",
+ status="completed",
+ content=[
+ ResponseOutputText(
+ annotations=[],
+ text="The weather in San Francisco is sunny, 20°C.",
+ type="output_text",
+ )
+ ]
+ )
+ ]
+ )
+
+ kwargs = {
+ "call_type": "responses",
+ "messages": [{"role": "user", "content": "What's the weather in San Francisco?"}],
+ "model": "gpt-4o",
+ "optional_params": {},
+ }
+
+ mock_span = MagicMock()
+
+ with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute:
+ LangfuseOtelLogger._set_langfuse_specific_attributes(mock_span, kwargs, response_obj)
+
+ # Verify observation output was set
+ output_calls = [
+ call for call in mock_safe_set_attribute.call_args_list
+ if call.args[1] == LangfuseSpanAttributes.OBSERVATION_OUTPUT.value
+ ]
+
+ assert len(output_calls) > 0, "observation.output should be set"
+ output_json = output_calls[0].args[2]
+ output_data = json.loads(output_json)
+
+ # Verify output contains reasoning and message
+ assert isinstance(output_data, list)
+ assert len(output_data) == 2
+
+ # Verify reasoning summary
+ assert output_data[0]["role"] == "reasoning_summary"
+ assert output_data[0]["content"] == "Let me analyze this problem step by step..."
+
+ # Verify message
+ assert output_data[1]["role"] == "assistant"
+ assert output_data[1]["content"] == "The weather in San Francisco is sunny, 20°C."
+
+ def test_responses_api_with_function_calls(self):
+ """Test Langfuse OTEL logger with Responses API function_call output."""
+ from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes
+ from openai.types.responses import ResponseFunctionToolCall
+
+ # Create Responses API response with function call
+ response_obj = ResponsesAPIResponse(
+ id="response-789",
+ created_at=1625247700,
+ output=[
+ ResponseFunctionToolCall(
+ id="fc-123",
+ type="function_call",
+ name="get_weather",
+ call_id="call-abc",
+ arguments='{"location": "San Francisco", "unit": "celsius"}',
+ status="completed"
+ )
+ ]
+ )
+
+ kwargs = {
+ "call_type": "responses",
+ "messages": [{"role": "user", "content": "What's the weather in San Francisco?"}],
+ "model": "gpt-4o",
+ "optional_params": {},
+ }
+
+ mock_span = MagicMock()
+
+ with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute:
+ LangfuseOtelLogger._set_langfuse_specific_attributes(mock_span, kwargs, response_obj)
+
+ # Verify observation output was set
+ output_calls = [
+ call for call in mock_safe_set_attribute.call_args_list
+ if call.args[1] == LangfuseSpanAttributes.OBSERVATION_OUTPUT.value
+ ]
+
+ assert len(output_calls) > 0, "observation.output should be set"
+ output_json = output_calls[0].args[2]
+ output_data = json.loads(output_json)
+
+ # Verify output contains function call
+ assert isinstance(output_data, list)
+ assert len(output_data) == 1
+
+ # Verify function call details
+ assert output_data[0]["type"] == "function_call"
+ assert output_data[0]["id"] == "fc-123"
+ assert output_data[0]["name"] == "get_weather"
+ assert output_data[0]["call_id"] == "call-abc"
+ assert output_data[0]["arguments"]["location"] == "San Francisco"
+ assert output_data[0]["arguments"]["unit"] == "celsius"
+
if __name__ == "__main__":
pytest.main([__file__])
\ No newline at end of file
diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py
index 68ca7bb1ecb..2774d168701 100644
--- a/tests/test_litellm/integrations/test_opentelemetry.py
+++ b/tests/test_litellm/integrations/test_opentelemetry.py
@@ -40,7 +40,7 @@ class TestOpenTelemetryGuardrails(unittest.TestCase):
}
# Create a kwargs dict with standard_logging_object containing guardrail information
- kwargs = {"standard_logging_object": {"guardrail_information": guardrail_info}}
+ kwargs = {"standard_logging_object": {"guardrail_information": [ guardrail_info ]}}
# Call the method
otel._create_guardrail_span(kwargs=kwargs, context=None)
@@ -156,7 +156,7 @@ class TestOpenTelemetry(unittest.TestCase):
}
# Create a kwargs dict with standard_logging_object containing guardrail information
- kwargs = {"standard_logging_object": {"guardrail_information": guardrail_info}}
+ kwargs = {"standard_logging_object": {"guardrail_information": [ guardrail_info ]}}
# Call the method
otel._create_guardrail_span(kwargs=kwargs, context=None)
diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/test_litellm/integrations/test_s3_v2.py
index d31d783ba0a..12ec8af1ef4 100644
--- a/tests/test_litellm/integrations/test_s3_v2.py
+++ b/tests/test_litellm/integrations/test_s3_v2.py
@@ -154,4 +154,145 @@ class TestS3V2UnitTests:
expected_download_url = "https://download.s3.endpoint.com/download-bucket/2025-09-14/download-test-key.json"
assert url_download == expected_download_url, f"Expected download URL {expected_download_url}, got {url_download}"
- assert result == {"downloaded": "data"}
\ No newline at end of file
+ assert result == {"downloaded": "data"}
+
+@pytest.mark.asyncio
+async def test_strip_base64_removes_file_and_nontext_entries():
+ logger = S3Logger(s3_strip_base64_files=True)
+
+ payload = {
+ "messages": [
+ {
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "Hello world"},
+ {"type": "image", "file": {"file_data": "data:image/png;base64,AAAA"}},
+ {"type": "file", "file": {"file_data": "data:application/pdf;base64,BBBB"}},
+ ],
+ },
+ {
+ "role": "assistant",
+ "content": [
+ {"type": "text", "text": "Response"},
+ {"type": "audio", "file": {"file_data": "data:audio/wav;base64,CCCC"}},
+ ],
+ },
+ ]
+ }
+
+ stripped = await logger._strip_base64_from_messages(payload)
+
+ # 1️⃣ File/image/audio entries are removed
+ assert len(stripped["messages"][0]["content"]) == 1
+ assert stripped["messages"][0]["content"][0]["text"] == "Hello world"
+
+ assert len(stripped["messages"][1]["content"]) == 1
+ assert stripped["messages"][1]["content"][0]["text"] == "Response"
+
+ # 2️⃣ No 'file' keys remain
+ for msg in stripped["messages"]:
+ for content in msg["content"]:
+ assert "file" not in content
+ assert content.get("type") == "text"
+
+
+@pytest.mark.asyncio
+async def test_strip_base64_keeps_non_file_content():
+ logger = S3Logger(s3_strip_base64_files=True)
+
+ payload = {
+ "messages": [
+ {
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "Just text"},
+ {"type": "text", "text": "Another message"},
+ ],
+ }
+ ]
+ }
+
+ stripped = await logger._strip_base64_from_messages(payload)
+
+ # Should not modify pure text messages
+ assert stripped["messages"][0]["content"] == payload["messages"][0]["content"]
+
+
+@pytest.mark.asyncio
+async def test_strip_base64_handles_empty_or_missing_messages():
+ logger = S3Logger(s3_strip_base64_files=True)
+
+ # Missing messages key
+ payload_no_messages = {}
+ stripped1 = await logger._strip_base64_from_messages(payload_no_messages)
+ assert stripped1 == payload_no_messages
+
+ # Empty messages list
+ payload_empty = {"messages": []}
+ stripped2 = await logger._strip_base64_from_messages(payload_empty)
+ assert stripped2 == payload_empty
+
+
+@pytest.mark.asyncio
+async def test_strip_base64_mixed_nested_objects():
+ """
+ Handles weird/nested content structures gracefully.
+ """
+ logger = S3Logger(s3_strip_base64_files=True)
+
+ payload = {
+ "messages": [
+ {
+ "role": "system",
+ "content": [
+ {"type": "text", "text": "Keep me"},
+ {"type": "custom", "metadata": "ignore but non-text"},
+ {"foo": "bar"},
+ {"file": {"file_data": "data:application/pdf;base64,XXX"}},
+ ],
+ "extra": {"trace_id": "123"},
+ }
+ ]
+ }
+
+ stripped = await logger._strip_base64_from_messages(payload)
+
+ # Custom/non-text and file entries removed
+ content = stripped["messages"][0]["content"]
+ assert len(content) == 2
+ assert {"type": "text", "text": "Keep me"} in content
+ assert {"foo": "bar"} in content
+ # Extra metadata preserved
+ assert stripped["messages"][0]["extra"]["trace_id"] == "123"
+
+
+@pytest.mark.asyncio
+async def test_strip_base64_recursive_redaction():
+ logger = S3Logger(s3_strip_base64_files=True)
+ payload = {
+ "messages": [
+ {
+ "content": [
+ {"type": "text", "text": "normal text"},
+ {"type": "text", "text": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUg"},
+ {"type": "text", "text": "Nested: {'data': 'data:application/pdf;base64,AAA...'}"},
+ {"file": {"file_data": "data:application/pdf;base64,AAAA"}},
+ {"metadata": {"preview": "data:audio/mp3;base64,AAAAA=="}},
+ ]
+ }
+ ]
+ }
+
+ result = await logger._strip_base64_from_messages(payload)
+ content = result["messages"][0]["content"]
+
+ # Dropped file-type entries
+ assert not any("file" in c for c in content)
+
+ # Base64 redacted globally
+ import json
+ for c in content:
+ if isinstance(c, dict):
+ s = json.dumps(c).lower()
+ # "[base64_redacted]" is fine, but raw base64 is not
+ assert "base64," not in s, f"Found real base64 blob in: {s}"
diff --git a/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py b/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py
index 6a0f4f8fa33..0962206476d 100644
--- a/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py
+++ b/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py
@@ -2,6 +2,7 @@ import io
import os
import pathlib
import sys
+from unittest.mock import MagicMock
import pytest
@@ -16,6 +17,7 @@ from litellm.llms.base_llm.audio_transcription.transformation import (
from litellm.llms.deepgram.audio_transcription.transformation import (
DeepgramAudioTranscriptionConfig,
)
+from litellm.types.utils import TranscriptionResponse
@pytest.fixture
@@ -238,3 +240,221 @@ def test_get_complete_url_with_detect_language_and_other_params():
assert "punctuate=true" in url
assert "diarize=false" in url
assert url.startswith("https://api.deepgram.com/v1/listen?")
+
+
+def test_transform_response_without_diarization():
+ """Test response transformation without diarization"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ # Mock response without diarization
+ mock_response = MagicMock()
+ mock_response.json.return_value = {
+ "metadata": {
+ "duration": 10.5,
+ },
+ "results": {
+ "channels": [
+ {
+ "alternatives": [
+ {
+ "transcript": "Hello this is a test.",
+ "confidence": 0.99,
+ "words": [
+ {"word": "Hello", "start": 0.0, "end": 0.5},
+ {"word": "this", "start": 0.6, "end": 0.8},
+ {"word": "is", "start": 0.9, "end": 1.1},
+ {"word": "a", "start": 1.2, "end": 1.3},
+ {"word": "test", "start": 1.4, "end": 1.8},
+ ],
+ }
+ ]
+ }
+ ]
+ },
+ }
+
+ result = handler.transform_audio_transcription_response(mock_response)
+
+ assert isinstance(result, TranscriptionResponse)
+ assert result.text == "Hello this is a test."
+ assert result["task"] == "transcribe"
+ assert result["duration"] == 10.5
+ assert len(result["words"]) == 5
+
+
+def test_transform_response_with_diarization_and_paragraphs():
+ """Test response transformation with diarization and paragraphs property"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ # Mock response with diarization and paragraphs
+ mock_response = MagicMock()
+ mock_response.json.return_value = {
+ "metadata": {
+ "duration": 15.0,
+ },
+ "results": {
+ "channels": [
+ {
+ "alternatives": [
+ {
+ "transcript": "Hello how are you I am fine thanks",
+ "paragraphs": {
+ "transcript": "\nSpeaker 0: Hello how are you\n\nSpeaker 1: I am fine thanks\n"
+ },
+ "words": [
+ {"word": "Hello", "start": 0.0, "end": 0.5, "speaker": 0},
+ {"word": "how", "start": 0.6, "end": 0.8, "speaker": 0},
+ {"word": "are", "start": 0.9, "end": 1.1, "speaker": 0},
+ {"word": "you", "start": 1.2, "end": 1.3, "speaker": 0},
+ {"word": "I", "start": 2.0, "end": 2.2, "speaker": 1},
+ {"word": "am", "start": 2.3, "end": 2.5, "speaker": 1},
+ {"word": "fine", "start": 2.6, "end": 2.9, "speaker": 1},
+ {"word": "thanks", "start": 3.0, "end": 3.5, "speaker": 1},
+ ],
+ }
+ ]
+ }
+ ]
+ },
+ }
+
+ result = handler.transform_audio_transcription_response(mock_response)
+
+ assert isinstance(result, TranscriptionResponse)
+ # Should use the pre-formatted paragraphs transcript
+ assert result.text == "\nSpeaker 0: Hello how are you\n\nSpeaker 1: I am fine thanks\n"
+ assert result["task"] == "transcribe"
+ assert result["duration"] == 15.0
+
+
+def test_transform_response_with_diarization_without_paragraphs():
+ """Test response transformation with diarization but no paragraphs property"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ # Mock response with diarization but without paragraphs
+ mock_response = MagicMock()
+ mock_response.json.return_value = {
+ "metadata": {
+ "duration": 15.0,
+ },
+ "results": {
+ "channels": [
+ {
+ "alternatives": [
+ {
+ "transcript": "Hello how are you I am fine thanks",
+ "words": [
+ {"word": "hello", "punctuated_word": "Hello", "start": 0.0, "end": 0.5, "speaker": 0},
+ {"word": "how", "punctuated_word": "how", "start": 0.6, "end": 0.8, "speaker": 0},
+ {"word": "are", "punctuated_word": "are", "start": 0.9, "end": 1.1, "speaker": 0},
+ {"word": "you", "punctuated_word": "you", "start": 1.2, "end": 1.3, "speaker": 0},
+ {"word": "i", "punctuated_word": "I", "start": 2.0, "end": 2.2, "speaker": 1},
+ {"word": "am", "punctuated_word": "am", "start": 2.3, "end": 2.5, "speaker": 1},
+ {"word": "fine", "punctuated_word": "fine", "start": 2.6, "end": 2.9, "speaker": 1},
+ {"word": "thanks", "punctuated_word": "thanks.", "start": 3.0, "end": 3.5, "speaker": 1},
+ ],
+ }
+ ]
+ }
+ ]
+ },
+ }
+
+ result = handler.transform_audio_transcription_response(mock_response)
+
+ assert isinstance(result, TranscriptionResponse)
+ # Should reconstruct from words using punctuated_word
+ expected_text = "Speaker 0: Hello how are you\n\nSpeaker 1: I am fine thanks.\n"
+ assert result.text == expected_text
+ assert result["task"] == "transcribe"
+ assert result["duration"] == 15.0
+
+
+def test_reconstruct_diarized_transcript_with_punctuated_words():
+ """Test reconstruction uses punctuated_word when available"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ words = [
+ {"word": "hello", "punctuated_word": "Hello", "speaker": 0},
+ {"word": "world", "punctuated_word": "world!", "speaker": 0},
+ {"word": "how", "punctuated_word": "How", "speaker": 1},
+ {"word": "are", "punctuated_word": "are", "speaker": 1},
+ {"word": "you", "punctuated_word": "you?", "speaker": 1},
+ ]
+
+ result = handler._reconstruct_diarized_transcript(words)
+
+ # Check that punctuated_word is used and speakers are properly separated
+ assert "Hello world!" in result
+ assert "How are you?" in result
+ assert "Speaker 0:" in result
+ assert "Speaker 1:" in result
+
+
+def test_reconstruct_diarized_transcript_fallback_to_word():
+ """Test reconstruction falls back to 'word' when punctuated_word is missing"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ words = [
+ {"word": "Hello", "speaker": 0}, # No punctuated_word
+ {"word": "world", "speaker": 0},
+ {"word": "test", "punctuated_word": "test.", "speaker": 1}, # Has punctuated_word
+ ]
+
+ result = handler._reconstruct_diarized_transcript(words)
+
+ # Should use 'word' when punctuated_word is not available
+ assert "Hello world" in result
+ assert "test." in result
+ assert "Speaker 0:" in result
+ assert "Speaker 1:" in result
+
+
+def test_reconstruct_diarized_transcript_empty_words():
+ """Test reconstruction with empty words list"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ result = handler._reconstruct_diarized_transcript([])
+
+ assert result == ""
+
+
+def test_reconstruct_diarized_transcript_single_speaker():
+ """Test reconstruction with single speaker"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ words = [
+ {"word": "This", "punctuated_word": "This", "speaker": 0},
+ {"word": "is", "punctuated_word": "is", "speaker": 0},
+ {"word": "a", "punctuated_word": "a", "speaker": 0},
+ {"word": "test", "punctuated_word": "test.", "speaker": 0},
+ ]
+
+ result = handler._reconstruct_diarized_transcript(words)
+
+ # Should have only one speaker segment
+ assert result.count("Speaker 0:") == 1
+ assert "This is a test." in result
+
+
+def test_reconstruct_diarized_transcript_multiple_speaker_changes():
+ """Test reconstruction with multiple speaker changes"""
+ handler = DeepgramAudioTranscriptionConfig()
+
+ words = [
+ {"word": "Hi", "speaker": 0},
+ {"word": "there", "speaker": 0},
+ {"word": "Hello", "speaker": 1},
+ {"word": "back", "speaker": 0}, # Speaker 0 again
+ {"word": "Thanks", "speaker": 1}, # Speaker 1 again
+ ]
+
+ result = handler._reconstruct_diarized_transcript(words)
+
+ # Should have 4 speaker segments (0, 1, 0, 1)
+ assert result.count("Speaker 0:") == 2
+ assert result.count("Speaker 1:") == 2
+ assert "Hi there" in result
+ assert "Hello" in result
+ assert "back" in result
+ assert "Thanks" in result
diff --git a/tests/test_litellm/llms/perplexity/chat/test_perplexity_chat_transformation.py b/tests/test_litellm/llms/perplexity/chat/test_perplexity_chat_transformation.py
index 784e6f6fe63..d83c1a9f659 100644
--- a/tests/test_litellm/llms/perplexity/chat/test_perplexity_chat_transformation.py
+++ b/tests/test_litellm/llms/perplexity/chat/test_perplexity_chat_transformation.py
@@ -7,7 +7,7 @@ from Perplexity API responses.
import os
import sys
-from unittest.mock import Mock
+from unittest.mock import Mock, patch
import pytest
@@ -707,4 +707,152 @@ class TestPerplexityChatTransformation:
# Check that no annotations were created (message content is None)
assert choice.message.content is None
# No annotations should be created since content is None
- assert not hasattr(choice.message, 'annotations') or choice.message.annotations is None
\ No newline at end of file
+ assert not hasattr(choice.message, 'annotations') or choice.message.annotations is None
+
+ # Tests for cost extraction functionality
+ def test_add_cost_to_usage_flat_structure(self):
+ """Test cost extraction from flat usage structure."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response with flat cost structure
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "total_cost": 0.00015
+ }
+ }
+
+ # Test cost extraction
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Check that cost was stored in hidden params
+ assert hasattr(model_response, "_hidden_params")
+ assert "additional_headers" in model_response._hidden_params
+ assert "llm_provider-x-litellm-response-cost" in model_response._hidden_params["additional_headers"]
+
+ cost = model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"]
+ assert cost == 0.00015
+
+ def test_add_cost_to_usage_nested_structure(self):
+ """Test cost extraction from nested usage structure."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response with nested cost structure
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "cost": {
+ "total_cost": 0.00025
+ }
+ }
+ }
+
+ # Test cost extraction
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Check that cost was stored in hidden params
+ assert hasattr(model_response, "_hidden_params")
+ assert "additional_headers" in model_response._hidden_params
+ assert "llm_provider-x-litellm-response-cost" in model_response._hidden_params["additional_headers"]
+
+ cost = model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"]
+ assert cost == 0.00025
+
+ def test_add_cost_to_usage_no_cost_data(self):
+ """Test handling when no cost data is present."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response without cost
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150
+ }
+ }
+
+ # Test cost extraction - should not raise error
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Should not have cost in hidden params
+ if hasattr(model_response, "_hidden_params"):
+ assert "llm_provider-x-litellm-response-cost" not in model_response._hidden_params.get("additional_headers", {})
+
+ def test_transform_response_includes_cost_extraction(self):
+ """Test that transform_response includes cost extraction."""
+ config = PerplexityChatConfig()
+
+ # Mock raw response
+ mock_response = Mock()
+ mock_response.json.return_value = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "total_cost": 0.00015
+ }
+ }
+ mock_response.headers = {}
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+ model_response.model = "perplexity/sonar-pro"
+
+ # Mock the parent transform_response to return our model_response
+ with patch.object(config.__class__.__bases__[0], 'transform_response', return_value=model_response):
+ result = config.transform_response(
+ model="perplexity/sonar-pro",
+ raw_response=mock_response,
+ model_response=model_response,
+ logging_obj=Mock(),
+ request_data={},
+ messages=[{"role": "user", "content": "Test"}],
+ optional_params={},
+ litellm_params={},
+ encoding=None,
+ )
+
+ # Check that cost was extracted and stored
+ assert hasattr(result, "_hidden_params")
+ assert "additional_headers" in result._hidden_params
+ assert "llm_provider-x-litellm-response-cost" in result._hidden_params["additional_headers"]
+
+ cost = result._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"]
+ assert cost == 0.00015
\ No newline at end of file
diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py
index f9a52100070..58b037269a7 100644
--- a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py
+++ b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py
@@ -17,7 +17,8 @@ import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
-from litellm.cost_calculator import completion_cost, cost_per_token
+from litellm import ModelResponse
+from litellm.cost_calculator import completion_cost, cost_per_token, response_cost_calculator
from litellm.llms.perplexity.cost_calculator import cost_per_token as perplexity_cost_per_token
from litellm.types.utils import Usage, PromptTokensDetailsWrapper
from litellm.utils import get_model_info
@@ -370,4 +371,116 @@ class TestPerplexityCostCalculator:
# Ensure costs are non-negative
assert prompt_cost >= 0
- assert completion_cost >= 0
\ No newline at end of file
+ assert completion_cost >= 0
+
+ def test_cost_extraction_priority_over_calculation(self):
+ """Test that extracted cost from API response takes priority over calculated cost."""
+ from litellm.cost_calculator import response_cost_calculator
+ from litellm.llms.perplexity.chat.transformation import PerplexityChatConfig
+
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse with extracted cost
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+ model_response.model = "perplexity/sonar-pro"
+
+ # Mock raw response with cost
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "total_cost": 0.00015 # This should be used instead of calculated cost
+ }
+ }
+
+ # Extract cost from API response
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Test response cost calculator - should use extracted cost
+ cost = response_cost_calculator(
+ response_object=model_response,
+ model="perplexity/sonar-pro",
+ custom_llm_provider="perplexity",
+ call_type="completion",
+ optional_params={}
+ )
+
+ # Should return the extracted cost, not calculated cost
+ assert cost == 0.00015
+
+ def test_cost_extraction_from_nested_structure(self):
+ """Test cost extraction from nested usage structure."""
+ from litellm.llms.perplexity.chat.transformation import PerplexityChatConfig
+
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response with nested cost structure
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "cost": {
+ "total_cost": 0.00025
+ }
+ }
+ }
+
+ # Test cost extraction
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Check that cost was stored in hidden params
+ assert hasattr(model_response, "_hidden_params")
+ assert "additional_headers" in model_response._hidden_params
+ assert "llm_provider-x-litellm-response-cost" in model_response._hidden_params["additional_headers"]
+
+ cost = model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"]
+ assert cost == 0.00025
+
+ def test_cost_extraction_error_handling(self):
+ """Test error handling during cost extraction."""
+ from litellm.llms.perplexity.chat.transformation import PerplexityChatConfig
+
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response with invalid cost data
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "cost": "invalid_cost" # Invalid cost type
+ }
+ }
+
+ # Test cost extraction - should not raise error
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Should not have cost in hidden params due to error
+ if hasattr(model_response, "_hidden_params"):
+ assert "llm_provider-x-litellm-response-cost" not in model_response._hidden_params.get("additional_headers", {})
\ No newline at end of file
diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
index a702e9ebc4c..b38e56ce8e8 100644
--- a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
+++ b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py
@@ -316,4 +316,169 @@ class TestPerplexityIntegration:
expected_completion_cost = (50 * 8e-6) + (1 / 1000 * 0.005)
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
- assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
\ No newline at end of file
+ assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
+
+ def test_cost_extraction_from_api_response(self):
+ """Test cost extraction from Perplexity API response."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+ model_response.model = "perplexity/sonar-pro"
+
+ # Mock raw response with cost
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "total_cost": 0.00015
+ }
+ }
+
+ # Test cost extraction
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Check that cost was stored in hidden params
+ assert hasattr(model_response, "_hidden_params")
+ assert "additional_headers" in model_response._hidden_params
+ assert "llm_provider-x-litellm-response-cost" in model_response._hidden_params["additional_headers"]
+
+ cost = model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"]
+ assert cost == 0.00015
+
+ def test_cost_extraction_integration_with_main_calculator(self):
+ """Test that extracted cost takes priority over calculated cost."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+ model_response.model = "perplexity/sonar-pro"
+
+ # Mock raw response with cost
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "total_cost": 0.00015
+ }
+ }
+
+ # Extract cost
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Test main cost calculator - should use extracted cost
+ from litellm.cost_calculator import response_cost_calculator
+ cost = response_cost_calculator(
+ response_object=model_response,
+ model="perplexity/sonar-pro",
+ custom_llm_provider="perplexity",
+ call_type="completion",
+ optional_params={}
+ )
+
+ # Should return the extracted cost, not calculated cost
+ assert cost == 0.00015
+
+ def test_cost_extraction_nested_structure(self):
+ """Test cost extraction from nested usage structure."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response with nested cost structure
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "cost": {
+ "total_cost": 0.00025
+ }
+ }
+ }
+
+ # Test cost extraction
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Check that cost was stored in hidden params
+ assert hasattr(model_response, "_hidden_params")
+ assert "additional_headers" in model_response._hidden_params
+ assert "llm_provider-x-litellm-response-cost" in model_response._hidden_params["additional_headers"]
+
+ cost = model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"]
+ assert cost == 0.00025
+
+ def test_cost_extraction_error_handling(self):
+ """Test error handling during cost extraction."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response with invalid cost data
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}],
+ "usage": {
+ "prompt_tokens": 100,
+ "completion_tokens": 50,
+ "total_tokens": 150,
+ "cost": "invalid_cost" # Invalid cost type
+ }
+ }
+
+ # Test cost extraction - should not raise error
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Should not have cost in hidden params due to error
+ if hasattr(model_response, "_hidden_params"):
+ assert "llm_provider-x-litellm-response-cost" not in model_response._hidden_params.get("additional_headers", {})
+
+ def test_cost_extraction_no_usage_data(self):
+ """Test handling when no usage data is present."""
+ config = PerplexityChatConfig()
+
+ # Create a ModelResponse
+ model_response = ModelResponse()
+ model_response.usage = Usage(
+ prompt_tokens=100,
+ completion_tokens=50,
+ total_tokens=150
+ )
+
+ # Mock raw response without usage
+ raw_response_json = {
+ "choices": [{"message": {"content": "Test response"}}]
+ }
+
+ # Test cost extraction - should not raise error
+ config.add_cost_to_usage(model_response, raw_response_json)
+
+ # Should not have cost in hidden params
+ if hasattr(model_response, "_hidden_params"):
+ assert "llm_provider-x-litellm-response-cost" not in model_response._hidden_params.get("additional_headers", {})
\ No newline at end of file
diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py
index 02dac0a93d6..db457999a03 100644
--- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py
+++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py
@@ -856,4 +856,121 @@ def test_get_token_url():
assert "v1beta1" not in url
assert "/v1/" in url
- pass
\ No newline at end of file
+ pass
+
+
+@pytest.mark.asyncio
+async def test_vertex_ai_token_counter_routes_partner_models():
+ """
+ Test that VertexAITokenCounter correctly routes partner models (Claude, Mistral, etc.)
+ to the partner models token counter instead of the Gemini token counter.
+ """
+ from unittest.mock import AsyncMock, patch
+ from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter
+ from litellm.types.utils import TokenCountResponse
+
+ token_counter = VertexAITokenCounter()
+
+ # Mock the partner models handler
+ with patch(
+ "litellm.llms.vertex_ai.vertex_ai_partner_models.main.VertexAIPartnerModels.count_tokens"
+ ) as mock_partner_count_tokens:
+ mock_partner_count_tokens.return_value = {
+ "input_tokens": 42,
+ "tokenizer_used": "vertex_ai_partner_models",
+ }
+
+ # Test with a Claude model (partner model)
+ result = await token_counter.count_tokens(
+ model_to_use="claude-3-5-sonnet-20241022",
+ messages=[{"role": "user", "content": "Hello"}],
+ contents=None,
+ deployment={
+ "litellm_params": {
+ "vertex_project": "test-project",
+ "vertex_location": "us-east5",
+ }
+ },
+ request_model="vertex_ai/claude-3-5-sonnet-20241022",
+ )
+
+ # Verify partner models handler was called
+ assert mock_partner_count_tokens.called
+ assert result is not None
+ assert isinstance(result, TokenCountResponse)
+ assert result.total_tokens == 42
+ assert result.tokenizer_type == "vertex_ai_partner_models"
+
+
+@pytest.mark.asyncio
+async def test_vertex_ai_token_counter_routes_gemini_models():
+ """
+ Test that VertexAITokenCounter correctly routes Gemini models
+ to the Gemini token counter (not partner models).
+ """
+ from unittest.mock import AsyncMock, patch
+ from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter
+ from litellm.types.utils import TokenCountResponse
+
+ token_counter = VertexAITokenCounter()
+
+ # Mock the Gemini handler (different import path)
+ with patch(
+ "litellm.llms.vertex_ai.count_tokens.handler.VertexAITokenCounter.acount_tokens"
+ ) as mock_gemini_count_tokens:
+ mock_gemini_count_tokens.return_value = {
+ "totalTokens": 50,
+ "tokenizer_used": "gemini",
+ }
+
+ # Test with a Gemini model (not a partner model)
+ result = await token_counter.count_tokens(
+ model_to_use="gemini-1.5-pro",
+ messages=[{"role": "user", "content": "Hello"}],
+ contents=None,
+ deployment={
+ "litellm_params": {
+ "vertex_project": "test-project",
+ "vertex_location": "us-central1",
+ }
+ },
+ request_model="vertex_ai/gemini-1.5-pro",
+ )
+
+ # Verify Gemini handler was called
+ assert mock_gemini_count_tokens.called
+ assert result is not None
+ assert isinstance(result, TokenCountResponse)
+ assert result.total_tokens == 50
+
+
+@pytest.mark.asyncio
+async def test_vertex_ai_partner_model_detection():
+ """
+ Test that VertexAIPartnerModels.is_vertex_partner_model correctly identifies
+ partner models (Claude, Mistral, Llama, etc.).
+ """
+ from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
+ VertexAIPartnerModels,
+ )
+
+ # Test Claude models (should be detected as partner model)
+ assert VertexAIPartnerModels.is_vertex_partner_model("claude-3-5-sonnet-20241022")
+ assert VertexAIPartnerModels.is_vertex_partner_model("claude-3-opus-20240229")
+ assert VertexAIPartnerModels.is_vertex_partner_model("claude-3-haiku-20240307")
+
+ # Test Mistral models
+ assert VertexAIPartnerModels.is_vertex_partner_model("mistral-large-2407")
+ assert VertexAIPartnerModels.is_vertex_partner_model("mistral-7b-instruct-v0.3")
+
+ # Test Meta/Llama models
+ assert VertexAIPartnerModels.is_vertex_partner_model("meta/llama-3.1-405b")
+
+ # Test Gemini models (should NOT be detected as partner model)
+ assert not VertexAIPartnerModels.is_vertex_partner_model("gemini-1.5-pro")
+ assert not VertexAIPartnerModels.is_vertex_partner_model("gemini-1.0-pro")
+ assert not VertexAIPartnerModels.is_vertex_partner_model("gemini-pro-vision")
+
+ # Test other non-partner models
+ assert not VertexAIPartnerModels.is_vertex_partner_model("text-bison-001")
+ assert not VertexAIPartnerModels.is_vertex_partner_model("chat-bison-001")
\ No newline at end of file
diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
index 42cec94322c..cc23318505d 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py
@@ -1965,6 +1965,83 @@ class TestProcessSSOJWTAccessToken:
assert result.team_ids == []
+class TestGenericResponseConvertorNestedAttributes:
+ """Test generic_response_convertor with nested attribute paths"""
+
+ def test_generic_response_convertor_with_nested_attributes(self):
+ """
+ Test that generic_response_convertor handles nested attributes with dotted notation
+ like "attributes.userId"
+ """
+ from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
+
+ # Mock JWT handler
+ mock_jwt_handler = MagicMock(spec=JWTHandler)
+ mock_jwt_handler.get_team_ids_from_jwt.return_value = []
+
+ # Payload with nested attributes structure
+ nested_payload = {
+ "sub": "user-sub-123",
+ "service": "test-service",
+ "auth_time": 1234567890,
+ "attributes": {
+ "given_name": "John",
+ "oauthClientId": "client-123",
+ "family_name": "Doe",
+ "userId": "nested-user-456",
+ "email": "john.doe@example.com",
+ },
+ "id": "top-level-id-789",
+ "client_id": "client-abc",
+ }
+
+ # Test with nested user ID attribute
+ with patch.dict(
+ os.environ,
+ {
+ "GENERIC_USER_ID_ATTRIBUTE": "attributes.userId",
+ "GENERIC_USER_EMAIL_ATTRIBUTE": "attributes.email",
+ "GENERIC_USER_FIRST_NAME_ATTRIBUTE": "attributes.given_name",
+ "GENERIC_USER_LAST_NAME_ATTRIBUTE": "attributes.family_name",
+ "GENERIC_USER_DISPLAY_NAME_ATTRIBUTE": "sub",
+ },
+ ):
+ # Act
+ result = generic_response_convertor(
+ response=nested_payload,
+ jwt_handler=mock_jwt_handler,
+ sso_jwt_handler=None,
+ )
+
+ # Assert
+ assert isinstance(result, CustomOpenID)
+
+ # Note: The current implementation uses response.get() which doesn't support
+ # dotted notation for nested attributes. This test documents the current behavior.
+ # If nested attribute support is needed, the implementation would need to be updated
+ # to handle dotted paths like "attributes.userId"
+
+ # Current behavior: returns None for nested paths
+ print(f"User ID result: {result.id}")
+ print(f"Email result: {result.email}")
+ print(f"First name result: {result.first_name}")
+ print(f"Last name result: {result.last_name}")
+ print(f"Display name result: {result.display_name}")
+
+ # Expected behavior with current implementation (no nested path support):
+ assert result.id == "nested-user-456"
+ assert (
+ result.email == "john.doe@example.com"
+ ) # Can't access "attributes.email" with .get()
+ assert (
+ result.first_name == "John"
+ ) # Can't access "attributes.given_name" with .get()
+ assert (
+ result.last_name == "Doe"
+ ) # Can't access "attributes.family_name" with .get()
+ assert result.display_name == "user-sub-123" # Top-level attribute works
+
+
class TestPKCEFunctionality:
"""Test PKCE (Proof Key for Code Exchange) functionality"""
@@ -1983,15 +2060,21 @@ class TestPKCEFunctionality:
# Assert
assert len(code_verifier) == 43
assert isinstance(code_verifier, str)
-
+
# Verify code_challenge is correctly generated from code_verifier
- expected_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest()
- expected_challenge = base64.urlsafe_b64encode(expected_challenge_bytes).decode('utf-8').rstrip('=')
+ expected_challenge_bytes = hashlib.sha256(
+ code_verifier.encode("utf-8")
+ ).digest()
+ expected_challenge = (
+ base64.urlsafe_b64encode(expected_challenge_bytes)
+ .decode("utf-8")
+ .rstrip("=")
+ )
assert code_challenge == expected_challenge
-
+
# Verify both are base64url encoded (no padding)
- assert '=' not in code_verifier
- assert '=' not in code_challenge
+ assert "=" not in code_verifier
+ assert "=" not in code_challenge
@pytest.mark.asyncio
async def test_prepare_token_exchange_parameters_with_pkce(self):
@@ -2013,17 +2096,20 @@ class TestPKCEFunctionality:
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
# Act
token_params = SSOAuthenticationHandler.prepare_token_exchange_parameters(
- request=mock_request,
- generic_include_client_id=False
+ request=mock_request, generic_include_client_id=False
)
# Assert
assert token_params["include_client_id"] is False
assert token_params["code_verifier"] == test_code_verifier
-
+
# Verify cache was accessed and deleted
- mock_cache.get_cache.assert_called_once_with(key=f"pkce_verifier:{test_state}")
- mock_cache.delete_cache.assert_called_once_with(key=f"pkce_verifier:{test_state}")
+ mock_cache.get_cache.assert_called_once_with(
+ key=f"pkce_verifier:{test_state}"
+ )
+ mock_cache.delete_cache.assert_called_once_with(
+ key=f"pkce_verifier:{test_state}"
+ )
@pytest.mark.asyncio
async def test_get_generic_sso_redirect_response_with_pkce(self):
@@ -2035,7 +2121,9 @@ class TestPKCEFunctionality:
# Mock SSO provider
mock_sso = MagicMock()
mock_redirect_response = MagicMock()
- original_location = "https://auth.example.com/authorize?state=test456&client_id=abc"
+ original_location = (
+ "https://auth.example.com/authorize?state=test456&client_id=abc"
+ )
mock_redirect_response.headers = {"location": original_location}
mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response)
mock_sso.__enter__ = MagicMock(return_value=mock_sso)
@@ -2050,7 +2138,7 @@ class TestPKCEFunctionality:
result = await SSOAuthenticationHandler.get_generic_sso_redirect_response(
generic_sso=mock_sso,
state=test_state,
- generic_authorization_endpoint="https://auth.example.com/authorize"
+ generic_authorization_endpoint="https://auth.example.com/authorize",
)
# Assert
diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
index 84e53a903a5..ea1017e1d5a 100644
--- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
+++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
@@ -21,6 +21,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
bedrock_llm_proxy_route,
create_pass_through_route,
llm_passthrough_factory_proxy_route,
+ milvus_proxy_route,
vertex_discovery_proxy_route,
vertex_proxy_route,
vllm_proxy_route,
@@ -179,11 +180,8 @@ class TestBaseOpenAIPassThroughHandler:
print("Verifying endpoint_func call parameters...")
mock_endpoint_func.assert_awaited_once()
assert mock_endpoint_func.await_args is not None
- call_kwargs = mock_endpoint_func.await_args[1]
- print(f"stream parameter: {call_kwargs['stream']}")
- print(f"query_params: {call_kwargs['query_params']}")
- assert call_kwargs["stream"] is False
- assert call_kwargs["query_params"] == {"model": "gpt-4"}
+ # The endpoint_func is called with request, fastapi_response, user_api_key_dict
+ # No longer checking for stream and query_params as they're handled differently
class TestVertexAIPassThroughHandler:
@@ -291,6 +289,7 @@ class TestVertexAIPassThroughHandler:
endpoint=endpoint,
target=f"https://{test_location}-aiplatform.googleapis.com/v1/projects/{test_project}/locations/{test_location}/publishers/google/models/gemini-1.5-flash:generateContent",
custom_headers={"Authorization": f"Bearer {test_token}"},
+ is_streaming_request=False,
)
@pytest.mark.asyncio
@@ -389,6 +388,7 @@ class TestVertexAIPassThroughHandler:
endpoint=endpoint,
target=f"https://aiplatform.googleapis.com/v1/projects/{test_project}/locations/{test_location}/publishers/google/models/gemini-1.5-flash:generateContent",
custom_headers={"Authorization": f"Bearer {test_token}"},
+ is_streaming_request=False,
)
@pytest.mark.parametrize(
@@ -481,6 +481,7 @@ class TestVertexAIPassThroughHandler:
endpoint=endpoint,
target=f"https://{default_location}-aiplatform.googleapis.com/v1/projects/{default_project}/locations/{default_location}/publishers/google/models/gemini-1.5-flash:generateContent",
custom_headers={"Authorization": f"Bearer {default_credentials}"},
+ is_streaming_request=False,
)
@pytest.mark.asyncio
@@ -559,6 +560,7 @@ class TestVertexAIPassThroughHandler:
endpoint=endpoint,
target=f"https://{test_location}-aiplatform.googleapis.com/v1/projects/{test_project}/locations/{test_location}/publishers/google/models/gemini-1.5-flash:generateContent",
custom_headers={"authorization": f"Bearer {test_token}"},
+ is_streaming_request=False,
)
@pytest.mark.asyncio
@@ -1134,17 +1136,17 @@ class TestBedrockLLMProxyRoute:
}
mock_llm_router = Mock()
-
+
# Mock ProxyBaseLLMRequestProcessing to raise the httpx error
with patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_passthrough_process_llm_request",
new_callable=AsyncMock,
- side_effect=mock_http_error
+ side_effect=mock_http_error,
):
mock_user_api_key_dict = Mock()
mock_user_api_key_dict.api_key = "test-key"
mock_user_api_key_dict.allowed_model_region = None
-
+
mock_proxy_logging_obj = Mock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
@@ -1226,6 +1228,7 @@ class TestLLMPassthroughFactoryProxyRoute:
endpoint="/chat/completions",
target="https://example.com/v1/chat/completions",
custom_headers={"x-api-key": "dummy"},
+ is_streaming_request=False,
)
mock_endpoint_func.assert_awaited_once()
@@ -1293,3 +1296,453 @@ class TestVLLMProxyRoute:
assert result == "factory_success"
mock_factory_route.assert_awaited_once()
+
+
+class TestMilvusProxyRoute:
+ """
+ Test cases for Milvus passthrough endpoint
+ """
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_success(self):
+ """
+ Test successful Milvus proxy route with valid managed vector store index
+ """
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ collection_name = "dall-e-6"
+ vector_store_name = "milvus-store-1"
+ vector_store_index = "collection_123"
+ api_base = "http://localhost:19530"
+
+ # Mock request
+ mock_request = MagicMock(spec=Request)
+ mock_request.method = "POST"
+ mock_request.headers = {"content-type": "application/json"}
+ mock_request.url = MagicMock()
+ mock_request.url.path = "/milvus/vectors/search"
+
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ # Mock vector store index object
+ mock_index_object = MagicMock()
+ mock_index_object.litellm_params.vector_store_name = vector_store_name
+ mock_index_object.litellm_params.vector_store_index = vector_store_index
+
+ # Mock vector store
+ mock_vector_store = {
+ "litellm_params": {
+ "api_base": api_base,
+ "api_key": "test-milvus-key",
+ }
+ }
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"collectionName": collection_name, "data": [[0.1, 0.2]]},
+ ) as mock_get_body, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
+ ) as mock_is_allowed, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"
+ ) as mock_safe_set, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route"
+ ) as mock_create_route, patch.object(
+ litellm, "vector_store_index_registry"
+ ) as mock_index_registry, patch.object(
+ litellm, "vector_store_registry"
+ ) as mock_vector_registry:
+ # Setup mocks
+ mock_provider_config = MagicMock()
+ mock_provider_config.get_auth_credentials.return_value = {
+ "headers": {"Authorization": "Bearer test-token"}
+ }
+ mock_provider_config.get_complete_url.return_value = api_base
+ mock_get_config.return_value = mock_provider_config
+
+ mock_index_registry.is_vector_store_index.return_value = True
+ mock_index_registry.get_vector_store_index_by_name.return_value = (
+ mock_index_object
+ )
+
+ mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = (
+ mock_vector_store
+ )
+
+ mock_endpoint_func = AsyncMock(
+ return_value={"results": [{"id": 1, "distance": 0.5}]}
+ )
+ mock_create_route.return_value = mock_endpoint_func
+
+ # Call the route
+ result = await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ # Verify calls
+ mock_get_body.assert_called_once()
+ mock_index_registry.is_vector_store_index.assert_called_once_with(
+ vector_store_index_name=collection_name
+ )
+ mock_is_allowed.assert_called_once()
+ mock_safe_set.assert_called_once()
+
+ # Verify collection name was updated to the actual index
+ set_body_call_args = mock_safe_set.call_args[0]
+ assert set_body_call_args[1]["collectionName"] == vector_store_index
+
+ # Verify create_pass_through_route was called with correct URL
+ mock_create_route.assert_called_once()
+ create_route_args = mock_create_route.call_args[1]
+ assert "vectors/search" in create_route_args["target"]
+ assert create_route_args["custom_headers"] == {
+ "Authorization": "Bearer test-token"
+ }
+
+ # Verify endpoint function was called
+ mock_endpoint_func.assert_awaited_once()
+ assert result == {"results": [{"id": 1, "distance": 0.5}]}
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_missing_collection_name(self):
+ """
+ Test that missing collection name raises HTTPException
+ """
+ from fastapi import HTTPException
+
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"data": [[0.1, 0.2]]}, # No collectionName
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config:
+ mock_get_config.return_value = MagicMock()
+
+ with pytest.raises(HTTPException) as exc_info:
+ await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ assert exc_info.value.status_code == 400
+ assert "Collection name is required" in str(exc_info.value.detail)
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_no_provider_config(self):
+ """
+ Test that missing provider config raises HTTPException
+ """
+ from fastapi import HTTPException
+
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config",
+ return_value=None,
+ ):
+ with pytest.raises(HTTPException) as exc_info:
+ await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ assert exc_info.value.status_code == 500
+ assert "Unable to find Milvus vector store config" in str(
+ exc_info.value.detail
+ )
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_no_index_registry(self):
+ """
+ Test that missing index registry raises HTTPException
+ """
+ from fastapi import HTTPException
+
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ collection_name = "test-collection"
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"collectionName": collection_name},
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config, patch.object(
+ litellm, "vector_store_index_registry", None
+ ):
+ mock_get_config.return_value = MagicMock()
+
+ with pytest.raises(HTTPException) as exc_info:
+ await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ assert exc_info.value.status_code == 500
+ assert "Unable to find Milvus vector store index registry" in str(
+ exc_info.value.detail
+ )
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_not_managed_index(self):
+ """
+ Test that non-managed vector store index raises HTTPException
+ """
+ from fastapi import HTTPException
+
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ collection_name = "unmanaged-collection"
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"collectionName": collection_name},
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config, patch.object(
+ litellm, "vector_store_index_registry"
+ ) as mock_index_registry, patch.object(
+ litellm, "vector_store_registry", MagicMock()
+ ):
+ mock_get_config.return_value = MagicMock()
+ mock_index_registry.is_vector_store_index.return_value = False
+
+ with pytest.raises(HTTPException) as exc_info:
+ await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ assert exc_info.value.status_code == 400
+ assert (
+ f"Collection {collection_name} is not a litellm managed vector store index"
+ in str(exc_info.value.detail)
+ )
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_vector_store_not_found(self):
+ """
+ Test that missing vector store raises Exception
+ """
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ collection_name = "test-collection"
+ vector_store_name = "missing-store"
+ vector_store_index = "collection_123"
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ mock_index_object = MagicMock()
+ mock_index_object.litellm_params.vector_store_name = vector_store_name
+ mock_index_object.litellm_params.vector_store_index = vector_store_index
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"collectionName": collection_name},
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"
+ ), patch.object(
+ litellm, "vector_store_index_registry"
+ ) as mock_index_registry, patch.object(
+ litellm, "vector_store_registry"
+ ) as mock_vector_registry:
+ mock_get_config.return_value = MagicMock()
+ mock_index_registry.is_vector_store_index.return_value = True
+ mock_index_registry.get_vector_store_index_by_name.return_value = (
+ mock_index_object
+ )
+ mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = (
+ None
+ )
+
+ with pytest.raises(Exception) as exc_info:
+ await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ assert f"Vector store not found for {vector_store_name}" in str(
+ exc_info.value
+ )
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_no_api_base(self):
+ """
+ Test that missing api_base raises Exception
+ """
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ collection_name = "test-collection"
+ vector_store_name = "milvus-store-1"
+ vector_store_index = "collection_123"
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ mock_index_object = MagicMock()
+ mock_index_object.litellm_params.vector_store_name = vector_store_name
+ mock_index_object.litellm_params.vector_store_index = vector_store_index
+
+ mock_vector_store = {"litellm_params": {}} # No api_base
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"collectionName": collection_name},
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"
+ ), patch.object(
+ litellm, "vector_store_index_registry"
+ ) as mock_index_registry, patch.object(
+ litellm, "vector_store_registry"
+ ) as mock_vector_registry:
+ mock_provider_config = MagicMock()
+ mock_provider_config.get_auth_credentials.return_value = {"headers": {}}
+ mock_provider_config.get_complete_url.return_value = None
+ mock_get_config.return_value = mock_provider_config
+
+ mock_index_registry.is_vector_store_index.return_value = True
+ mock_index_registry.get_vector_store_index_by_name.return_value = (
+ mock_index_object
+ )
+ mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = (
+ mock_vector_store
+ )
+
+ with pytest.raises(Exception) as exc_info:
+ await milvus_proxy_route(
+ endpoint="vectors/search",
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ assert (
+ f"api_base not found in vector store configuration for {vector_store_name}"
+ in str(exc_info.value)
+ )
+
+ @pytest.mark.asyncio
+ async def test_milvus_proxy_route_endpoint_without_leading_slash(self):
+ """
+ Test that endpoint without leading slash is handled correctly
+ """
+ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
+ milvus_proxy_route,
+ )
+
+ collection_name = "test-collection"
+ vector_store_name = "milvus-store-1"
+ vector_store_index = "collection_123"
+ api_base = "http://localhost:19530"
+
+ mock_request = MagicMock(spec=Request)
+ mock_response = MagicMock(spec=Response)
+ mock_user_api_key_dict = MagicMock()
+
+ mock_index_object = MagicMock()
+ mock_index_object.litellm_params.vector_store_name = vector_store_name
+ mock_index_object.litellm_params.vector_store_index = vector_store_index
+
+ mock_vector_store = {"litellm_params": {"api_base": api_base}}
+
+ with patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
+ return_value={"collectionName": collection_name},
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
+ ) as mock_get_config, patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"
+ ), patch(
+ "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route"
+ ) as mock_create_route, patch.object(
+ litellm, "vector_store_index_registry"
+ ) as mock_index_registry, patch.object(
+ litellm, "vector_store_registry"
+ ) as mock_vector_registry:
+ mock_provider_config = MagicMock()
+ mock_provider_config.get_auth_credentials.return_value = {"headers": {}}
+ mock_provider_config.get_complete_url.return_value = api_base
+ mock_get_config.return_value = mock_provider_config
+
+ mock_index_registry.is_vector_store_index.return_value = True
+ mock_index_registry.get_vector_store_index_by_name.return_value = (
+ mock_index_object
+ )
+ mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = (
+ mock_vector_store
+ )
+
+ mock_endpoint_func = AsyncMock(return_value={"status": "success"})
+ mock_create_route.return_value = mock_endpoint_func
+
+ # Call with endpoint without leading slash
+ await milvus_proxy_route(
+ endpoint="vectors/search", # No leading slash
+ request=mock_request,
+ fastapi_response=mock_response,
+ user_api_key_dict=mock_user_api_key_dict,
+ )
+
+ # Verify that the target URL has correct path
+ create_route_args = mock_create_route.call_args[1]
+ assert "/vectors/search" in create_route_args["target"]
diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx
new file mode 100644
index 00000000000..a74c6d283ca
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx
@@ -0,0 +1,101 @@
+import { render, renderHook } from "@testing-library/react";
+import { describe, it, vi, expect } from "vitest";
+import { Form } from "antd";
+import AddModelTab from "./add_model_tab";
+import { Providers } from "../provider_info_helpers";
+import type { Team } from "../key_team_helpers/key_list";
+import type { CredentialItem } from "../networking";
+import type { UploadProps } from "antd/es/upload";
+
+// Mock the networking module
+vi.mock("../networking", async () => {
+ const actual = await vi.importActual("../networking");
+ return {
+ ...actual,
+ getGuardrailsList: vi.fn().mockResolvedValue({
+ guardrails: [{ guardrail_name: "test-guardrail-1" }, { guardrail_name: "test-guardrail-2" }],
+ }),
+ tagListCall: vi.fn().mockResolvedValue({}),
+ modelAvailableCall: vi.fn().mockResolvedValue({
+ data: [{ id: "model-group-1" }, { id: "model-group-2" }],
+ }),
+ };
+});
+
+describe("Add Model Tab", () => {
+ it("should render", () => {
+ // Create a form instance using renderHook
+ const { result } = renderHook(() => Form.useForm());
+ const [form] = result.current;
+
+ // Mock functions
+ const handleOk = vi.fn();
+ const setSelectedProvider = vi.fn();
+ const setProviderModelsFn = vi.fn();
+ const getPlaceholder = vi.fn((provider: Providers) => `Enter ${provider} model name`);
+ const setShowAdvancedSettings = vi.fn();
+
+ // Mock data
+ const selectedProvider = Providers.OpenAI;
+ const providerModels = ["gpt-4", "gpt-3.5-turbo"];
+ const showAdvancedSettings = false;
+
+ const teams: Team[] = [
+ {
+ team_id: "team-1",
+ team_alias: "Test Team",
+ models: ["gpt-4"],
+ max_budget: 100,
+ budget_duration: "monthly",
+ tpm_limit: null,
+ rpm_limit: null,
+ organization_id: "org-1",
+ created_at: "2024-01-01T00:00:00Z",
+ keys: [],
+ members_with_roles: [],
+ },
+ ];
+
+ const credentials: CredentialItem[] = [
+ {
+ credential_name: "test-credential",
+ credential_values: {},
+ credential_info: {
+ custom_llm_provider: "openai",
+ description: "Test credential",
+ },
+ },
+ ];
+
+ const uploadProps: UploadProps = {
+ beforeUpload: () => false,
+ showUploadList: false,
+ };
+
+ const accessToken = "test-access-token";
+ const userRole = "Admin";
+ const premiumUser = true;
+
+ const { getByRole } = render(
+ ,
+ );
+ // Check for the heading specifically
+ expect(getByRole("heading", { name: "Add Model" })).toBeInTheDocument();
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx
index faf14c96a89..678ba741ca8 100644
--- a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx
+++ b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx
@@ -240,7 +240,7 @@ const AddModelTab: React.FC = ({
-
+
= ({
/>
-
-
@@ -277,13 +271,18 @@ const AddModelTab: React.FC = ({
console.log("🔑 Credential Name Changed:", credentialName);
// Only show provider specific fields if no credentials selected
if (!credentialName) {
- return ;
+ return (
+ <>
+
+
+ >
+ );
}
- return (
-
- Using existing credentials - no additional provider fields needed
-
- );
+ return null;
}}
diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_info.test.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_info.test.tsx
index 301045103bc..bd2b45d6f34 100644
--- a/ui/litellm-dashboard/src/components/guardrails/guardrail_info.test.tsx
+++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_info.test.tsx
@@ -110,4 +110,41 @@ describe("Guardrail Info", () => {
});
}
});
+
+ it("should render the guardrail info", async () => {
+ // Mock the network responses
+ vi.mocked(networking.getGuardrailInfo).mockResolvedValue({
+ guardrail_id: "123",
+ guardrail_name: "Test Guardrail",
+ litellm_params: {
+ guardrail: "presidio",
+ mode: "pre_call",
+ default_on: true,
+ pii_entities_config: {
+ PERSON: "MASK",
+ EMAIL: "REDACT",
+ },
+ },
+ created_at: "2024-01-01T00:00:00Z",
+ updated_at: "2024-01-01T00:00:00Z",
+ guardrail_definition_location: "database",
+ });
+
+ vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({
+ supported_entities: ["PERSON", "EMAIL"],
+ supported_actions: ["MASK", "REDACT"],
+ pii_entity_categories: [],
+ supported_modes: ["pre_call", "post_call"],
+ });
+
+ vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
+
+ const { getByText } = render(
+
{}} accessToken="123" isAdmin={true} />,
+ );
+
+ await waitFor(() => {
+ expect(getByText("PII Entity Configuration")).toBeInTheDocument();
+ });
+ });
});
diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx
index d783108de3c..d314cbc5ebe 100644
--- a/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx
+++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_info.tsx
@@ -14,7 +14,7 @@ import {
TextInput,
} from "@tremor/react";
import { Button, Form, Input, Select, Divider, Tooltip } from "antd";
-import { InfoCircleOutlined } from "@ant-design/icons";
+import { InfoCircleOutlined, EyeInvisibleOutlined, StopOutlined } from "@ant-design/icons";
import {
getGuardrailInfo,
updateGuardrailCall,
@@ -414,21 +414,35 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,
)}
- {guardrailData.guardrail_info && Object.keys(guardrailData.guardrail_info).length > 0 && (
-
- Guardrail Info
-
- {Object.entries(guardrailData.guardrail_info).map(([key, value]) => (
-
-
{key}
-
- {typeof value === "object" ? JSON.stringify(value, null, 2) : String(value)}
-
+ {guardrailData.litellm_params?.pii_entities_config &&
+ Object.keys(guardrailData.litellm_params.pii_entities_config).length > 0 && (
+
+ PII Entity Configuration
+
+
+ Entity Type
+ Configuration
- ))}
-
-
- )}
+
+ {Object.entries(guardrailData.litellm_params?.pii_entities_config).map(([key, value]) => (
+
+ {key}
+
+
+ {value === "MASK" ? : }
+ {String(value)}
+
+
+
+ ))}
+
+
+
+ )}
{/* Settings Panel (only for admins) */}
diff --git a/ui/litellm-dashboard/src/components/team/team_info.test.tsx b/ui/litellm-dashboard/src/components/team/team_info.test.tsx
new file mode 100644
index 00000000000..526f0972d98
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/team/team_info.test.tsx
@@ -0,0 +1,81 @@
+import { afterEach, describe, expect, it, vi } from "vitest";
+import TeamInfoView from "./team_info";
+import { render, waitFor } from "@testing-library/react";
+import * as networking from "@/components/networking";
+
+// Mock the networking module
+vi.mock("@/components/networking", () => ({
+ teamInfoCall: vi.fn(),
+ teamMemberDeleteCall: vi.fn(),
+ teamMemberAddCall: vi.fn(),
+ teamMemberUpdateCall: vi.fn(),
+ teamUpdateCall: vi.fn(),
+ getGuardrailsList: vi.fn(),
+ fetchMCPAccessGroups: vi.fn(),
+ getTeamPermissionsCall: vi.fn(),
+}));
+
+describe("TeamInfoView", () => {
+ afterEach(() => {
+ vi.clearAllMocks();
+ });
+
+ it("should render", async () => {
+ // Mock the team info response
+ vi.mocked(networking.teamInfoCall).mockResolvedValue({
+ team_id: "123",
+ team_info: {
+ team_alias: "Test Team",
+ team_id: "123",
+ organization_id: null,
+ admins: ["admin@test.com"],
+ members: ["user1@test.com", "user2@test.com"],
+ members_with_roles: [
+ {
+ user_id: "user1@test.com",
+ user_email: "user1@test.com",
+ role: "member",
+ spend: 0,
+ budget_id: "budget1",
+ },
+ ],
+ metadata: {},
+ tpm_limit: null,
+ rpm_limit: null,
+ max_budget: null,
+ budget_duration: null,
+ models: [],
+ blocked: false,
+ spend: 0,
+ max_parallel_requests: null,
+ budget_reset_at: null,
+ model_id: null,
+ litellm_model_table: null,
+ created_at: "2024-01-01T00:00:00Z",
+ team_member_budget_table: null,
+ },
+ keys: [],
+ team_memberships: [],
+ });
+
+ vi.mocked(networking.getGuardrailsList).mockResolvedValue([]);
+ vi.mocked(networking.fetchMCPAccessGroups).mockResolvedValue([]);
+
+ const { getByText } = render(
+
{}}
+ onClose={() => {}}
+ accessToken="123"
+ is_team_admin={true}
+ is_proxy_admin={true}
+ userModels={[]}
+ editTeam={false}
+ premiumUser={false}
+ />,
+ );
+ await waitFor(() => {
+ expect(getByText("User ID")).toBeInTheDocument();
+ });
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/team/team_info.tsx b/ui/litellm-dashboard/src/components/team/team_info.tsx
index 2aec5a3b767..34446c51664 100644
--- a/ui/litellm-dashboard/src/components/team/team_info.tsx
+++ b/ui/litellm-dashboard/src/components/team/team_info.tsx
@@ -25,7 +25,7 @@ import {
teamUpdateCall,
getGuardrailsList,
} from "@/components/networking";
-import { Button, Form, Input, Select, message, Tooltip } from "antd";
+import { Button, Form, Input, Select, message, Modal, Tooltip } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { ArrowLeftIcon } from "@heroicons/react/outline";
import MemberModal from "./edit_membership";
@@ -138,6 +138,8 @@ const TeamInfoView: React.FC = ({
const [mcpAccessGroupsLoaded, setMcpAccessGroupsLoaded] = useState(false);
const [copiedStates, setCopiedStates] = useState>({});
const [guardrailsList, setGuardrailsList] = useState([]);
+ const [memberToDelete, setMemberToDelete] = useState(null);
+ const [isDeleting, setIsDeleting] = useState(false);
console.log("userModels in team info", userModels);
@@ -268,13 +270,16 @@ const TeamInfoView: React.FC = ({
}
};
- const handleMemberDelete = async (member: Member) => {
- try {
- if (accessToken == null) {
- return;
- }
+ const handleMemberDelete = (member: Member) => {
+ setMemberToDelete(member);
+ };
- await teamMemberDeleteCall(accessToken, teamId, member);
+ const handleDeleteConfirm = async () => {
+ if (!memberToDelete || !accessToken) return;
+
+ setIsDeleting(true);
+ try {
+ await teamMemberDeleteCall(accessToken, teamId, memberToDelete);
NotificationsManager.success("Team member removed successfully");
@@ -287,9 +292,16 @@ const TeamInfoView: React.FC = ({
} catch (error) {
NotificationsManager.fromBackend("Failed to remove team member");
console.error("Error removing team member:", error);
+ } finally {
+ setIsDeleting(false);
+ setMemberToDelete(null);
}
};
+ const handleDeleteCancel = () => {
+ setMemberToDelete(null);
+ };
+
const handleTeamUpdate = async (values: any) => {
try {
if (!accessToken) return;
@@ -885,6 +897,30 @@ const TeamInfoView: React.FC = ({
onSubmit={handleMemberCreate}
accessToken={accessToken}
/>
+
+ {/* Delete Member Confirmation Modal */}
+ {memberToDelete && (
+
+ Are you sure you want to remove this member from the team?
+
+ User ID: {memberToDelete.user_id}
+
+ {memberToDelete.user_email && (
+
+ Email: {memberToDelete.user_email}
+
+ )}
+ This action cannot be undone.
+
+ )}
);
};
diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx
index 4145283db41..8aa37f007de 100644
--- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx
@@ -36,30 +36,150 @@ interface GuardrailInformation {
}
interface GuardrailViewerProps {
- data: GuardrailInformation;
+ data: GuardrailInformation | GuardrailInformation[];
}
+interface GuardrailDetailsProps {
+ entry: GuardrailInformation;
+ index: number;
+ total: number;
+}
+
+const formatTime = (timestamp: number) => {
+ const date = new Date(timestamp * 1000);
+ return date.toLocaleString();
+};
+
+const GuardrailDetails = ({ entry, index, total }: GuardrailDetailsProps) => {
+ const guardrailProvider = entry.guardrail_provider ?? "presidio";
+ const statusLabel = entry.guardrail_status ?? "unknown";
+ const isSuccess = statusLabel.toLowerCase() === "success";
+ const maskedEntityCount = entry.masked_entity_count || {};
+ const totalMaskedEntities = Object.values(maskedEntityCount).reduce(
+ (sum, count) => sum + (typeof count === "number" ? count : 0),
+ 0,
+ );
+
+ const guardrailResponse = entry.guardrail_response;
+ const presidioEntities = Array.isArray(guardrailResponse) ? guardrailResponse : [];
+ const bedrockResponse =
+ guardrailProvider === "bedrock" &&
+ guardrailResponse !== null &&
+ typeof guardrailResponse === "object" &&
+ !Array.isArray(guardrailResponse)
+ ? (guardrailResponse as BedrockGuardrailResponse)
+ : undefined;
+
+ return (
+
+ {total > 1 && (
+
+
+ Guardrail #{index + 1}
+ {entry.guardrail_name}
+
+
+ {guardrailProvider}
+
+
+ )}
+
+
+
+
+ Guardrail Name:
+ {entry.guardrail_name}
+
+
+ Mode:
+ {entry.guardrail_mode}
+
+
+ Status:
+
+
+ {statusLabel}
+
+
+
+
+
+
+
+ Start Time:
+ {formatTime(entry.start_time)}
+
+
+ End Time:
+ {formatTime(entry.end_time)}
+
+
+ Duration:
+ {entry.duration.toFixed(4)}s
+
+
+
+
+ {totalMaskedEntities > 0 && (
+
+
Masked Entity Summary
+
+ {Object.entries(maskedEntityCount).map(([entityType, count]) => (
+
+ {entityType}: {count}
+
+ ))}
+
+
+ )}
+
+ {guardrailProvider === "presidio" && presidioEntities.length > 0 && (
+
+ )}
+
+ {guardrailProvider === "bedrock" && bedrockResponse && (
+
+
+
+ )}
+
+ );
+};
+
const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
+ const guardrailEntries = Array.isArray(data)
+ ? data.filter((entry): entry is GuardrailInformation => Boolean(entry))
+ : data
+ ? [data]
+ : [];
+
+ if (guardrailEntries.length === 0) {
+ return null;
+ }
+
const [sectionExpanded, setSectionExpanded] = useState(true);
- // Default to presidio for backwards compatibility
- const guardrailProvider = data.guardrail_provider ?? "presidio";
+ const primaryName = guardrailEntries.length === 1 ? guardrailEntries[0].guardrail_name : `${guardrailEntries.length} guardrails`;
+ const statuses = Array.from(new Set(guardrailEntries.map((entry) => entry.guardrail_status)));
+ const allSucceeded = statuses.every((status) => (status ?? "").toLowerCase() === "success");
+ const aggregatedStatus = allSucceeded ? "success" : "failure";
+ const totalMaskedEntities = guardrailEntries.reduce((sum, entry) => {
+ return (
+ sum +
+ Object.values(entry.masked_entity_count || {}).reduce((acc, count) => acc + (typeof count === "number" ? count : 0), 0)
+ );
+ }, 0);
- if (!data) return null;
-
- const isSuccess = typeof data.guardrail_status === "string" && data.guardrail_status.toLowerCase() === "success";
-
- const tooltipTitle = isSuccess ? null : "Guardrail failed to run.";
-
- // Calculate total masked entities
- const totalMaskedEntities = data.masked_entity_count
- ? Object.values(data.masked_entity_count).reduce((sum, count) => sum + count, 0)
- : 0;
-
- const formatTime = (timestamp: number): string => {
- const date = new Date(timestamp * 1000);
- return date.toLocaleString();
- };
+ const tooltipTitle = allSucceeded ? null : "Guardrail failed to run.";
return (
@@ -67,9 +187,9 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
className="flex justify-between items-center p-4 border-b cursor-pointer hover:bg-gray-50"
onClick={() => setSectionExpanded(!sectionExpanded)}
>
-
+
{
Guardrail Information
- {/* Header status chip with tooltip */}
- {data.guardrail_status}
+ {aggregatedStatus}
+ {primaryName}
+
{totalMaskedEntities > 0 && (
-
+
{totalMaskedEntities} masked {totalMaskedEntities === 1 ? "entity" : "entities"}
)}
@@ -99,76 +220,15 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
{sectionExpanded && (
-
-
-
-
-
- Guardrail Name:
- {data.guardrail_name}
-
-
- Mode:
- {data.guardrail_mode}
-
-
- Status:
-
-
- {data.guardrail_status}
-
-
-
-
-
-
-
- Start Time:
- {formatTime(data.start_time)}
-
-
- End Time:
- {formatTime(data.end_time)}
-
-
- Duration:
- {data.duration.toFixed(4)}s
-
-
-
-
- {/* Masked Entity Summary */}
- {data.masked_entity_count && Object.keys(data.masked_entity_count).length > 0 && (
-
-
Masked Entity Summary
-
- {Object.entries(data.masked_entity_count).map(([entityType, count]) => (
-
- {entityType}: {count}
-
- ))}
-
-
- )}
-
-
- {/* Provider-specific Detected Entities */}
- {guardrailProvider === "presidio" && (data.guardrail_response as GuardrailEntity[])?.length > 0 && (
-
- )}
-
- {guardrailProvider === "bedrock" && data.guardrail_response && (
-
-
-
- )}
+
+ {guardrailEntries.map((entry, index) => (
+
+ ))}
)}
diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx
index bfc57a84f7d..fb833bdfca3 100644
--- a/ui/litellm-dashboard/src/components/view_logs/index.tsx
+++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx
@@ -790,20 +790,34 @@ export function RequestViewer({ row }: { row: Row
}) {
metadata.vector_store_request_metadata.length > 0;
// Extract guardrail information from metadata if available
- const hasGuardrailData = row.original.metadata && row.original.metadata.guardrail_information;
+ const guardrailInfo = row.original.metadata?.guardrail_information;
+ const guardrailEntries = Array.isArray(guardrailInfo)
+ ? guardrailInfo
+ : guardrailInfo
+ ? [guardrailInfo]
+ : [];
+ const hasGuardrailData = guardrailEntries.length > 0;
// Calculate total masked entities if guardrail data exists
- const getTotalMaskedEntities = (): number => {
- if (!hasGuardrailData || !row.original.metadata?.guardrail_information.masked_entity_count) {
- return 0;
+ const totalMaskedEntities = guardrailEntries.reduce((sum, entry) => {
+ const maskedCounts = entry?.masked_entity_count;
+ if (!maskedCounts) {
+ return sum;
}
- return Object.values(row.original.metadata.guardrail_information.masked_entity_count).reduce(
- (sum: number, count: any) => sum + (typeof count === "number" ? count : 0),
- 0,
+ return (
+ sum +
+ Object.values(maskedCounts).reduce(
+ (acc, count) => (typeof count === "number" ? acc + count : acc),
+ 0,
+ )
);
- };
+ }, 0);
- const totalMaskedEntities = getTotalMaskedEntities();
+ const primaryGuardrailLabel = guardrailEntries.length === 1
+ ? guardrailEntries[0]?.guardrail_name ?? "-"
+ : guardrailEntries.length > 1
+ ? `${guardrailEntries.length} guardrails`
+ : "-";
return (
@@ -850,7 +864,7 @@ export function RequestViewer({ row }: { row: Row
}) {
Guardrail:
- {row.original.metadata!.guardrail_information.guardrail_name}
+ {primaryGuardrailLabel}
{totalMaskedEntities > 0 && (
{totalMaskedEntities} masked
@@ -934,7 +948,7 @@ export function RequestViewer({ row }: { row: Row }) {
{/* Guardrail Data - Show only if present */}
- {hasGuardrailData &&
}
+ {hasGuardrailData &&
}
{/* Vector Store Request Data - Show only if present */}
{hasVectorStoreData &&
}