Merge branch 'main' of https://github.com/BerriAI/litellm into litellm_non_root_model_hub

This commit is contained in:
yuneng-jiang 2025-11-03 11:31:31 -08:00
commit b6f2b9a8be
76 changed files with 4609 additions and 916 deletions

View file

@ -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).

View file

@ -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:

View file

@ -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

View file

@ -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 |

View file

@ -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:

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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.

View file

@ -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(

View file

@ -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

View file

@ -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),
)

View file

@ -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):

View file

@ -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)

View file

@ -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 = ""

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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],

View file

@ -0,0 +1 @@
"""Perplexity chat completion transformations."""

View file

@ -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)
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

View file

@ -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
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

View file

@ -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(

View file

@ -0,0 +1 @@
# Count tokens handler for Vertex AI Partner Models (Anthropic, Mistral, etc.)

View file

@ -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",
}

View file

@ -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))

View file

@ -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"

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

View file

@ -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"
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

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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):

View file

@ -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]

View file

@ -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"}
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"}

View file

@ -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
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"

View file

@ -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}")

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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.

View file

@ -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:

View file

@ -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"

View file

@ -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__])

View file

@ -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)

View file

@ -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"}
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}"

View file

@ -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

View file

@ -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
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

View file

@ -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
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", {})

View file

@ -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)
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", {})

View file

@ -856,4 +856,121 @@ def test_get_token_url():
assert "v1beta1" not in url
assert "/v1/" in url
pass
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")

View file

@ -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

View file

@ -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"]

View file

@ -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(
<AddModelTab
form={form}
handleOk={handleOk}
selectedProvider={selectedProvider}
setSelectedProvider={setSelectedProvider}
providerModels={providerModels}
setProviderModelsFn={setProviderModelsFn}
getPlaceholder={getPlaceholder}
uploadProps={uploadProps}
showAdvancedSettings={showAdvancedSettings}
setShowAdvancedSettings={setShowAdvancedSettings}
teams={teams}
credentials={credentials}
accessToken={accessToken}
userRole={userRole}
premiumUser={premiumUser}
/>,
);
// Check for the heading specifically
expect(getByRole("heading", { name: "Add Model" })).toBeInTheDocument();
});
});

View file

@ -240,7 +240,7 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
</Typography.Text>
</div>
<Form.Item label="Existing Credentials" name="litellm_credential_name">
<Form.Item label="Existing Credentials" name="litellm_credential_name" initialValue={null}>
<AntdSelect
showSearch
placeholder="Select or search for existing credentials"
@ -259,12 +259,6 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
/>
</Form.Item>
<div className="flex items-center my-4">
<div className="flex-grow border-t border-gray-200"></div>
<span className="px-4 text-gray-500 text-sm">OR</span>
<div className="flex-grow border-t border-gray-200"></div>
</div>
<Form.Item
noStyle
shouldUpdate={(prevValues, currentValues) =>
@ -277,13 +271,18 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
console.log("🔑 Credential Name Changed:", credentialName);
// Only show provider specific fields if no credentials selected
if (!credentialName) {
return <ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />;
return (
<>
<div className="flex items-center my-4">
<div className="flex-grow border-t border-gray-200"></div>
<span className="px-4 text-gray-500 text-sm">OR</span>
<div className="flex-grow border-t border-gray-200"></div>
</div>
<ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />
</>
);
}
return (
<div className="text-gray-500 text-sm text-center">
Using existing credentials - no additional provider fields needed
</div>
);
return null;
}}
</Form.Item>
<div className="flex items-center my-4">

View file

@ -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(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
await waitFor(() => {
expect(getByText("PII Entity Configuration")).toBeInTheDocument();
});
});
});

View file

@ -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<GuardrailInfoProps> = ({ guardrailId, onClose,
</Card>
)}
{guardrailData.guardrail_info && Object.keys(guardrailData.guardrail_info).length > 0 && (
<Card className="mt-6">
<Text>Guardrail Info</Text>
<div className="mt-2 space-y-2">
{Object.entries(guardrailData.guardrail_info).map(([key, value]) => (
<div key={key} className="flex">
<Text className="font-medium w-1/3">{key}</Text>
<Text className="w-2/3">
{typeof value === "object" ? JSON.stringify(value, null, 2) : String(value)}
</Text>
{guardrailData.litellm_params?.pii_entities_config &&
Object.keys(guardrailData.litellm_params.pii_entities_config).length > 0 && (
<Card className="mt-6">
<Text className="mb-4 text-lg font-semibold">PII Entity Configuration</Text>
<div className="border rounded-lg overflow-hidden shadow-sm">
<div className="bg-gray-50 px-5 py-3 border-b flex">
<Text className="flex-1 font-semibold text-gray-700">Entity Type</Text>
<Text className="flex-1 font-semibold text-gray-700">Configuration</Text>
</div>
))}
</div>
</Card>
)}
<div className="max-h-[400px] overflow-y-auto">
{Object.entries(guardrailData.litellm_params?.pii_entities_config).map(([key, value]) => (
<div key={key} className="px-5 py-3 flex border-b hover:bg-gray-50 transition-colors">
<Text className="flex-1 font-medium text-gray-900">{key}</Text>
<Text className="flex-1">
<span
className={`inline-flex items-center gap-1.5 ${
value === "MASK" ? "text-blue-600" : "text-red-600"
}`}
>
{value === "MASK" ? <EyeInvisibleOutlined /> : <StopOutlined />}
{String(value)}
</span>
</Text>
</div>
))}
</div>
</div>
</Card>
)}
</TabPanel>
{/* Settings Panel (only for admins) */}

