mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'main' of https://github.com/BerriAI/litellm into litellm_non_root_model_hub
This commit is contained in:
commit
b6f2b9a8be
76 changed files with 4609 additions and 916 deletions
|
|
@ -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).
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
110
litellm/integrations/langfuse/langfuse_otel_attributes.py
Normal file
110
litellm/integrations/langfuse/langfuse_otel_attributes.py
Normal 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),
|
||||
)
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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 = ""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
1
litellm/llms/perplexity/chat/__init__.py
Normal file
1
litellm/llms/perplexity/chat/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""Perplexity chat completion transformations."""
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
# Count tokens handler for Vertex AI Partner Models (Anthropic, Mistral, etc.)
|
||||
|
|
@ -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",
|
||||
}
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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", {})
|
||||
|
|
@ -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", {})
|
||||
|
|
@ -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")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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) */}
|
||||
|
|
|
|||
81
ui/litellm-dashboard/src/components/team/team_info.test.tsx
Normal file
81
ui/litellm-dashboard/src/components/team/team_info.test.tsx
Normal 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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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} />}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue