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

Are you sure you want to remove this member from the team?

+

+ User ID: {memberToDelete.user_id} +

+ {memberToDelete.user_email && ( +

+ Email: {memberToDelete.user_email} +

+ )} +

This action cannot be undone.

+
+ )}
); }; diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx index 4145283db41..8aa37f007de 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx @@ -36,30 +36,150 @@ interface GuardrailInformation { } interface GuardrailViewerProps { - data: GuardrailInformation; + data: GuardrailInformation | GuardrailInformation[]; } +interface GuardrailDetailsProps { + entry: GuardrailInformation; + index: number; + total: number; +} + +const formatTime = (timestamp: number) => { + const date = new Date(timestamp * 1000); + return date.toLocaleString(); +}; + +const GuardrailDetails = ({ entry, index, total }: GuardrailDetailsProps) => { + const guardrailProvider = entry.guardrail_provider ?? "presidio"; + const statusLabel = entry.guardrail_status ?? "unknown"; + const isSuccess = statusLabel.toLowerCase() === "success"; + const maskedEntityCount = entry.masked_entity_count || {}; + const totalMaskedEntities = Object.values(maskedEntityCount).reduce( + (sum, count) => sum + (typeof count === "number" ? count : 0), + 0, + ); + + const guardrailResponse = entry.guardrail_response; + const presidioEntities = Array.isArray(guardrailResponse) ? guardrailResponse : []; + const bedrockResponse = + guardrailProvider === "bedrock" && + guardrailResponse !== null && + typeof guardrailResponse === "object" && + !Array.isArray(guardrailResponse) + ? (guardrailResponse as BedrockGuardrailResponse) + : undefined; + + return ( +
+ {total > 1 && ( +
+

+ Guardrail #{index + 1} + {entry.guardrail_name} +

+ + {guardrailProvider} + +
+ )} + +
+
+
+ Guardrail Name: + {entry.guardrail_name} +
+
+ Mode: + {entry.guardrail_mode} +
+
+ Status: + + + {statusLabel} + + +
+
+ +
+
+ Start Time: + {formatTime(entry.start_time)} +
+
+ End Time: + {formatTime(entry.end_time)} +
+
+ Duration: + {entry.duration.toFixed(4)}s +
+
+
+ + {totalMaskedEntities > 0 && ( +
+
Masked Entity Summary
+
+ {Object.entries(maskedEntityCount).map(([entityType, count]) => ( + + {entityType}: {count} + + ))} +
+
+ )} + + {guardrailProvider === "presidio" && presidioEntities.length > 0 && ( +
+ +
+ )} + + {guardrailProvider === "bedrock" && bedrockResponse && ( +
+ +
+ )} +
+ ); +}; + const GuardrailViewer = ({ data }: GuardrailViewerProps) => { + const guardrailEntries = Array.isArray(data) + ? data.filter((entry): entry is GuardrailInformation => Boolean(entry)) + : data + ? [data] + : []; + + if (guardrailEntries.length === 0) { + return null; + } + const [sectionExpanded, setSectionExpanded] = useState(true); - // Default to presidio for backwards compatibility - const guardrailProvider = data.guardrail_provider ?? "presidio"; + const primaryName = guardrailEntries.length === 1 ? guardrailEntries[0].guardrail_name : `${guardrailEntries.length} guardrails`; + const statuses = Array.from(new Set(guardrailEntries.map((entry) => entry.guardrail_status))); + const allSucceeded = statuses.every((status) => (status ?? "").toLowerCase() === "success"); + const aggregatedStatus = allSucceeded ? "success" : "failure"; + const totalMaskedEntities = guardrailEntries.reduce((sum, entry) => { + return ( + sum + + Object.values(entry.masked_entity_count || {}).reduce((acc, count) => acc + (typeof count === "number" ? count : 0), 0) + ); + }, 0); - if (!data) return null; - - const isSuccess = typeof data.guardrail_status === "string" && data.guardrail_status.toLowerCase() === "success"; - - const tooltipTitle = isSuccess ? null : "Guardrail failed to run."; - - // Calculate total masked entities - const totalMaskedEntities = data.masked_entity_count - ? Object.values(data.masked_entity_count).reduce((sum, count) => sum + count, 0) - : 0; - - const formatTime = (timestamp: number): string => { - const date = new Date(timestamp * 1000); - return date.toLocaleString(); - }; + const tooltipTitle = allSucceeded ? null : "Guardrail failed to run."; return (
@@ -67,9 +187,9 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => { className="flex justify-between items-center p-4 border-b cursor-pointer hover:bg-gray-50" onClick={() => setSectionExpanded(!sectionExpanded)} > -
+
{

Guardrail Information

- {/* Header status chip with tooltip */} - {data.guardrail_status} + {aggregatedStatus} + {primaryName} + {totalMaskedEntities > 0 && ( - + {totalMaskedEntities} masked {totalMaskedEntities === 1 ? "entity" : "entities"} )} @@ -99,76 +220,15 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
{sectionExpanded && ( -
-
-
-
-
- Guardrail Name: - {data.guardrail_name} -
-
- Mode: - {data.guardrail_mode} -
-
- Status: - - - {data.guardrail_status} - - -
-
- -
-
- Start Time: - {formatTime(data.start_time)} -
-
- End Time: - {formatTime(data.end_time)} -
-
- Duration: - {data.duration.toFixed(4)}s -
-
-
- - {/* Masked Entity Summary */} - {data.masked_entity_count && Object.keys(data.masked_entity_count).length > 0 && ( -
-

Masked Entity Summary

-
- {Object.entries(data.masked_entity_count).map(([entityType, count]) => ( - - {entityType}: {count} - - ))} -
-
- )} -
- - {/* Provider-specific Detected Entities */} - {guardrailProvider === "presidio" && (data.guardrail_response as GuardrailEntity[])?.length > 0 && ( - - )} - - {guardrailProvider === "bedrock" && data.guardrail_response && ( -
- -
- )} +
+ {guardrailEntries.map((entry, index) => ( + + ))}
)}
diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index bfc57a84f7d..fb833bdfca3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -790,20 +790,34 @@ export function RequestViewer({ row }: { row: Row }) { metadata.vector_store_request_metadata.length > 0; // Extract guardrail information from metadata if available - const hasGuardrailData = row.original.metadata && row.original.metadata.guardrail_information; + const guardrailInfo = row.original.metadata?.guardrail_information; + const guardrailEntries = Array.isArray(guardrailInfo) + ? guardrailInfo + : guardrailInfo + ? [guardrailInfo] + : []; + const hasGuardrailData = guardrailEntries.length > 0; // Calculate total masked entities if guardrail data exists - const getTotalMaskedEntities = (): number => { - if (!hasGuardrailData || !row.original.metadata?.guardrail_information.masked_entity_count) { - return 0; + const totalMaskedEntities = guardrailEntries.reduce((sum, entry) => { + const maskedCounts = entry?.masked_entity_count; + if (!maskedCounts) { + return sum; } - return Object.values(row.original.metadata.guardrail_information.masked_entity_count).reduce( - (sum: number, count: any) => sum + (typeof count === "number" ? count : 0), - 0, + return ( + sum + + Object.values(maskedCounts).reduce( + (acc, count) => (typeof count === "number" ? acc + count : acc), + 0, + ) ); - }; + }, 0); - const totalMaskedEntities = getTotalMaskedEntities(); + const primaryGuardrailLabel = guardrailEntries.length === 1 + ? guardrailEntries[0]?.guardrail_name ?? "-" + : guardrailEntries.length > 1 + ? `${guardrailEntries.length} guardrails` + : "-"; return (
@@ -850,7 +864,7 @@ export function RequestViewer({ row }: { row: Row }) {
Guardrail:
- {row.original.metadata!.guardrail_information.guardrail_name} + {primaryGuardrailLabel} {totalMaskedEntities > 0 && ( {totalMaskedEntities} masked @@ -934,7 +948,7 @@ export function RequestViewer({ row }: { row: Row }) {
{/* Guardrail Data - Show only if present */} - {hasGuardrailData && } + {hasGuardrailData && } {/* Vector Store Request Data - Show only if present */} {hasVectorStoreData && }