View file

@ -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(
<TeamInfoView
teamId="123"
onUpdate={() => {}}
onClose={() => {}}
accessToken="123"
is_team_admin={true}
is_proxy_admin={true}
userModels={[]}
editTeam={false}
premiumUser={false}
/>,
);
await waitFor(() => {
expect(getByText("User ID")).toBeInTheDocument();
});
});
});

View file

@ -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<TeamInfoProps> = ({
const [mcpAccessGroupsLoaded, setMcpAccessGroupsLoaded] = useState(false);
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({});
const [guardrailsList, setGuardrailsList] = useState<string[]>([]);
const [memberToDelete, setMemberToDelete] = useState<Member | null>(null);
const [isDeleting, setIsDeleting] = useState(false);
console.log("userModels in team info", userModels);
@ -268,13 +270,16 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
}
};
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<TeamInfoProps> = ({
} 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<TeamInfoProps> = ({
onSubmit={handleMemberCreate}
accessToken={accessToken}
/>
{/* Delete Member Confirmation Modal */}
{memberToDelete && (
<Modal
title="Delete Team Member"
open={memberToDelete !== null}
onOk={handleDeleteConfirm}
onCancel={handleDeleteCancel}
confirmLoading={isDeleting}
okText={isDeleting ? "Deleting..." : "Delete"}
okButtonProps={{ danger: true }}
>
<p>Are you sure you want to remove this member from the team?</p>
<p className="mt-2">
<strong>User ID:</strong> {memberToDelete.user_id}
</p>
{memberToDelete.user_email && (
<p>
<strong>Email:</strong> {memberToDelete.user_email}
</p>
)}
<p className="mt-2 text-red-600">This action cannot be undone.</p>
</Modal>
)}
</div>
);
};

View file

@ -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 (
<div className="bg-white rounded-lg border border-gray-200 p-4">
{total > 1 && (
<div className="flex items-center justify-between mb-4">
<h4 className="text-base font-semibold">
Guardrail #{index + 1}
<span className="ml-2 font-mono text-sm text-gray-600">{entry.guardrail_name}</span>
</h4>
<span className="px-2 py-0.5 bg-gray-100 text-gray-600 rounded-md text-xs capitalize">
{guardrailProvider}
</span>
</div>
)}
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Guardrail Name:</span>
<span className="font-mono break-words">{entry.guardrail_name}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Mode:</span>
<span className="font-mono break-words">{entry.guardrail_mode}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Status:</span>
<Tooltip title={isSuccess ? null : "Guardrail failed to run."} placement="top" arrow destroyTooltipOnHide>
<span
className={`px-2 py-1 rounded-md text-xs font-medium inline-block ${
isSuccess ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800 cursor-help"
}`}
>
{statusLabel}
</span>
</Tooltip>
</div>
</div>
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Start Time:</span>
<span>{formatTime(entry.start_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">End Time:</span>
<span>{formatTime(entry.end_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Duration:</span>
<span>{entry.duration.toFixed(4)}s</span>
</div>
</div>
</div>
{totalMaskedEntities > 0 && (
<div className="mt-4 pt-4 border-t">
<h5 className="font-medium mb-2">Masked Entity Summary</h5>
<div className="flex flex-wrap gap-2">
{Object.entries(maskedEntityCount).map(([entityType, count]) => (
<span
key={entityType}
className="px-3 py-1.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium"
>
{entityType}: {count}
</span>
))}
</div>
</div>
)}
{guardrailProvider === "presidio" && presidioEntities.length > 0 && (
<div className="mt-4">
<PresidioDetectedEntities entities={presidioEntities} />
</div>
)}
{guardrailProvider === "bedrock" && bedrockResponse && (
<div className="mt-4">
<BedrockGuardrailDetails response={bedrockResponse} />
</div>
)}
</div>
);
};
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 (
<div className="bg-white rounded-lg shadow mb-6">
@ -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)}
>
<div className="flex items-center">
<div className="flex items-center gap-2">
<svg
className={`w-5 h-5 mr-2 text-gray-600 transition-transform ${sectionExpanded ? "transform rotate-90" : ""}`}
className={`w-5 h-5 text-gray-600 transition-transform ${sectionExpanded ? "transform rotate-90" : ""}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
@ -78,19 +198,20 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
</svg>
<h3 className="text-lg font-medium">Guardrail Information</h3>
{/* Header status chip with tooltip */}
<Tooltip title={tooltipTitle} placement="top" arrow destroyTooltipOnHide>
<span
className={`ml-3 px-2 py-1 rounded-md text-xs font-medium inline-block ${
isSuccess ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800 cursor-help"
className={`ml-2 px-2 py-1 rounded-md text-xs font-medium inline-block ${
allSucceeded ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800 cursor-help"
}`}
>
{data.guardrail_status}
{aggregatedStatus}
</span>
</Tooltip>
<span className="ml-2 font-mono text-sm text-gray-600">{primaryName}</span>
{totalMaskedEntities > 0 && (
<span className="ml-3 px-2 py-1 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
<span className="ml-2 px-2 py-1 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{totalMaskedEntities} masked {totalMaskedEntities === 1 ? "entity" : "entities"}
</span>
)}
@ -99,76 +220,15 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
</div>
{sectionExpanded && (
<div className="p-4">
<div className="bg-white rounded-lg border p-4 mb-4">
<div className="grid grid-cols-2 gap-4">
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Guardrail Name:</span>
<span className="font-mono">{data.guardrail_name}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Mode:</span>
<span className="font-mono">{data.guardrail_mode}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Status:</span>
<Tooltip title={tooltipTitle} placement="top" arrow destroyTooltipOnHide>
<span
className={`px-2 py-1 rounded-md text-xs font-medium inline-block ${
isSuccess ? "bg-green-100 text-green-800" : "bg-red-100 text-red-800 cursor-help"
}`}
>
{data.guardrail_status}
</span>
</Tooltip>
</div>
</div>
<div className="space-y-2">
<div className="flex">
<span className="font-medium w-1/3">Start Time:</span>
<span>{formatTime(data.start_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">End Time:</span>
<span>{formatTime(data.end_time)}</span>
</div>
<div className="flex">
<span className="font-medium w-1/3">Duration:</span>
<span>{data.duration.toFixed(4)}s</span>
</div>
</div>
</div>
{/* Masked Entity Summary */}
{data.masked_entity_count && Object.keys(data.masked_entity_count).length > 0 && (
<div className="mt-4 pt-4 border-t">
<h4 className="font-medium mb-2">Masked Entity Summary</h4>
<div className="flex flex-wrap gap-2">
{Object.entries(data.masked_entity_count).map(([entityType, count]) => (
<span
key={entityType}
className="px-3 py-1.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium"
>
{entityType}: {count}
</span>
))}
</div>
</div>
)}
</div>
{/* Provider-specific Detected Entities */}
{guardrailProvider === "presidio" && (data.guardrail_response as GuardrailEntity[])?.length > 0 && (
<PresidioDetectedEntities entities={data.guardrail_response as GuardrailEntity[]} />
)}
{guardrailProvider === "bedrock" && data.guardrail_response && (
<div className="mt-4">
<BedrockGuardrailDetails response={data.guardrail_response as BedrockGuardrailResponse} />
</div>
)}
<div className="p-4 space-y-6">
{guardrailEntries.map((entry, index) => (
<GuardrailDetails
key={`${entry.guardrail_name ?? "guardrail"}-${index}`}
entry={entry}
index={index}
total={guardrailEntries.length}
/>
))}
</div>
)}
</div>

View file

@ -790,20 +790,34 @@ export function RequestViewer({ row }: { row: Row<LogEntry> }) {
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<number>(
(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 (
<div className="p-6 bg-gray-50 space-y-6 w-full max-w-full overflow-hidden box-border">
@ -850,7 +864,7 @@ export function RequestViewer({ row }: { row: Row<LogEntry> }) {
<div className="flex">
<span className="font-medium w-1/3">Guardrail:</span>
<div>
<span className="font-mono">{row.original.metadata!.guardrail_information.guardrail_name}</span>
<span className="font-mono">{primaryGuardrailLabel}</span>
{totalMaskedEntities > 0 && (
<span className="ml-2 px-2 py-0.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{totalMaskedEntities} masked
@ -934,7 +948,7 @@ export function RequestViewer({ row }: { row: Row<LogEntry> }) {
</div>
{/* Guardrail Data - Show only if present */}
{hasGuardrailData && <GuardrailViewer data={row.original.metadata!.guardrail_information} />}
{hasGuardrailData && <GuardrailViewer data={guardrailInfo} />}
{/* Vector Store Request Data - Show only if present */}
{hasVectorStoreData && <VectorStoreViewer data={metadata.vector_store_request_metadata} />}