diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 9a4e9a315ea..0e804cbfd12 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -48,7 +48,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime USER root # Install runtime dependencies -RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip libsndfile +RUN apk add --no-cache bash openssl tzdata nodejs npm python3 py3-pip WORKDIR /app # Copy the current directory contents into the container at /app diff --git a/docs/my-website/docs/mcp_control.md b/docs/my-website/docs/mcp_control.md index a7d66a6b7fc..f6ab4a73087 100644 --- a/docs/my-website/docs/mcp_control.md +++ b/docs/my-website/docs/mcp_control.md @@ -108,7 +108,7 @@ Some MCP servers are meant to be shared broadly—think internal knowledge bases 3. Toggle **Allow All LiteLLM Keys** on. MCP server configuration in Admin UI @@ -634,18 +634,3 @@ Control which tools different teams can access from the same MCP server. For exa This video shows how to set allowed tools for a Key, Team, or Organization. - - -## Dashboard View Modes - -Proxy admins can also control what non-admins see inside the MCP dashboard via `general_settings.user_mcp_management_mode`: - -- `restricted` *(default)* – users only see servers that their team explicitly has access to. -- `view_all` – every dashboard user can see the full MCP server list. - -```yaml title="Config example" -general_settings: - user_mcp_management_mode: view_all -``` - -This is useful when you want discoverability for MCP offerings without granting additional execution privileges. diff --git a/docs/my-website/docs/providers/gigachat.md b/docs/my-website/docs/providers/gigachat.md deleted file mode 100644 index 13eec298c25..00000000000 --- a/docs/my-website/docs/providers/gigachat.md +++ /dev/null @@ -1,283 +0,0 @@ -import Tabs from '@theme/Tabs'; -import TabItem from '@theme/TabItem'; - -# GigaChat -https://developers.sber.ru/docs/ru/gigachat/api/overview - -GigaChat is Sber AI's large language model, Russia's leading LLM provider. - -:::tip - -**We support ALL GigaChat models, just set `model=gigachat/` as a prefix when sending litellm requests** - -::: - -:::warning - -GigaChat API uses self-signed SSL certificates. You must pass `ssl_verify=False` in your requests. - -::: - -## Supported Features - -| Feature | Supported | -|---------|-----------| -| Chat Completion | Yes | -| Streaming | Yes | -| Async | Yes | -| Function Calling / Tools | Yes | -| Structured Output (JSON Schema) | Yes (via function call emulation) | -| Image Input | Yes (base64 and URL) - GigaChat-2-Max, GigaChat-2-Pro only | -| Embeddings | Yes | - -## API Key - -GigaChat uses OAuth authentication. Set your credentials as environment variables: - -```python -import os - -# Required: Set credentials (base64-encoded client_id:client_secret) -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -# Optional: Set scope (default is GIGACHAT_API_PERS for personal use) -os.environ['GIGACHAT_SCOPE'] = "GIGACHAT_API_PERS" # or GIGACHAT_API_B2B for business -``` - -Get your credentials at: https://developers.sber.ru/studio/ - -## Sample Usage - -```python -from litellm import completion -import os - -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -response = completion( - model="gigachat/GigaChat-2-Max", - messages=[ - {"role": "user", "content": "Hello from LiteLLM!"} - ], - ssl_verify=False, # Required for GigaChat -) -print(response) -``` - -## Sample Usage - Streaming - -```python -from litellm import completion -import os - -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -response = completion( - model="gigachat/GigaChat-2-Max", - messages=[ - {"role": "user", "content": "Hello from LiteLLM!"} - ], - stream=True, - ssl_verify=False, # Required for GigaChat -) - -for chunk in response: - print(chunk) -``` - -## Sample Usage - Function Calling - -```python -from litellm import completion -import os - -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -tools = [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a city", - "parameters": { - "type": "object", - "properties": { - "city": {"type": "string", "description": "City name"} - }, - "required": ["city"] - } - } -}] - -response = completion( - model="gigachat/GigaChat-2-Max", - messages=[{"role": "user", "content": "What's the weather in Moscow?"}], - tools=tools, - ssl_verify=False, # Required for GigaChat -) -print(response) -``` - -## Sample Usage - Structured Output - -GigaChat supports structured output via JSON schema (emulated through function calling): - -```python -from litellm import completion -import os - -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -response = completion( - model="gigachat/GigaChat-2-Max", - messages=[{"role": "user", "content": "Extract info: John is 30 years old"}], - response_format={ - "type": "json_schema", - "json_schema": { - "name": "person", - "schema": { - "type": "object", - "properties": { - "name": {"type": "string"}, - "age": {"type": "integer"} - } - } - } - }, - ssl_verify=False, # Required for GigaChat -) -print(response) # Returns JSON: {"name": "John", "age": 30} -``` - -## Sample Usage - Image Input - -GigaChat supports image input via base64 or URL (GigaChat-2-Max and GigaChat-2-Pro only): - -```python -from litellm import completion -import os - -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -response = completion( - model="gigachat/GigaChat-2-Max", # Vision requires GigaChat-2-Max or GigaChat-2-Pro - messages=[{ - "role": "user", - "content": [ - {"type": "text", "text": "What's in this image?"}, - {"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}} - ] - }], - ssl_verify=False, # Required for GigaChat -) -print(response) -``` - -## Sample Usage - Embeddings - -```python -from litellm import embedding -import os - -os.environ['GIGACHAT_CREDENTIALS'] = "your-credentials-here" - -response = embedding( - model="gigachat/Embeddings", - input=["Hello world", "How are you?"], - ssl_verify=False, # Required for GigaChat -) -print(response) -``` - -## Usage with LiteLLM Proxy - -### 1. Set GigaChat Models on config.yaml - -```yaml -model_list: - - model_name: gigachat - litellm_params: - model: gigachat/GigaChat-2-Max - api_key: "os.environ/GIGACHAT_CREDENTIALS" - ssl_verify: false - - model_name: gigachat-lite - litellm_params: - model: gigachat/GigaChat-2-Lite - api_key: "os.environ/GIGACHAT_CREDENTIALS" - ssl_verify: false - - model_name: gigachat-embeddings - litellm_params: - model: gigachat/Embeddings - api_key: "os.environ/GIGACHAT_CREDENTIALS" - ssl_verify: false -``` - -### 2. Start Proxy - -```bash -litellm --config config.yaml -``` - -### 3. Test it - - - - -```shell -curl --location 'http://0.0.0.0:4000/chat/completions' \ ---header 'Content-Type: application/json' \ ---data '{ - "model": "gigachat", - "messages": [ - { - "role": "user", - "content": "Hello!" - } - ] -}' -``` - - - -```python -import openai -client = openai.OpenAI( - api_key="anything", - base_url="http://0.0.0.0:4000" -) - -response = client.chat.completions.create( - model="gigachat", - messages=[{"role": "user", "content": "Hello!"}] -) -print(response) -``` - - - -## Supported Models - -### Chat Models - -| Model Name | Context Window | Vision | Description | -|------------|----------------|--------|-------------| -| gigachat/GigaChat-2-Lite | 128K | No | Fast, lightweight model | -| gigachat/GigaChat-2-Pro | 128K | Yes | Professional model with vision | -| gigachat/GigaChat-2-Max | 128K | Yes | Maximum capability model | - -### Embedding Models - -| Model Name | Max Input | Dimensions | Description | -|------------|-----------|------------|-------------| -| gigachat/Embeddings | 512 | 1024 | Standard embeddings | -| gigachat/Embeddings-2 | 512 | 1024 | Updated embeddings | -| gigachat/EmbeddingsGigaR | 4096 | 2560 | High-dimensional embeddings | - -:::note -Available models may vary depending on your API access level (personal or business). -::: - -## Limitations - -- Only one function call per request (GigaChat API limitation) -- Maximum 1 image per message, 10 images total per conversation -- GigaChat API uses self-signed SSL certificates - `ssl_verify=False` is required diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index f4359a86ba9..87064d442ae 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -111,7 +111,6 @@ general_settings: master_key: string maximum_spend_logs_retention_period: 30d # The maximum time to retain spend logs before deletion. maximum_spend_logs_retention_interval: 1d # interval in which the spend log cleanup task should run in. - user_mcp_management_mode: restricted # or "view_all" # Database Settings database_url: string @@ -231,7 +230,6 @@ router_settings: | image_generation_model | str | The default model to use for image generation - ignores model set in request | | store_model_in_db | boolean | If true, enables storing model + credential information in the DB. | | supported_db_objects | List[str] | Fine-grained control over which object types to load from the database when `store_model_in_db` is True. Available types: `"models"`, `"mcp"`, `"guardrails"`, `"vector_stores"`, `"pass_through_endpoints"`, `"prompts"`, `"model_cost_map"`. If not set, all object types are loaded (default behavior). Example: `supported_db_objects: ["mcp"]` to only load MCP servers from DB. | -| user_mcp_management_mode | string | Controls what non-admins can see on the MCP dashboard. `restricted` (default) only lists MCP servers that the user’s teams are explicitly allowed to access. `view_all` lets every user see the full MCP server list. Tool list/call always respects per-key permissions, so users still cannot run MCP calls without access. | | store_prompts_in_spend_logs | boolean | If true, allows prompts and responses to be stored in the spend logs table. | | max_request_size_mb | int | The maximum size for requests in MB. Requests above this size will be rejected. | | max_response_size_mb | int | The maximum size for responses in MB. LLM Responses above this size will not be sent. | @@ -671,7 +669,6 @@ router_settings: | LANGSMITH_DEFAULT_RUN_NAME | Default name for Langsmith run | LANGSMITH_PROJECT | Project name for Langsmith integration | LANGSMITH_SAMPLING_RATE | Sampling rate for Langsmith logging -| LANGSMITH_TENANT_ID | Tenant ID for Langsmith multi-tenant deployments | LANGTRACE_API_KEY | API key for Langtrace service | LASSO_API_BASE | Base URL for Lasso API | LASSO_API_KEY | API key for Lasso service @@ -710,7 +707,6 @@ router_settings: | LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development) | LITELLM_NON_ROOT | Flag to run LiteLLM in non-root mode for enhanced security in Docker containers | LITELLM_RATE_LIMIT_WINDOW_SIZE | Rate limit window size for LiteLLM. Default is 60 -| LITELLM_REASONING_AUTO_SUMMARY | If set to "true", automatically enables detailed reasoning summaries for reasoning models (e.g., o1, o3-mini, deepseek-reasoner). When enabled, adds `summary: "detailed"` to reasoning effort configurations. Default is "false" | LITELLM_SALT_KEY | Salt key for encryption in LiteLLM | LITELLM_SSL_CIPHERS | SSL/TLS cipher configuration for faster handshakes. Controls cipher suite preferences for OpenSSL connections. | LITELLM_SECRET_AWS_KMS_LITELLM_LICENSE | AWS KMS encrypted license for LiteLLM @@ -778,7 +774,6 @@ router_settings: | OTEL_EXPORTER_OTLP_HEADERS | Headers for OpenTelemetry requests | OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry | OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing -| OTEL_LOGS_EXPORTER | Exporter type for OpenTelemetry logs (e.g., console) | PAGERDUTY_API_KEY | API key for PagerDuty Alerting | PANW_PRISMA_AIRS_API_KEY | API key for PANW Prisma AIRS service | PANW_PRISMA_AIRS_API_BASE | Base URL for PANW Prisma AIRS service @@ -893,4 +888,4 @@ router_settings: | DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL | Time-to-live in seconds for health check lock in shared health check mode. Default is 60 (1 minute) | ZSCALER_AI_GUARD_API_KEY | API key for Zscaler AI Guard service | ZSCALER_AI_GUARD_POLICY_ID | Policy ID for Zscaler AI Guard guardrails -| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy +| ZSCALER_AI_GUARD_URL | Base URL for Zscaler AI Guard API. Default is https://api.us1.zseclipse.net/v1/detection/execute-policy \ No newline at end of file diff --git a/docs/my-website/docs/reasoning_content.md b/docs/my-website/docs/reasoning_content.md index 04c6d7ee6cc..fca3df638c7 100644 --- a/docs/my-website/docs/reasoning_content.md +++ b/docs/my-website/docs/reasoning_content.md @@ -591,68 +591,3 @@ Expected Response - -## OpenAI Responses API - Auto-Summary Control - -When using OpenAI Responses API models (like `gpt-5`) via `/chat/completions` with `reasoning_effort`, you can control whether `summary="detailed"` is automatically added to the reasoning parameter. - -### Enabling Auto-Summary - -You can enable automatic `summary="detailed"` in two ways: - - - - -```python -import litellm - -# Enable auto-summary globally -litellm.reasoning_auto_summary = True - -response = litellm.completion( - model="openai/responses/gpt-5-mini", - messages=[{"role": "user", "content": "What is the capital of France?"}], - reasoning_effort="low", # Will automatically add summary="detailed" -) -``` - - - - - -```bash -# Set environment variable -export LITELLM_REASONING_AUTO_SUMMARY=true - -# Or in your .env file -LITELLM_REASONING_AUTO_SUMMARY=true -``` - - - - - -```yaml -litellm_settings: - reasoning_auto_summary: true # Enable auto-summary for all requests - -model_list: - - model_name: gpt-5-mini - litellm_params: - model: openai/responses/gpt-5-mini -``` - - - - -### Manual Control (Recommended) - -For fine-grained control, pass `reasoning_effort` as a dictionary: - -```python -response = litellm.completion( - model="openai/responses/gpt-5-mini", - messages=[{"role": "user", "content": "What is the capital of France?"}], - reasoning_effort={"effort": "low", "summary": "detailed"}, # Explicit control -) -``` diff --git a/docs/my-website/docs/response_api_compact.md b/docs/my-website/docs/response_api_compact.md deleted file mode 100644 index f5caa32ea33..00000000000 --- a/docs/my-website/docs/response_api_compact.md +++ /dev/null @@ -1,104 +0,0 @@ -import Tabs from '@theme/Tabs'; -import TabItem from '@theme/TabItem'; - -# /responses/compact - -Compress conversation history using OpenAI's `/responses/compact` endpoint. - -| Feature | Supported | -|---------|-----------| -| Supported LiteLLM Versions | 1.72.0+ | -| Supported Providers | `openai` | - -## Usage - -### LiteLLM Python SDK - -```python showLineNumbers title="Compact Response" -import litellm - -response = litellm.compact_responses( - model="openai/gpt-4o", - input=[{"role": "user", "content": "Hello, how are you?"}], - instructions="Be helpful", - previous_response_id="resp_abc123" # optional -) - -print(response.id) -print(response.object) # "response.compaction" -print(response.output) -``` - -### LiteLLM Proxy - - - - -```bash showLineNumbers title="Compact Request" -curl http://localhost:4000/v1/responses/compact \ - -H "Content-Type: application/json" \ - -H "Authorization: Bearer sk-1234" \ - -d '{ - "model": "openai/gpt-4o", - "input": [{"role": "user", "content": "Hello"}], - "instructions": "Be helpful" - }' -``` - - - - -```python showLineNumbers title="Compact with OpenAI SDK" -import httpx - -response = httpx.post( - "http://localhost:4000/v1/responses/compact", - headers={"Authorization": "Bearer sk-1234"}, - json={ - "model": "openai/gpt-4o", - "input": [{"role": "user", "content": "Hello"}], - "instructions": "Be helpful" - } -) - -print(response.json()) -``` - - - - -## Request Parameters - -| Parameter | Type | Required | Description | -|-----------|------|----------|-------------| -| `model` | string | Yes | Model to use for compaction | -| `input` | string or array | Yes | Input messages to compact | -| `instructions` | string | No | System instructions | -| `previous_response_id` | string | No | ID of previous response to continue from | - -## Response Format - -```json -{ - "id": "resp_abc123", - "object": "response.compaction", - "created_at": 1734366691, - "output": [ - { - "type": "message", - "role": "assistant", - "content": [...] - }, - { - "type": "compaction", - "encrypted_content": "..." - } - ], - "usage": { - "input_tokens": 100, - "output_tokens": 50, - "total_tokens": 150 - } -} -``` - diff --git a/docs/my-website/img/mcp_allow_all_ui.png b/docs/my-website/img/mcp_allow_all_ui.png deleted file mode 100644 index f074deb801e..00000000000 Binary files a/docs/my-website/img/mcp_allow_all_ui.png and /dev/null differ diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 11effc82fe7..3793aec037f 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -541,14 +541,7 @@ const sidebars = { }, "realtime", "rerank", - { - type: "category", - label: "/responses", - items: [ - "response_api", - "response_api_compact", - ] - }, + "response_api", { type: "category", label: "/search", diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index a1fa4ef8440..7ffbe95be13 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -10,7 +10,6 @@ from pathlib import Path from typing import Optional from litellm_proxy_extras._logging import logger -from litellm.caching.redis_cache import RedisCache def str_to_bool(value: Optional[str]) -> bool: @@ -19,103 +18,6 @@ def str_to_bool(value: Optional[str]) -> bool: return value.lower() in ("true", "1", "t", "y", "yes") -class MigrationLockManager: - """Redis-based lock manager for database migrations""" - - MIGRATION_LOCK_KEY = "migration_lock" - LOCK_TTL_SECONDS = 300 # 5 minutes TTL - - def __init__(self, redis_cache: Optional[RedisCache] = None): - self.redis_cache = redis_cache - self.lock_acquired = False - self.pod_id = f"pod_{os.getpid()}_{int(time.time())}" - - def _get_redis_lock_key(self) -> str: - """Get Redis lock key for migration""" - return f"migration_lock:{self.MIGRATION_LOCK_KEY}" - - def acquire_lock(self) -> bool: - """Acquire migration lock""" - if self.redis_cache is None: - logger.warning( - "Redis cache is not available, running migration without lock protection" - ) - self.lock_acquired = True - return True - - try: - lock_key = self._get_redis_lock_key() - - # Redis SET with NX (only if not exists) and EX (expiration) - acquired = self.redis_cache.set_cache( - key=lock_key, value=self.pod_id, nx=True, ttl=self.LOCK_TTL_SECONDS - ) - - if acquired: - self.lock_acquired = True - logger.info(f"Migration lock acquired by pod {self.pod_id}") - return True - else: - logger.info("Migration lock is already held by another pod") - return False - - except Exception as e: - logger.warning(f"Failed to acquire migration lock: {e}") - return False - - def wait_for_lock_release( - self, check_interval: int = 5, max_wait: int = 300 - ) -> bool: - """Wait for another process to release the lock""" - if self.redis_cache is None: - logger.warning("Redis cache is not available, cannot wait for lock") - return False - - logger.info(f"Waiting for migration lock to be released (max {max_wait}s)...") - start_time = time.time() - - while time.time() - start_time < max_wait: - # Try to acquire lock using the public acquire_lock method - if self.acquire_lock(): - logger.info( - f"Migration lock acquired after waiting by pod {self.pod_id}" - ) - return True - - time.sleep(check_interval) - - logger.warning(f"Failed to acquire migration lock within {max_wait} seconds") - return False - - def release_lock(self): - """Release migration lock""" - if not self.lock_acquired or self.redis_cache is None: - return - - try: - lock_key = self._get_redis_lock_key() - - # Verify current pod owns the lock - current_value = self.redis_cache.get_cache(lock_key) - if current_value and str(current_value) == self.pod_id: - self.redis_cache.delete_cache(lock_key) - logger.info(f"Migration lock released by pod {self.pod_id}") - else: - logger.warning(f"Pod {self.pod_id} cannot release lock (not owner)") - - except Exception as e: - logger.warning(f"Failed to release migration lock: {e}") - finally: - self.lock_acquired = False - - def __enter__(self): - """Context manager entry - acquire lock when entering with statement""" - self.acquire_lock() - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Context manager exit - release lock when exiting with statement""" - self.release_lock() def _get_prisma_env() -> dict: """Get environment variables for Prisma, handling offline mode if configured.""" @@ -444,50 +346,19 @@ class ProxyExtrasDBManager: ) @staticmethod - def setup_database( - use_migrate: bool = False, redis_cache: Optional[RedisCache] = None - ) -> bool: + def setup_database(use_migrate: bool = False) -> bool: """ Set up the database using either prisma migrate or prisma db push - Uses migrations from litellm-proxy-extras package. - In multi-instance environment, use redis lock to prevent concurrent execution. + Uses migrations from litellm-proxy-extras package Args: schema_path (str): Path to the Prisma schema file use_migrate (bool): Whether to use prisma migrate instead of db push - redis_cache: Redis cache instance for distributed locking Returns: bool: True if setup was successful, False otherwise """ schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma" - - database_url = os.getenv("DATABASE_URL") - if not database_url: - logger.error("DATABASE_URL environment variable is not set") - return False - - # Use MigrationLockManager to prevent concurrent migration execution - with MigrationLockManager(redis_cache) as lock_manager: - # Lock is already acquired in __enter__, check if it was successful - if not lock_manager.lock_acquired: - # Cannot acquire lock, another process is running migration - logger.info( - "Another pod is running migration, waiting for completion..." - ) - - # Wait for other process to complete migration - if not lock_manager.wait_for_lock_release(): - logger.error("Failed to acquire migration lock after waiting") - return False - - # Successfully acquired lock, proceed with migration - logger.info("Acquired migration lock, proceeding with migration") - return ProxyExtrasDBManager._execute_migration(use_migrate, schema_path) - - @staticmethod - def _execute_migration(use_migrate: bool, schema_path: str) -> bool: - """Execute the actual migration""" for attempt in range(4): original_dir = os.getcwd() migrations_dir = ProxyExtrasDBManager._get_prisma_dir() diff --git a/litellm/__init__.py b/litellm/__init__.py index 7f7ee21f692..dfe959bf747 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -197,7 +197,6 @@ retry = True api_key: Optional[str] = None openai_key: Optional[str] = None groq_key: Optional[str] = None -gigachat_key: Optional[str] = None databricks_key: Optional[str] = None openai_like_key: Optional[str] = None azure_key: Optional[str] = None @@ -276,7 +275,6 @@ banned_keywords_list: Optional[Union[str, List]] = None llm_guard_mode: Literal["all", "key-specific", "request-specific"] = "all" guardrail_name_config_map: Dict[str, GuardrailItem] = {} include_cost_in_streaming_usage: bool = False -reasoning_auto_summary: bool = False ### PROMPTS #### from litellm.types.prompts.init_prompts import PromptSpec @@ -1442,8 +1440,6 @@ if TYPE_CHECKING: from .llms.github_copilot.chat.transformation import GithubCopilotConfig as GithubCopilotConfig from .llms.github_copilot.responses.transformation import GithubCopilotResponsesAPIConfig as GithubCopilotResponsesAPIConfig from .llms.github_copilot.embedding.transformation import GithubCopilotEmbeddingConfig as GithubCopilotEmbeddingConfig - from .llms.gigachat.chat.transformation import GigaChatConfig as GigaChatConfig - from .llms.gigachat.embedding.transformation import GigaChatEmbeddingConfig as GigaChatEmbeddingConfig from .llms.nebius.chat.transformation import NebiusConfig as NebiusConfig from .llms.wandb.chat.transformation import WandbConfig as WandbConfig from .llms.dashscope.chat.transformation import DashScopeChatConfig as DashScopeChatConfig diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 26133ebc222..8c8266d5ca4 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -255,8 +255,6 @@ LLM_CONFIG_NAMES = ( "GithubCopilotEmbeddingConfig", "NebiusConfig", "WandbConfig", - "GigaChatConfig", - "GigaChatEmbeddingConfig", "DashScopeChatConfig", "MoonshotChatConfig", "DockerModelRunnerChatConfig", @@ -646,8 +644,6 @@ _LLM_CONFIGS_IMPORT_MAP = { "GithubCopilotEmbeddingConfig": (".llms.github_copilot.embedding.transformation", "GithubCopilotEmbeddingConfig"), "NebiusConfig": (".llms.nebius.chat.transformation", "NebiusConfig"), "WandbConfig": (".llms.wandb.chat.transformation", "WandbConfig"), - "GigaChatConfig": (".llms.gigachat.chat.transformation", "GigaChatConfig"), - "GigaChatEmbeddingConfig": (".llms.gigachat.embedding.transformation", "GigaChatEmbeddingConfig"), "DashScopeChatConfig": (".llms.dashscope.chat.transformation", "DashScopeChatConfig"), "MoonshotChatConfig": (".llms.moonshot.chat.transformation", "MoonshotChatConfig"), "DockerModelRunnerChatConfig": (".llms.docker_model_runner.chat.transformation", "DockerModelRunnerChatConfig"), diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index a89efc4e82b..5b206317b29 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -3,7 +3,6 @@ Handler for transforming /chat/completions api requests to litellm.responses req """ import json -import os from typing import ( TYPE_CHECKING, Any, @@ -23,7 +22,6 @@ from typing import ( from openai.types.responses.tool_param import FunctionToolParam from pydantic import BaseModel -import litellm from litellm import ModelResponse from litellm._logging import verbose_logger from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator @@ -693,26 +691,19 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if isinstance(reasoning_effort, dict): return Reasoning(**reasoning_effort) # type: ignore[typeddict-item] - # Check if auto-summary is enabled via flag or environment variable - # Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var - auto_summary_enabled = ( - litellm.reasoning_auto_summary - or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true" - ) - - # If string is passed, map with optional summary based on flag/env var + # If string is passed, map with summary="detailed" if reasoning_effort == "none": - return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") # type: ignore + return Reasoning(effort="none", summary="detailed") # type: ignore elif reasoning_effort == "high": - return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high") + return Reasoning(effort="high", summary="detailed") elif reasoning_effort == "xhigh": - return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item] + return Reasoning(effort="xhigh", summary="detailed") # type: ignore[typeddict-item] elif reasoning_effort == "medium": - return Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium") + return Reasoning(effort="medium", summary="detailed") elif reasoning_effort == "low": - return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low") + return Reasoning(effort="low", summary="detailed") elif reasoning_effort == "minimal": - return Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal") + return Reasoning(effort="minimal", summary="detailed") return None def _transform_response_format_to_text_format( diff --git a/litellm/constants.py b/litellm/constants.py index 1cd2da549ca..e8524a87c41 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -375,7 +375,6 @@ LITELLM_CHAT_PROVIDERS = [ "perplexity", "mistral", "groq", - "gigachat", "nvidia_nim", "cerebras", "baseten", diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 6b30b6b736e..88f7908e9a2 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -187,12 +187,6 @@ "ui_name": "Sampling Rate", "description": "Sampling rate for logging (0.0 to 1.0, default: 1.0)", "required": false - }, - "langsmith_tenant_id": { - "type": "text", - "ui_name": "Tenant ID", - "description": "LangSmith tenant ID for organization-scoped API keys (required when using org-scoped keys)", - "required": false } }, "description": "Langsmith Logging Integration" diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 7e62613a7e4..f0f355b4895 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -50,42 +50,6 @@ else: Langfuse = Any -def _extract_cache_read_input_tokens(usage_obj) -> int: - """ - Extract cache_read_input_tokens from usage object. - - Checks both: - 1. Top-level cache_read_input_tokens (Anthropic format) - 2. prompt_tokens_details.cached_tokens (Gemini, OpenAI format) - - See: https://github.com/BerriAI/litellm/issues/18520 - - Args: - usage_obj: Usage object from LLM response - - Returns: - int: Number of cached tokens read, defaults to 0 - """ - cache_read_input_tokens = usage_obj.get("cache_read_input_tokens") or 0 - - # Check prompt_tokens_details.cached_tokens (used by Gemini and other providers) - if hasattr(usage_obj, "prompt_tokens_details"): - prompt_tokens_details = getattr(usage_obj, "prompt_tokens_details", None) - if ( - prompt_tokens_details is not None - and hasattr(prompt_tokens_details, "cached_tokens") - ): - cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None) - if ( - cached_tokens is not None - and isinstance(cached_tokens, (int, float)) - and cached_tokens > 0 - ): - cache_read_input_tokens = cached_tokens - - return cache_read_input_tokens - - class LangFuseLogger: # Class variables or attributes def __init__( @@ -793,8 +757,8 @@ class LangFuseLogger: cache_creation_input_tokens = ( _usage_obj.get("cache_creation_input_tokens") or 0 ) - cache_read_input_tokens = _extract_cache_read_input_tokens( - _usage_obj + cache_read_input_tokens = ( + _usage_obj.get("cache_read_input_tokens") or 0 ) usage = { diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py index 570b78f2927..cc9b361b69d 100644 --- a/litellm/integrations/langsmith.py +++ b/litellm/integrations/langsmith.py @@ -40,7 +40,6 @@ class LangsmithLogger(CustomBatchLogger): langsmith_project: Optional[str] = None, langsmith_base_url: Optional[str] = None, langsmith_sampling_rate: Optional[float] = None, - langsmith_tenant_id: Optional[str] = None, **kwargs, ): self.flush_lock = asyncio.Lock() @@ -49,7 +48,6 @@ class LangsmithLogger(CustomBatchLogger): langsmith_api_key=langsmith_api_key, langsmith_project=langsmith_project, langsmith_base_url=langsmith_base_url, - langsmith_tenant_id=langsmith_tenant_id, ) self.sampling_rate: float = ( langsmith_sampling_rate @@ -78,7 +76,6 @@ class LangsmithLogger(CustomBatchLogger): langsmith_api_key: Optional[str] = None, langsmith_project: Optional[str] = None, langsmith_base_url: Optional[str] = None, - langsmith_tenant_id: Optional[str] = None, ) -> LangsmithCredentialsObject: _credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY") _credentials_project = ( @@ -89,13 +86,11 @@ class LangsmithLogger(CustomBatchLogger): or os.getenv("LANGSMITH_BASE_URL") or "https://api.smith.langchain.com" ) - _credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID") return LangsmithCredentialsObject( LANGSMITH_API_KEY=_credentials_api_key, LANGSMITH_BASE_URL=_credentials_base_url, LANGSMITH_PROJECT=_credentials_project, - LANGSMITH_TENANT_ID=_credentials_tenant_id, ) def _prepare_log_data( @@ -370,11 +365,8 @@ class LangsmithLogger(CustomBatchLogger): """ langsmith_api_base = credentials["LANGSMITH_BASE_URL"] langsmith_api_key = credentials["LANGSMITH_API_KEY"] - langsmith_tenant_id = credentials.get("LANGSMITH_TENANT_ID") url = self._add_endpoint_to_url(langsmith_api_base, "runs/batch") headers = {"x-api-key": langsmith_api_key} - if langsmith_tenant_id: - headers["x-tenant-id"] = langsmith_tenant_id elements_to_log = [queue_object["data"] for queue_object in queue_objects] try: @@ -426,7 +418,6 @@ class LangsmithLogger(CustomBatchLogger): api_key=credentials["LANGSMITH_API_KEY"], project=credentials["LANGSMITH_PROJECT"], base_url=credentials["LANGSMITH_BASE_URL"], - tenant_id=credentials.get("LANGSMITH_TENANT_ID"), ) if key not in log_queue_by_credentials: @@ -475,9 +466,6 @@ class LangsmithLogger(CustomBatchLogger): langsmith_base_url=standard_callback_dynamic_params.get( "langsmith_base_url", None ), - langsmith_tenant_id=standard_callback_dynamic_params.get( - "langsmith_tenant_id", None - ), ) else: credentials = self.default_credentials @@ -503,16 +491,13 @@ class LangsmithLogger(CustomBatchLogger): def get_run_by_id(self, run_id): langsmith_api_key = self.default_credentials["LANGSMITH_API_KEY"] + langsmith_api_base = self.default_credentials["LANGSMITH_BASE_URL"] - langsmith_tenant_id = self.default_credentials.get("LANGSMITH_TENANT_ID") url = f"{langsmith_api_base}/runs/{run_id}" - headers = {"x-api-key": langsmith_api_key} - if langsmith_tenant_id: - headers["x-tenant-id"] = langsmith_tenant_id response = litellm.module_level_client.get( url=url, - headers=headers, + headers={"x-api-key": langsmith_api_key}, ) return response.json() diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index a7d2326d938..12e60bc25bb 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -196,88 +196,50 @@ class OpenTelemetry(CustomLogger): litellm.service_callback.append(self) setattr(proxy_server, "open_telemetry_logger", self) - def _get_or_create_provider( - self, - provider, - provider_name: str, - get_existing_provider_fn, - sdk_provider_class, - create_new_provider_fn, - set_provider_fn, - ): - """ - Generic helper to get or create an OpenTelemetry provider (Tracer, Meter, or Logger). - - Args: - provider: The provider instance passed to the init function (can be None) - provider_name: Name for logging (e.g., "TracerProvider") - get_existing_provider_fn: Function to get the existing global provider - sdk_provider_class: The SDK provider class to check for (e.g., TracerProvider from SDK) - create_new_provider_fn: Function to create a new provider instance - set_provider_fn: Function to set the provider globally - - Returns: - The provider to use (either existing, new, or explicitly provided) - """ - if provider is not None: - # Provider explicitly provided (e.g., for testing) - # Do NOT call set_provider_fn - the caller is responsible for managing global state - # If they want it to be global, they've already set it before passing it to us - verbose_logger.debug( - "OpenTelemetry: Using provided TracerProvider: %s", - type(provider).__name__, - ) - return provider - - # Check if a provider is already set globally - try: - existing_provider = get_existing_provider_fn() - - # If a real SDK provider exists (set by another SDK like Langfuse), use it - # This uses a positive check for SDK providers instead of a negative check for proxy providers - if isinstance(existing_provider, sdk_provider_class): - verbose_logger.debug( - "OpenTelemetry: Using existing %s: %s", - provider_name, - type(existing_provider).__name__, - ) - provider = existing_provider - # Don't call set_provider to preserve existing context - else: - # Default proxy provider or unknown type, create our own - verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name) - provider = create_new_provider_fn() - set_provider_fn(provider) - except Exception as e: - # Fallback: create a new provider if something goes wrong - verbose_logger.debug( - "OpenTelemetry: Exception checking existing %s, creating new one: %s", - provider_name, - str(e), - ) - provider = create_new_provider_fn() - set_provider_fn(provider) - - return provider - def _init_tracing(self, tracer_provider): from opentelemetry import trace from opentelemetry.sdk.trace import TracerProvider from opentelemetry.trace import SpanKind - def create_tracer_provider(): - provider = TracerProvider(resource=_get_litellm_resource()) - provider.add_span_processor(self._get_span_processor()) - return provider + # use provided tracer or create a new one + if tracer_provider is None: + # Check if a TracerProvider is already set globally (e.g., by Langfuse SDK) + try: + from opentelemetry.trace import ProxyTracerProvider - tracer_provider = self._get_or_create_provider( - provider=tracer_provider, - provider_name="TracerProvider", - get_existing_provider_fn=trace.get_tracer_provider, - sdk_provider_class=TracerProvider, - create_new_provider_fn=create_tracer_provider, - set_provider_fn=trace.set_tracer_provider, - ) + 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__, + ) + tracer_provider = existing_provider + # Don't call set_tracer_provider to preserve existing context + else: + # No real provider exists yet, create our own + verbose_logger.debug("OpenTelemetry: Creating new TracerProvider") + tracer_provider = TracerProvider(resource=_get_litellm_resource()) + tracer_provider.add_span_processor(self._get_span_processor()) + trace.set_tracer_provider(tracer_provider) + except Exception as e: + # Fallback: create a new provider if something goes wrong + verbose_logger.debug( + "OpenTelemetry: Exception checking existing provider, creating new one: %s", + str(e), + ) + tracer_provider = TracerProvider(resource=_get_litellm_resource()) + tracer_provider.add_span_processor(self._get_span_processor()) + trace.set_tracer_provider(tracer_provider) + else: + # Tracer provider explicitly provided (e.g., for testing) + # Do NOT call set_tracer_provider - the caller is responsible for managing global state + # If they want it to be global, they've already set it before passing it to us + verbose_logger.debug( + "OpenTelemetry: Using provided TracerProvider: %s", + type(tracer_provider).__name__, + ) # Grab our tracer from the TracerProvider (not from global context) # This ensures we use the provided TracerProvider (e.g., for testing) @@ -295,24 +257,39 @@ class OpenTelemetry(CustomLogger): return from opentelemetry import metrics - from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.metrics import Histogram, MeterProvider - def create_meter_provider(): - metric_reader = self._get_metric_reader() - return MeterProvider( - metric_readers=[metric_reader], resource=_get_litellm_resource() + # Only create OTLP infrastructure if no custom meter provider is provided + if meter_provider is None: + from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( + OTLPMetricExporter, + ) + from opentelemetry.sdk.metrics.export import ( + AggregationTemporality, + PeriodicExportingMetricReader, ) - meter_provider = self._get_or_create_provider( - provider=meter_provider, - provider_name="MeterProvider", - get_existing_provider_fn=metrics.get_meter_provider, - sdk_provider_class=MeterProvider, - create_new_provider_fn=create_meter_provider, - set_provider_fn=metrics.set_meter_provider, - ) + normalized_endpoint = self._normalize_otel_endpoint( + self.config.endpoint, "metrics" + ) + _metric_exporter = OTLPMetricExporter( + endpoint=normalized_endpoint, + headers=OpenTelemetry._get_headers_dictionary(self.config.headers), + preferred_temporality={Histogram: AggregationTemporality.DELTA}, + ) + _metric_reader = PeriodicExportingMetricReader( + _metric_exporter, export_interval_millis=10000 + ) - meter = meter_provider.get_meter(__name__) + meter_provider = MeterProvider( + metric_readers=[_metric_reader], resource=_get_litellm_resource() + ) + meter = meter_provider.get_meter(__name__) + else: + # Use the provided meter provider as-is, without creating additional OTLP infrastructure + meter = meter_provider.get_meter(__name__) + + metrics.set_meter_provider(meter_provider) self._operation_duration_histogram = meter.create_histogram( name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38 @@ -350,26 +327,22 @@ class OpenTelemetry(CustomLogger): if not self.config.enable_events: return - from opentelemetry._logs import get_logger_provider, set_logger_provider + from opentelemetry._logs import set_logger_provider from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider from opentelemetry.sdk._logs.export import BatchLogRecordProcessor - def create_logger_provider(): - provider = OTLoggerProvider(resource=_get_litellm_resource()) + # set up log pipeline + if logger_provider is None: + litellm_resource = _get_litellm_resource() + logger_provider = OTLoggerProvider(resource=litellm_resource) + # Only add OTLP exporter if we created the logger provider ourselves log_exporter = self._get_log_exporter() - provider.add_log_record_processor( - BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type] - ) - return provider + if log_exporter: + logger_provider.add_log_record_processor( + BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type] + ) - self._get_or_create_provider( - provider=logger_provider, - provider_name="LoggerProvider", - get_existing_provider_fn=get_logger_provider, - sdk_provider_class=OTLoggerProvider, - create_new_provider_fn=create_logger_provider, - set_provider_fn=set_logger_provider, - ) + set_logger_provider(logger_provider) def log_success_event(self, kwargs, response_obj, start_time, end_time): self._handle_success(kwargs, response_obj, start_time, end_time) @@ -971,15 +944,6 @@ class OpenTelemetry(CustomLogger): if not self.config.enable_events: return - # NOTE: Semantic logs (gen_ai.content.prompt/completion events) have compatibility issues - # with OTEL SDK >= 1.39.0 due to breaking changes in PR #4676: - # - LogRecord moved from opentelemetry.sdk._logs to opentelemetry.sdk._logs._internal - # - LogRecord constructor no longer accepts 'resource' parameter (now inherited from LoggerProvider) - # - LogData class was removed entirely - # These logs work correctly in OTEL SDK < 1.39.0 but may fail in >= 1.39.0. - # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676 - # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords - from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider from opentelemetry.sdk._logs import LogRecord as SdkLogRecord @@ -1843,8 +1807,7 @@ class OpenTelemetry(CustomLogger): ) return self.OTEL_EXPORTER - otel_logs_exporter = os.getenv("OTEL_LOGS_EXPORTER") - if self.OTEL_EXPORTER == "console" or otel_logs_exporter == "console": + if self.OTEL_EXPORTER == "console": from opentelemetry.sdk._logs.export import ConsoleLogExporter verbose_logger.debug( @@ -1891,67 +1854,6 @@ class OpenTelemetry(CustomLogger): return ConsoleLogExporter() - def _get_metric_reader(self): - """ - Get the appropriate metric reader based on the configuration. - """ - from opentelemetry.sdk.metrics import Histogram - from opentelemetry.sdk.metrics.export import ( - AggregationTemporality, - ConsoleMetricExporter, - PeriodicExportingMetricReader, - ) - - verbose_logger.debug( - "OpenTelemetry Logger, initializing metric reader\nself.OTEL_EXPORTER: %s\nself.OTEL_ENDPOINT: %s\nself.OTEL_HEADERS: %s", - self.OTEL_EXPORTER, - self.OTEL_ENDPOINT, - self.OTEL_HEADERS, - ) - - _split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS) - normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "metrics") - - if self.OTEL_EXPORTER == "console": - exporter = ConsoleMetricExporter() - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) - - elif ( - self.OTEL_EXPORTER == "otlp_http" - or self.OTEL_EXPORTER == "http/protobuf" - or self.OTEL_EXPORTER == "http/json" - ): - from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( - OTLPMetricExporter, - ) - - exporter = OTLPMetricExporter( - endpoint=normalized_endpoint, - headers=_split_otel_headers, - preferred_temporality={Histogram: AggregationTemporality.DELTA}, - ) - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) - - elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc": - from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( - OTLPMetricExporter, - ) - - exporter = OTLPMetricExporter( - endpoint=normalized_endpoint, - headers=_split_otel_headers, - preferred_temporality={Histogram: AggregationTemporality.DELTA}, - ) - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) - - else: - verbose_logger.warning( - "OpenTelemetry: Unknown metric exporter '%s', defaulting to console. Supported: console, otlp_http, otlp_grpc", - self.OTEL_EXPORTER, - ) - exporter = ConsoleMetricExporter() - return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) - def _normalize_otel_endpoint( self, endpoint: Optional[str], signal_type: str ) -> Optional[str]: diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 6baaae7ae3f..d92af417175 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -2000,56 +2000,24 @@ class CustomStreamWrapper: ) ## Map to OpenAI Exception try: - mapped_exception = exception_type( + raise exception_type( model=self.model, custom_llm_provider=self.custom_llm_provider, original_exception=e, completion_kwargs={}, extra_kwargs={}, ) - except Exception as mapping_error: - mapped_exception = mapping_error + except Exception as e: + from litellm.exceptions import MidStreamFallbackError - def _normalize_status_code(exc: Exception) -> Optional[int]: - """ - Best-effort status_code extraction. - Uses status_code on the exception, then falls back to the response. - """ - try: - code = getattr(exc, "status_code", None) - if code is not None: - return int(code) - except Exception: - pass - - response = getattr(exc, "response", None) - if response is not None: - try: - status_code = getattr(response, "status_code", None) - if status_code is not None: - return int(status_code) - except Exception: - pass - return None - - mapped_status_code = _normalize_status_code(mapped_exception) - original_status_code = _normalize_status_code(e) - - if mapped_status_code is not None and 400 <= mapped_status_code < 500: - raise mapped_exception - if original_status_code is not None and 400 <= original_status_code < 500: - raise mapped_exception - - from litellm.exceptions import MidStreamFallbackError - - raise MidStreamFallbackError( - message=str(mapped_exception), - model=self.model, - llm_provider=self.custom_llm_provider or "anthropic", - original_exception=mapped_exception, - generated_content=self.response_uptil_now, - is_pre_first_chunk=not self.sent_first_chunk, - ) + raise MidStreamFallbackError( + message=str(e), + model=self.model, + llm_provider=self.custom_llm_provider or "anthropic", + original_exception=e, + generated_content=self.response_uptil_now, + is_pre_first_chunk=not self.sent_first_chunk, + ) @staticmethod def _strip_sse_data_from_chunk(chunk: Optional[str]) -> Optional[str]: diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 7a4da985528..facabbda72a 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -242,30 +242,3 @@ class BaseResponsesAPIConfig(ABC): ######################################################### ########## END CANCEL RESPONSE API TRANSFORMATION ####### ######################################################### - - ######################################################### - ########## COMPACT RESPONSE API TRANSFORMATION ########## - ######################################################### - @abstractmethod - def transform_compact_response_api_request( - self, - model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, - api_base: str, - litellm_params: GenericLiteLLMParams, - headers: dict, - ) -> Tuple[str, Dict]: - pass - - @abstractmethod - def transform_compact_response_api_response( - self, - raw_response: httpx.Response, - logging_obj: LiteLLMLoggingObj, - ) -> ResponsesAPIResponse: - pass - - ######################################################### - ########## END COMPACT RESPONSE API TRANSFORMATION ###### - ######################################################### diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ea740400664..34ea598a655 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -91,7 +91,6 @@ from litellm.types.rerank import RerankResponse from litellm.types.responses.main import DeleteResponseResult from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( - CallTypes, EmbeddingResponse, FileTypes, LiteLLMBatch, @@ -851,9 +850,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client( - params={"ssl_verify": litellm_params.get("ssl_verify", None)} - ) + sync_httpx_client = _get_httpx_client() else: sync_httpx_client = client @@ -899,8 +896,7 @@ class BaseLLMHTTPHandler: ) -> EmbeddingResponse: if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( - llm_provider=litellm.LlmProviders(custom_llm_provider), - params={"ssl_verify": litellm_params.get("ssl_verify", None)}, + llm_provider=litellm.LlmProviders(custom_llm_provider) ) else: async_httpx_client = client @@ -2008,10 +2004,6 @@ class BaseLLMHTTPHandler: """ Handles responses API requests. When _is_async=True, returns a coroutine instead of making the call directly. - - Keeps the pre-transform request context for streaming so post-call hooks/metadata - (added for Responses API parity with chat) receive the original params instead of - the provider-shaped body that caused them to be skipped before. """ if _is_async: @@ -2068,18 +2060,6 @@ class BaseLLMHTTPHandler: if extra_body: data.update(extra_body) - # Preserve the OpenAI-style request context (not sent to the provider) for streaming - # hooks/metadata; the streaming iterator now consumes this to run deployment hooks - # with the same info as chat, including litellm_params. - request_context: Dict[str, Any] = {"input": input} - try: - request_context.update(response_api_optional_request_params) - except Exception: - pass - # Needed by streaming callbacks/metadata helpers to reconstruct api_base/model_id - # but never included in the outbound provider payload. - request_context["litellm_params"] = dict(litellm_params) - ## LOGGING logging_obj.pre_call( input=input, @@ -2117,8 +2097,6 @@ class BaseLLMHTTPHandler: responses_api_provider_config=responses_api_provider_config, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, - request_data=request_context, - call_type=CallTypes.responses.value, ) return SyncResponsesAPIStreamingIterator( @@ -2128,8 +2106,6 @@ class BaseLLMHTTPHandler: responses_api_provider_config=responses_api_provider_config, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, - request_data=request_context, - call_type=CallTypes.responses.value, ) else: # For non-streaming requests @@ -2213,18 +2189,6 @@ class BaseLLMHTTPHandler: if extra_body: data.update(extra_body) - # Preserve the OpenAI-style request context (not sent to the provider) for streaming - # hooks/metadata; the streaming iterator now consumes this to run deployment hooks - # with the same info as chat, including litellm_params. - request_context: Dict[str, Any] = {"input": input} - try: - request_context.update(response_api_optional_request_params) - except Exception: - pass - # Needed by streaming callbacks/metadata helpers to reconstruct api_base/model_id - # but never included in the outbound provider payload. - request_context["litellm_params"] = dict(litellm_params) - ## LOGGING logging_obj.pre_call( input=input, @@ -2263,8 +2227,6 @@ class BaseLLMHTTPHandler: responses_api_provider_config=responses_api_provider_config, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, - request_data=request_context, - call_type=CallTypes.responses.value, ) # Return the streaming iterator @@ -2275,8 +2237,6 @@ class BaseLLMHTTPHandler: responses_api_provider_config=responses_api_provider_config, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, - request_data=request_context, - call_type=CallTypes.responses.value, ) else: # For non-streaming, proceed as before @@ -3566,174 +3526,6 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, ) - def compact_response_api_handler( - self, - model: str, - input: Union[str, "ResponseInputParam"], - responses_api_provider_config: BaseResponsesAPIConfig, - response_api_optional_request_params: Dict, - litellm_params: GenericLiteLLMParams, - logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - _is_async: bool = False, - shared_session: Optional["ClientSession"] = None, - ) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: - """ - Handler for the compact responses API. - """ - if _is_async: - return self.async_compact_response_api_handler( - model=model, - input=input, - responses_api_provider_config=responses_api_provider_config, - response_api_optional_request_params=response_api_optional_request_params, - litellm_params=litellm_params, - logging_obj=logging_obj, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - extra_body=extra_body, - timeout=timeout, - client=client, - shared_session=shared_session, - ) - if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client( - params={"ssl_verify": litellm_params.get("ssl_verify", None)} - ) - else: - sync_httpx_client = client - - headers = responses_api_provider_config.validate_environment( - headers=extra_headers or {}, model=model, litellm_params=litellm_params - ) - - if extra_headers: - headers.update(extra_headers) - - api_base = responses_api_provider_config.get_complete_url( - api_base=litellm_params.api_base, - litellm_params=dict(litellm_params), - ) - - url, data = responses_api_provider_config.transform_compact_response_api_request( - model=model, - input=input, - response_api_optional_request_params=response_api_optional_request_params, - api_base=api_base, - litellm_params=litellm_params, - headers=headers, - ) - - ## LOGGING - logging_obj.pre_call( - input=input, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": url, - "headers": headers, - }, - ) - - try: - response = sync_httpx_client.post( - url=url, headers=headers, json=data, timeout=timeout - ) - - except Exception as e: - raise self._handle_error( - e=e, - provider_config=responses_api_provider_config, - ) - - return responses_api_provider_config.transform_compact_response_api_response( - raw_response=response, - logging_obj=logging_obj, - ) - - async def async_compact_response_api_handler( - self, - model: str, - input: Union[str, "ResponseInputParam"], - responses_api_provider_config: BaseResponsesAPIConfig, - response_api_optional_request_params: Dict, - litellm_params: GenericLiteLLMParams, - logging_obj: LiteLLMLoggingObj, - custom_llm_provider: Optional[str], - extra_headers: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - _is_async: bool = False, - shared_session: Optional["ClientSession"] = None, - ) -> ResponsesAPIResponse: - """ - Async version of the compact response API handler. - """ - if client is None or not isinstance(client, AsyncHTTPHandler): - verbose_logger.debug( - f"Creating HTTP client for compact_response with shared_session: {id(shared_session) if shared_session else None}" - ) - async_httpx_client = get_async_httpx_client( - llm_provider=litellm.LlmProviders(custom_llm_provider), - params={"ssl_verify": litellm_params.get("ssl_verify", None)}, - shared_session=shared_session, - ) - else: - async_httpx_client = client - - headers = responses_api_provider_config.validate_environment( - headers=extra_headers or {}, model=model, litellm_params=litellm_params - ) - - if extra_headers: - headers.update(extra_headers) - - api_base = responses_api_provider_config.get_complete_url( - api_base=litellm_params.api_base, - litellm_params=dict(litellm_params), - ) - - url, data = responses_api_provider_config.transform_compact_response_api_request( - model=model, - input=input, - response_api_optional_request_params=response_api_optional_request_params, - api_base=api_base, - litellm_params=litellm_params, - headers=headers, - ) - - ## LOGGING - logging_obj.pre_call( - input=input, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": url, - "headers": headers, - }, - ) - - try: - response = await async_httpx_client.post( - url=url, headers=headers, json=data, timeout=timeout - ) - - except Exception as e: - raise self._handle_error( - e=e, - provider_config=responses_api_provider_config, - ) - - return responses_api_provider_config.transform_compact_response_api_response( - raw_response=response, - logging_obj=logging_obj, - ) - def list_files(self): """ Lists all files @@ -8496,4 +8288,4 @@ class BaseLLMHTTPHandler: return skills_api_provider_config.transform_delete_skill_response( raw_response=response, logging_obj=logging_obj, - ) + ) \ No newline at end of file diff --git a/litellm/llms/gigachat/__init__.py b/litellm/llms/gigachat/__init__.py deleted file mode 100644 index 3ddbd7864d9..00000000000 --- a/litellm/llms/gigachat/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -""" -GigaChat Provider for LiteLLM - -GigaChat is Sber AI's large language model (Russia's leading LLM). -Supports: -- Chat completions (sync/async) -- Streaming (sync/async) -- Function calling / Tools -- Structured output via JSON schema (emulated through function calls) -- Image input (base64 and URL) -- Embeddings - -API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/overview -""" - -from .chat.transformation import GigaChatConfig, GigaChatError -from .embedding.transformation import GigaChatEmbeddingConfig - -__all__ = [ - "GigaChatConfig", - "GigaChatEmbeddingConfig", - "GigaChatError", -] diff --git a/litellm/llms/gigachat/authenticator.py b/litellm/llms/gigachat/authenticator.py deleted file mode 100644 index e61015a4a21..00000000000 --- a/litellm/llms/gigachat/authenticator.py +++ /dev/null @@ -1,241 +0,0 @@ -""" -GigaChat OAuth Authenticator - -Handles OAuth 2.0 token management for GigaChat API. -Based on official GigaChat SDK authentication flow. -""" - -import time -import uuid -from typing import Optional, Tuple - -import httpx - -from litellm._logging import verbose_logger -from litellm.caching.caching import InMemoryCache -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.llms.custom_httpx.http_handler import ( - HTTPHandler, - _get_httpx_client, - get_async_httpx_client, -) -from litellm.secret_managers.main import get_secret_str -from litellm.types.utils import LlmProviders - -# GigaChat OAuth endpoint -GIGACHAT_AUTH_URL = "https://ngw.devices.sberbank.ru:9443/api/v2/oauth" - -# Default scope for personal API access -GIGACHAT_SCOPE = "GIGACHAT_API_PERS" - -# Token expiry buffer in milliseconds (refresh token 60s before expiry) -TOKEN_EXPIRY_BUFFER_MS = 60000 - -# Cache for access tokens -_token_cache = InMemoryCache() - - -class GigaChatAuthError(BaseLLMException): - """GigaChat authentication error.""" - - pass - - -def _get_credentials() -> Optional[str]: - """Get GigaChat credentials from environment.""" - return get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY") - - -def _get_auth_url() -> str: - """Get GigaChat auth URL from environment or use default.""" - return get_secret_str("GIGACHAT_AUTH_URL") or GIGACHAT_AUTH_URL - - -def _get_scope() -> str: - """Get GigaChat scope from environment or use default.""" - return get_secret_str("GIGACHAT_SCOPE") or GIGACHAT_SCOPE - - -def _get_http_client() -> HTTPHandler: - """Get cached httpx client with SSL verification disabled.""" - return _get_httpx_client(params={"ssl_verify": False}) - - -def get_access_token( - credentials: Optional[str] = None, - scope: Optional[str] = None, - auth_url: Optional[str] = None, -) -> str: - """ - Get valid access token, using cache if available. - - Args: - credentials: Base64-encoded credentials (client_id:client_secret) - scope: API scope (GIGACHAT_API_PERS, GIGACHAT_API_CORP, etc.) - auth_url: OAuth endpoint URL - - Returns: - Access token string - - Raises: - GigaChatAuthError: If authentication fails - """ - credentials = credentials or _get_credentials() - if not credentials: - raise GigaChatAuthError( - status_code=401, - message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.", - ) - - scope = scope or _get_scope() - auth_url = auth_url or _get_auth_url() - - # Check cache - cache_key = f"gigachat_token:{credentials[:16]}" - cached = _token_cache.get_cache(cache_key) - if cached: - token, expires_at = cached - # Check if token is still valid (with buffer) - if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS: - verbose_logger.debug("Using cached GigaChat access token") - return token - - # Request new token - token, expires_at = _request_token_sync(credentials, scope, auth_url) - - # Cache token - ttl_seconds = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000) - if ttl_seconds > 0: - _token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds) - - return token - - -async def get_access_token_async( - credentials: Optional[str] = None, - scope: Optional[str] = None, - auth_url: Optional[str] = None, -) -> str: - """Async version of get_access_token.""" - credentials = credentials or _get_credentials() - if not credentials: - raise GigaChatAuthError( - status_code=401, - message="GigaChat credentials not provided. Set GIGACHAT_CREDENTIALS or GIGACHAT_API_KEY environment variable.", - ) - - scope = scope or _get_scope() - auth_url = auth_url or _get_auth_url() - - # Check cache - cache_key = f"gigachat_token:{credentials[:16]}" - cached = _token_cache.get_cache(cache_key) - if cached: - token, expires_at = cached - if time.time() * 1000 < expires_at - TOKEN_EXPIRY_BUFFER_MS: - verbose_logger.debug("Using cached GigaChat access token") - return token - - # Request new token - token, expires_at = await _request_token_async(credentials, scope, auth_url) - - # Cache token - ttl_seconds = max(0, (expires_at - TOKEN_EXPIRY_BUFFER_MS - time.time() * 1000) / 1000) - if ttl_seconds > 0: - _token_cache.set_cache(cache_key, (token, expires_at), ttl=ttl_seconds) - - return token - - -def _request_token_sync( - credentials: str, - scope: str, - auth_url: str, -) -> Tuple[str, int]: - """ - Request new access token from GigaChat OAuth endpoint (sync). - - Returns: - Tuple of (access_token, expires_at_ms) - """ - headers = { - "Authorization": f"Basic {credentials}", - "RqUID": str(uuid.uuid4()), - "Content-Type": "application/x-www-form-urlencoded", - } - data = {"scope": scope} - - verbose_logger.debug(f"Requesting GigaChat access token from {auth_url}") - - try: - client = _get_http_client() - response = client.post(auth_url, headers=headers, data=data, timeout=30) - response.raise_for_status() - return _parse_token_response(response) - except httpx.HTTPStatusError as e: - raise GigaChatAuthError( - status_code=e.response.status_code, - message=f"GigaChat authentication failed: {e.response.text}", - ) - except httpx.RequestError as e: - raise GigaChatAuthError( - status_code=500, - message=f"GigaChat authentication request failed: {str(e)}", - ) - - -async def _request_token_async( - credentials: str, - scope: str, - auth_url: str, -) -> Tuple[str, int]: - """Async version of _request_token_sync.""" - headers = { - "Authorization": f"Basic {credentials}", - "RqUID": str(uuid.uuid4()), - "Content-Type": "application/x-www-form-urlencoded", - } - data = {"scope": scope} - - verbose_logger.debug(f"Requesting GigaChat access token from {auth_url}") - - try: - client = get_async_httpx_client( - llm_provider=LlmProviders.GIGACHAT, - params={"ssl_verify": False}, - ) - response = await client.post(auth_url, headers=headers, data=data, timeout=30) - response.raise_for_status() - return _parse_token_response(response) - except httpx.HTTPStatusError as e: - raise GigaChatAuthError( - status_code=e.response.status_code, - message=f"GigaChat authentication failed: {e.response.text}", - ) - except httpx.RequestError as e: - raise GigaChatAuthError( - status_code=500, - message=f"GigaChat authentication request failed: {str(e)}", - ) - - -def _parse_token_response(response: httpx.Response) -> Tuple[str, int]: - """Parse OAuth token response.""" - data = response.json() - - # GigaChat returns either 'tok'/'exp' or 'access_token'/'expires_at' - access_token = data.get("tok") or data.get("access_token") - expires_at = data.get("exp") or data.get("expires_at") - - if not access_token: - raise GigaChatAuthError( - status_code=500, - message=f"Invalid token response: {data}", - ) - - # expires_at is in milliseconds - if isinstance(expires_at, str): - expires_at = int(expires_at) - - verbose_logger.debug("GigaChat access token obtained successfully") - return access_token, expires_at diff --git a/litellm/llms/gigachat/chat/__init__.py b/litellm/llms/gigachat/chat/__init__.py deleted file mode 100644 index 3e030497a1a..00000000000 --- a/litellm/llms/gigachat/chat/__init__.py +++ /dev/null @@ -1,12 +0,0 @@ -""" -GigaChat Chat Module -""" - -from .transformation import GigaChatConfig, GigaChatError -from .streaming import GigaChatModelResponseIterator - -__all__ = [ - "GigaChatConfig", - "GigaChatError", - "GigaChatModelResponseIterator", -] diff --git a/litellm/llms/gigachat/chat/streaming.py b/litellm/llms/gigachat/chat/streaming.py deleted file mode 100644 index 3565559e43c..00000000000 --- a/litellm/llms/gigachat/chat/streaming.py +++ /dev/null @@ -1,134 +0,0 @@ -""" -GigaChat Streaming Response Handler -""" - -import json -import uuid -from typing import Any, Optional - -from litellm.types.llms.openai import ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk -from litellm.types.utils import GenericStreamingChunk - - -class GigaChatModelResponseIterator: - """Iterator for GigaChat streaming responses.""" - - def __init__( - self, - streaming_response: Any, - sync_stream: bool, - json_mode: Optional[bool] = False, - ): - self.streaming_response = streaming_response - self.response_iterator = self.streaming_response - self.json_mode = json_mode - - def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: - """Parse a single streaming chunk from GigaChat.""" - text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None - is_finished = False - finish_reason: Optional[str] = None - - choices = chunk.get("choices", []) - if not choices: - return GenericStreamingChunk( - text="", - tool_use=None, - is_finished=False, - finish_reason="", - usage=None, - index=0, - ) - - choice = choices[0] - delta = choice.get("delta", {}) - finish_reason = choice.get("finish_reason") - - # Extract text content - text = delta.get("content", "") or "" - - # Handle function_call in stream - if finish_reason == "function_call" and delta.get("function_call"): - func_call = delta["function_call"] - args = func_call.get("arguments", {}) - - if isinstance(args, dict): - args = json.dumps(args, ensure_ascii=False) - - tool_use = ChatCompletionToolCallChunk( - id=f"call_{uuid.uuid4().hex[:24]}", - type="function", - function=ChatCompletionToolCallFunctionChunk( - name=func_call.get("name", ""), - arguments=args, - ), - index=0, - ) - finish_reason = "tool_calls" - - if finish_reason is not None: - is_finished = True - - return GenericStreamingChunk( - text=text, - tool_use=tool_use, - is_finished=is_finished, - finish_reason=finish_reason or "", - usage=None, - index=choice.get("index", 0), - ) - - def __iter__(self): - return self - - def __next__(self) -> GenericStreamingChunk: - try: - chunk = self.response_iterator.__next__() - if isinstance(chunk, str): - # Parse SSE format: data: {...} - if chunk.startswith("data: "): - chunk = chunk[6:] - if chunk.strip() == "[DONE]": - raise StopIteration - try: - chunk = json.loads(chunk) - except json.JSONDecodeError: - return GenericStreamingChunk( - text="", - tool_use=None, - is_finished=False, - finish_reason="", - usage=None, - index=0, - ) - return self.chunk_parser(chunk) - except StopIteration: - raise - - def __aiter__(self): - return self - - async def __anext__(self) -> GenericStreamingChunk: - try: - chunk = await self.response_iterator.__anext__() - if isinstance(chunk, str): - # Parse SSE format - if chunk.startswith("data: "): - chunk = chunk[6:] - if chunk.strip() == "[DONE]": - raise StopAsyncIteration - try: - chunk = json.loads(chunk) - except json.JSONDecodeError: - return GenericStreamingChunk( - text="", - tool_use=None, - is_finished=False, - finish_reason="", - usage=None, - index=0, - ) - return self.chunk_parser(chunk) - except StopAsyncIteration: - raise diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py deleted file mode 100644 index 4ce333a1309..00000000000 --- a/litellm/llms/gigachat/chat/transformation.py +++ /dev/null @@ -1,473 +0,0 @@ -""" -GigaChat Chat Transformation - -Transforms OpenAI-format requests to GigaChat format and back. -""" - -import json -import time -import uuid -from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional, Union - -import httpx - -from litellm._logging import verbose_logger -from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException -from litellm.secret_managers.main import get_secret_str -from litellm.types.llms.openai import AllMessageValues -from litellm.types.utils import Choices, Message, ModelResponse, Usage - -from ..authenticator import get_access_token -from ..file_handler import upload_file_sync - -if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj - - LiteLLMLoggingObj = _LiteLLMLoggingObj -else: - LiteLLMLoggingObj = Any - -# GigaChat API endpoint -GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1" - - -class GigaChatError(BaseLLMException): - """GigaChat API error.""" - - pass - - -class GigaChatConfig(BaseConfig): - """ - Configuration class for GigaChat API. - - GigaChat is Sber's (Russia's largest bank) LLM API. - - Supported parameters: - temperature: Sampling temperature (0-2, default 0.87) - top_p: Nucleus sampling parameter - max_tokens: Maximum tokens to generate - repetition_penalty: Repetition penalty factor - profanity_check: Enable content filtering - stream: Enable streaming - """ - - temperature: Optional[float] = None - top_p: Optional[float] = None - max_tokens: Optional[int] = None - repetition_penalty: Optional[float] = None - profanity_check: Optional[bool] = None - - def __init__( - self, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - max_tokens: Optional[int] = None, - repetition_penalty: Optional[float] = None, - profanity_check: Optional[bool] = None, - ) -> None: - locals_ = locals().copy() - for key, value in locals_.items(): - if key != "self" and value is not None: - setattr(self.__class__, key, value) - # Instance variables for current request context - self._current_credentials: Optional[str] = None - self._current_api_base: Optional[str] = None - - def get_complete_url( - self, - api_base: Optional[str], - api_key: Optional[str], - model: str, - optional_params: dict, - litellm_params: dict, - stream: Optional[bool] = None, - ) -> str: - """Get complete API URL for chat completions.""" - base = api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL - return f"{base}/chat/completions" - - def validate_environment( - self, - headers: dict, - model: str, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> dict: - """ - Set up headers with OAuth token. - """ - # Get access token - credentials = api_key or get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY") - access_token = get_access_token(credentials=credentials) - - # Store credentials for image uploads - self._current_credentials = credentials - self._current_api_base = api_base - - headers["Authorization"] = f"Bearer {access_token}" - headers["Content-Type"] = "application/json" - headers["Accept"] = "application/json" - - return headers - - def get_supported_openai_params(self, model: str) -> List[str]: - """Return list of supported OpenAI parameters.""" - return [ - "stream", - "temperature", - "top_p", - "max_tokens", - "max_completion_tokens", - "stop", - "tools", - "tool_choice", - "functions", - "function_call", - "response_format", - ] - - def map_openai_params( - self, - non_default_params: dict, - optional_params: dict, - model: str, - drop_params: bool, - ) -> dict: - """Map OpenAI parameters to GigaChat parameters.""" - for param, value in non_default_params.items(): - if param == "stream": - optional_params["stream"] = value - elif param == "temperature": - # GigaChat: temperature 0 means use top_p=0 instead - if value == 0: - optional_params["top_p"] = 0 - else: - optional_params["temperature"] = value - elif param == "top_p": - optional_params["top_p"] = value - elif param in ("max_tokens", "max_completion_tokens"): - optional_params["max_tokens"] = value - elif param == "stop": - # GigaChat doesn't support stop sequences - pass - elif param == "tools": - # Convert tools to functions format - optional_params["functions"] = self._convert_tools_to_functions(value) - elif param == "tool_choice": - if isinstance(value, dict) and value.get("function"): - optional_params["function_call"] = {"name": value["function"]["name"]} - elif value == "auto": - pass # Default behavior - elif value == "required": - # GigaChat doesn't have 'required', handled differently - pass - elif param == "functions": - optional_params["functions"] = value - elif param == "function_call": - optional_params["function_call"] = value - elif param == "response_format": - # Handle structured output via function calling - if value.get("type") == "json_schema": - json_schema = value.get("json_schema", {}) - schema_name = json_schema.get("name", "structured_output") - schema = json_schema.get("schema", {}) - - function_def = { - "name": schema_name, - "description": f"Output structured response: {schema_name}", - "parameters": schema, - } - - if "functions" not in optional_params: - optional_params["functions"] = [] - optional_params["functions"].append(function_def) - optional_params["function_call"] = {"name": schema_name} - optional_params["_structured_output"] = True - - return optional_params - - def _convert_tools_to_functions(self, tools: List[dict]) -> List[dict]: - """Convert OpenAI tools format to GigaChat functions format.""" - functions = [] - for tool in tools: - if tool.get("type") == "function": - func = tool.get("function", {}) - functions.append({ - "name": func.get("name", ""), - "description": func.get("description", ""), - "parameters": func.get("parameters", {}), - }) - return functions - - def _upload_image(self, image_url: str) -> Optional[str]: - """ - Upload image to GigaChat and return file_id. - - Args: - image_url: URL or base64 data URL of the image - - Returns: - file_id string or None if upload failed - """ - try: - return upload_file_sync( - image_url=image_url, - credentials=self._current_credentials, - api_base=self._current_api_base, - ) - except Exception as e: - verbose_logger.error(f"Failed to upload image: {e}") - return None - - def transform_request( - self, - model: str, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - headers: dict, - ) -> dict: - """Transform OpenAI request to GigaChat format.""" - # Transform messages - giga_messages = self._transform_messages(messages) - - # Build request - request_data = { - "model": model.replace("gigachat/", ""), - "messages": giga_messages, - } - - # Add optional params - for key in ["temperature", "top_p", "max_tokens", "stream", - "repetition_penalty", "profanity_check"]: - if key in optional_params: - request_data[key] = optional_params[key] - - # Add functions if present - if "functions" in optional_params: - request_data["functions"] = optional_params["functions"] - if "function_call" in optional_params: - request_data["function_call"] = optional_params["function_call"] - - return request_data - - def _transform_messages(self, messages: List[AllMessageValues]) -> List[dict]: - """Transform OpenAI messages to GigaChat format.""" - transformed = [] - - for i, msg in enumerate(messages): - message = dict(msg) - - # Remove unsupported fields - message.pop("name", None) - - # Transform roles - role = message.get("role", "user") - if role == "developer": - message["role"] = "system" - elif role == "system" and i > 0: - # GigaChat only allows system message as first message - message["role"] = "user" - elif role == "tool": - message["role"] = "function" - content = message.get("content", "") - if not isinstance(content, str): - message["content"] = json.dumps(content, ensure_ascii=False) - - # Handle None content - if message.get("content") is None: - message["content"] = "" - - # Handle list content (multimodal) - extract text and images - content = message.get("content") - if isinstance(content, list): - texts = [] - attachments = [] - for part in content: - if isinstance(part, dict): - if part.get("type") == "text": - texts.append(part.get("text", "")) - elif part.get("type") == "image_url": - # Extract image URL and upload to GigaChat - image_url = part.get("image_url", {}) - if isinstance(image_url, str): - url = image_url - else: - url = image_url.get("url", "") - if url: - file_id = self._upload_image(url) - if file_id: - attachments.append(file_id) - message["content"] = "\n".join(texts) if texts else "" - if attachments: - message["attachments"] = attachments - - # Transform tool_calls to function_call - tool_calls = message.get("tool_calls") - if tool_calls and isinstance(tool_calls, list) and len(tool_calls) > 0: - tool_call = tool_calls[0] - func = tool_call.get("function", {}) - args = func.get("arguments", "{}") - if isinstance(args, str): - try: - args = json.loads(args) - except json.JSONDecodeError: - args = {} - message["function_call"] = { - "name": func.get("name", ""), - "arguments": args, - } - message.pop("tool_calls", None) - - transformed.append(message) - - # Collapse consecutive user messages - return self._collapse_user_messages(transformed) - - def _collapse_user_messages(self, messages: List[dict]) -> List[dict]: - """Collapse consecutive user messages into one.""" - collapsed: List[dict] = [] - prev_user_msg: Optional[dict] = None - content_parts: List[str] = [] - - for msg in messages: - if msg.get("role") == "user" and prev_user_msg is not None: - content_parts.append(msg.get("content", "")) - else: - if content_parts and prev_user_msg: - prev_user_msg["content"] = "\n".join( - [prev_user_msg.get("content", "")] + content_parts - ) - content_parts = [] - collapsed.append(msg) - prev_user_msg = msg if msg.get("role") == "user" else None - - if content_parts and prev_user_msg: - prev_user_msg["content"] = "\n".join( - [prev_user_msg.get("content", "")] + content_parts - ) - - return collapsed - - def transform_response( - self, - model: str, - raw_response: httpx.Response, - model_response: ModelResponse, - logging_obj: LiteLLMLoggingObj, - request_data: dict, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - encoding: Any, - api_key: Optional[str] = None, - json_mode: Optional[bool] = None, - ) -> ModelResponse: - """Transform GigaChat response to OpenAI format.""" - try: - response_json = raw_response.json() - except Exception: - raise GigaChatError( - status_code=raw_response.status_code, - message=f"Invalid JSON response: {raw_response.text}", - ) - - is_structured_output = optional_params.get("_structured_output", False) - - choices = [] - for choice in response_json.get("choices", []): - message_data = choice.get("message", {}) - finish_reason = choice.get("finish_reason", "stop") - - # Transform function_call to tool_calls or content - if finish_reason == "function_call" and message_data.get("function_call"): - func_call = message_data["function_call"] - args = func_call.get("arguments", {}) - - if is_structured_output: - # Convert to content for structured output - if isinstance(args, dict): - content = json.dumps(args, ensure_ascii=False) - else: - content = str(args) - message_data["content"] = content - message_data.pop("function_call", None) - message_data.pop("functions_state_id", None) - finish_reason = "stop" - else: - # Convert to tool_calls format - if isinstance(args, dict): - args = json.dumps(args, ensure_ascii=False) - message_data["tool_calls"] = [{ - "id": f"call_{uuid.uuid4().hex[:24]}", - "type": "function", - "function": { - "name": func_call.get("name", ""), - "arguments": args, - } - }] - message_data.pop("function_call", None) - finish_reason = "tool_calls" - - # Clean up GigaChat-specific fields - message_data.pop("functions_state_id", None) - - choices.append( - Choices( - index=choice.get("index", 0), - message=Message( - role=message_data.get("role", "assistant"), - content=message_data.get("content"), - tool_calls=message_data.get("tool_calls"), - ), - finish_reason=finish_reason, - ) - ) - - # Build usage - usage_data = response_json.get("usage", {}) - usage = Usage( - prompt_tokens=usage_data.get("prompt_tokens", 0), - completion_tokens=usage_data.get("completion_tokens", 0), - total_tokens=usage_data.get("total_tokens", 0), - ) - - model_response.id = response_json.get("id", f"chatcmpl-{uuid.uuid4().hex[:12]}") - model_response.created = response_json.get("created", int(time.time())) - model_response.model = model - model_response.choices = choices # type: ignore - setattr(model_response, "usage", usage) - - return model_response - - def get_error_class( - self, - error_message: str, - status_code: int, - headers: Union[dict, httpx.Headers], - ) -> BaseLLMException: - """Return GigaChat error class.""" - return GigaChatError( - status_code=status_code, - message=error_message, - headers=headers, - ) - - def get_model_response_iterator( - self, - streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], - sync_stream: bool, - json_mode: Optional[bool] = False, - ): - """Return streaming response iterator.""" - from .streaming import GigaChatModelResponseIterator - - return GigaChatModelResponseIterator( - streaming_response=streaming_response, - sync_stream=sync_stream, - json_mode=json_mode, - ) diff --git a/litellm/llms/gigachat/embedding/__init__.py b/litellm/llms/gigachat/embedding/__init__.py deleted file mode 100644 index af237e49aab..00000000000 --- a/litellm/llms/gigachat/embedding/__init__.py +++ /dev/null @@ -1,7 +0,0 @@ -""" -GigaChat Embedding Module -""" - -from .transformation import GigaChatEmbeddingConfig - -__all__ = ["GigaChatEmbeddingConfig"] diff --git a/litellm/llms/gigachat/embedding/transformation.py b/litellm/llms/gigachat/embedding/transformation.py deleted file mode 100644 index 0da6565050e..00000000000 --- a/litellm/llms/gigachat/embedding/transformation.py +++ /dev/null @@ -1,212 +0,0 @@ -""" -GigaChat Embedding Transformation - -Transforms OpenAI /v1/embeddings format to GigaChat format. -API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/rest/post-embeddings -""" - -import types -from typing import List, Optional, Tuple, Union - -import httpx - -from litellm import LlmProviders -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues -from litellm.types.utils import EmbeddingResponse - -from ..authenticator import get_access_token - -# GigaChat API endpoint -GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1" - - -class GigaChatEmbeddingError(BaseLLMException): - """GigaChat Embedding API error.""" - - pass - - -class GigaChatEmbeddingConfig(BaseEmbeddingConfig): - """ - Configuration class for GigaChat Embeddings API. - - GigaChat embeddings endpoint: POST /api/v1/embeddings - """ - - def __init__(self) -> None: - pass - - @classmethod - def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } - - def get_supported_openai_params(self, model: str) -> List[str]: - """GigaChat embeddings don't support additional parameters.""" - return [] - - def map_openai_params( - self, - non_default_params: dict, - optional_params: dict, - model: str, - drop_params: bool, - ) -> dict: - """Map OpenAI params to GigaChat format (no special mapping needed).""" - return optional_params - - def _get_openai_compatible_provider_info( - self, - api_base: Optional[str], - api_key: Optional[str], - ) -> Tuple[str, Optional[str], Optional[str]]: - """ - Returns provider info for GigaChat. - - Returns: - Tuple of (custom_llm_provider, api_base, dynamic_api_key) - """ - api_base = api_base or GIGACHAT_BASE_URL - return LlmProviders.GIGACHAT.value, api_base, api_key - - def get_complete_url( - self, - api_base: Optional[str], - api_key: Optional[str], - model: str, - optional_params: dict, - litellm_params: dict, - stream: Optional[bool] = None, - ) -> str: - """Get the complete URL for embeddings endpoint.""" - base = api_base or GIGACHAT_BASE_URL - return f"{base}/embeddings" - - def transform_embedding_request( - self, - model: str, - input: AllEmbeddingInputValues, - optional_params: dict, - headers: dict, - ) -> dict: - """ - Transform OpenAI embedding request to GigaChat format. - - GigaChat format: - { - "model": "Embeddings", - "input": ["text1", "text2", ...] - } - """ - # Normalize input to list - if isinstance(input, str): - input_list: list = [input] - elif isinstance(input, list): - input_list = input - else: - input_list = [input] - - # Remove gigachat/ prefix from model if present - if model.startswith("gigachat/"): - model = model[9:] - - return { - "model": model, - "input": input_list, - } - - def transform_embedding_response( - self, - model: str, - raw_response: httpx.Response, - model_response: EmbeddingResponse, - logging_obj: LiteLLMLoggingObj, - api_key: Optional[str], - request_data: dict, - optional_params: dict, - litellm_params: dict, - ) -> EmbeddingResponse: - """ - Transform GigaChat embedding response to OpenAI format. - - GigaChat returns: - { - "object": "list", - "data": [{"object": "embedding", "embedding": [...], "index": 0, "usage": {...}}], - "model": "Embeddings" - } - """ - response_json = raw_response.json() - - # Log response - logging_obj.post_call( - input=request_data.get("input"), - api_key=api_key, - additional_args={"complete_input_dict": request_data}, - original_response=response_json, - ) - - # Calculate total tokens from individual embeddings - total_tokens = 0 - if "data" in response_json: - for emb in response_json["data"]: - if "usage" in emb and "prompt_tokens" in emb["usage"]: - total_tokens += emb["usage"]["prompt_tokens"] - # Remove usage from individual embeddings (not part of OpenAI format) - if "usage" in emb: - del emb["usage"] - - # Set overall usage - response_json["usage"] = { - "prompt_tokens": total_tokens, - "total_tokens": total_tokens, - } - - return EmbeddingResponse(**response_json) - - def validate_environment( - self, - headers: dict, - model: str, - messages: List[AllMessageValues], - optional_params: dict, - litellm_params: dict, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - ) -> dict: - """ - Set up headers with OAuth token for GigaChat. - """ - # Get access token via OAuth - access_token = get_access_token(api_key) - - default_headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {access_token}", - } - return {**default_headers, **headers} - - def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] - ) -> BaseLLMException: - """Return GigaChat-specific error class.""" - return GigaChatEmbeddingError( - status_code=status_code, - message=error_message, - ) diff --git a/litellm/llms/gigachat/file_handler.py b/litellm/llms/gigachat/file_handler.py deleted file mode 100644 index 200428a747a..00000000000 --- a/litellm/llms/gigachat/file_handler.py +++ /dev/null @@ -1,211 +0,0 @@ -""" -GigaChat File Handler - -Handles file uploads to GigaChat API for image processing. -GigaChat requires files to be uploaded first, then referenced by file_id. -""" - -import base64 -import hashlib -import re -import uuid -from typing import Dict, Optional, Tuple - -from litellm._logging import verbose_logger -from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, - get_async_httpx_client, -) -from litellm.types.utils import LlmProviders - -from .authenticator import get_access_token, get_access_token_async - -# GigaChat API endpoint -GIGACHAT_BASE_URL = "https://gigachat.devices.sberbank.ru/api/v1" - -# Simple in-memory cache for file IDs -_file_cache: Dict[str, str] = {} - - -def _get_url_hash(url: str) -> str: - """Generate hash for URL to use as cache key.""" - return hashlib.sha256(url.encode()).hexdigest() - - -def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]: - """ - Parse data URL (base64 image). - - Returns: - Tuple of (content_bytes, content_type, extension) or None - """ - match = re.match(r"data:([^;]+);base64,(.+)", data_url) - if not match: - return None - - content_type = match.group(1) - base64_data = match.group(2) - content_bytes = base64.b64decode(base64_data) - ext = content_type.split("/")[-1].split(";")[0] or "jpg" - - return content_bytes, content_type, ext - - -def _download_image_sync(url: str) -> Tuple[bytes, str, str]: - """Download image from URL synchronously.""" - client = _get_httpx_client(params={"ssl_verify": False}) - response = client.get(url) - response.raise_for_status() - - content_type = response.headers.get("content-type", "image/jpeg") - ext = content_type.split("/")[-1].split(";")[0] or "jpg" - - return response.content, content_type, ext - - -async def _download_image_async(url: str) -> Tuple[bytes, str, str]: - """Download image from URL asynchronously.""" - client = get_async_httpx_client( - llm_provider=LlmProviders.GIGACHAT, - params={"ssl_verify": False}, - ) - response = await client.get(url) - response.raise_for_status() - - content_type = response.headers.get("content-type", "image/jpeg") - ext = content_type.split("/")[-1].split(";")[0] or "jpg" - - return response.content, content_type, ext - - -def upload_file_sync( - image_url: str, - credentials: Optional[str] = None, - api_base: Optional[str] = None, -) -> Optional[str]: - """ - Upload file to GigaChat and return file_id (sync). - - Args: - image_url: URL or base64 data URL of the image - credentials: GigaChat credentials for auth - api_base: Optional custom API base URL - - Returns: - file_id string or None if upload failed - """ - url_hash = _get_url_hash(image_url) - - # Check cache - if url_hash in _file_cache: - verbose_logger.debug(f"Image found in cache: {url_hash[:16]}...") - return _file_cache[url_hash] - - try: - # Get image data - parsed = _parse_data_url(image_url) - if parsed: - content_bytes, content_type, ext = parsed - verbose_logger.debug("Decoded base64 image") - else: - verbose_logger.debug(f"Downloading image from URL: {image_url[:80]}...") - content_bytes, content_type, ext = _download_image_sync(image_url) - - filename = f"{uuid.uuid4()}.{ext}" - - # Get access token - access_token = get_access_token(credentials) - - # Upload to GigaChat - base_url = api_base or GIGACHAT_BASE_URL - upload_url = f"{base_url}/files" - - client = _get_httpx_client(params={"ssl_verify": False}) - response = client.post( - upload_url, - headers={"Authorization": f"Bearer {access_token}"}, - files={"file": (filename, content_bytes, content_type)}, - data={"purpose": "general"}, - timeout=60, - ) - response.raise_for_status() - result = response.json() - - file_id = result.get("id") - if file_id: - _file_cache[url_hash] = file_id - verbose_logger.debug(f"File uploaded successfully, file_id: {file_id}") - - return file_id - - except Exception as e: - verbose_logger.error(f"Error uploading file to GigaChat: {e}") - return None - - -async def upload_file_async( - image_url: str, - credentials: Optional[str] = None, - api_base: Optional[str] = None, -) -> Optional[str]: - """ - Upload file to GigaChat and return file_id (async). - - Args: - image_url: URL or base64 data URL of the image - credentials: GigaChat credentials for auth - api_base: Optional custom API base URL - - Returns: - file_id string or None if upload failed - """ - url_hash = _get_url_hash(image_url) - - # Check cache - if url_hash in _file_cache: - verbose_logger.debug(f"Image found in cache: {url_hash[:16]}...") - return _file_cache[url_hash] - - try: - # Get image data - parsed = _parse_data_url(image_url) - if parsed: - content_bytes, content_type, ext = parsed - verbose_logger.debug("Decoded base64 image") - else: - verbose_logger.debug(f"Downloading image from URL: {image_url[:80]}...") - content_bytes, content_type, ext = await _download_image_async(image_url) - - filename = f"{uuid.uuid4()}.{ext}" - - # Get access token - access_token = await get_access_token_async(credentials) - - # Upload to GigaChat - base_url = api_base or GIGACHAT_BASE_URL - upload_url = f"{base_url}/files" - - client = get_async_httpx_client( - llm_provider=LlmProviders.GIGACHAT, - params={"ssl_verify": False}, - ) - response = await client.post( - upload_url, - headers={"Authorization": f"Bearer {access_token}"}, - files={"file": (filename, content_bytes, content_type)}, - data={"purpose": "general"}, - timeout=60, - ) - response.raise_for_status() - result = response.json() - - file_id = result.get("id") - if file_id: - _file_cache[url_hash] = file_id - verbose_logger.debug(f"File uploaded successfully, file_id: {file_id}") - - return file_id - - except Exception as e: - verbose_logger.error(f"Error uploading file to GigaChat: {e}") - return None diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index cc2439b431a..96598c1dfe6 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -500,69 +500,3 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): response._hidden_params["headers"] = raw_response_headers return response - - ######################################################### - ########## COMPACT RESPONSE API TRANSFORMATION ########## - ######################################################### - def transform_compact_response_api_request( - self, - model: str, - input: Union[str, ResponseInputParam], - response_api_optional_request_params: Dict, - api_base: str, - litellm_params: GenericLiteLLMParams, - headers: dict, - ) -> Tuple[str, Dict]: - """ - Transform the compact response API request into a URL and data - - OpenAI API expects the following request - - POST /v1/responses/compact - """ - url = f"{api_base}/compact" - - input = self._validate_input_param(input) - data = dict( - ResponsesAPIRequestParams( - model=model, input=input, **response_api_optional_request_params - ) - ) - - return url, data - - def transform_compact_response_api_response( - self, - raw_response: httpx.Response, - logging_obj: LiteLLMLoggingObj, - ) -> ResponsesAPIResponse: - """ - Transform the compact response API response into a ResponsesAPIResponse - """ - try: - logging_obj.post_call( - original_response=raw_response.text, - additional_args={"complete_input_dict": {}}, - ) - raw_response_json = raw_response.json() - raw_response_json["created_at"] = _safe_convert_created_field( - raw_response_json["created_at"] - ) - except Exception: - raise OpenAIError( - message=raw_response.text, status_code=raw_response.status_code - ) - raw_response_headers = dict(raw_response.headers) - processed_headers = process_response_headers(raw_response_headers) - - try: - response = ResponsesAPIResponse(**raw_response_json) - except Exception: - verbose_logger.debug( - f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct" - ) - response = ResponsesAPIResponse.model_construct(**raw_response_json) - - response._hidden_params["additional_headers"] = processed_headers - response._hidden_params["headers"] = raw_response_headers - - return response diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index d23c698cd7a..826f151df35 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -6,7 +6,7 @@ Handles Authentication and generating request urls for Vertex AI and Google AI S import json import os -from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple import litellm from litellm._logging import verbose_logger @@ -168,6 +168,7 @@ class VertexBase: ) def _credentials_from_default_auth(self, scopes): + import google.auth as google_auth return google_auth.default(scopes=scopes) @@ -391,7 +392,7 @@ class VertexBase: Returns token, url """ - version: Optional[Literal["v1", "v1beta1"]] = None + version: Optional[Literal["v1beta1", "v1"]] = None if custom_llm_provider == "gemini": url, endpoint = _get_gemini_url( mode=mode, @@ -414,7 +415,7 @@ class VertexBase: stream=stream, vertex_project=vertex_project, vertex_location=vertex_location, - vertex_api_version=cast(Literal["v1", "v1beta1"], version), + vertex_api_version=version, ) return self._check_custom_proxy( diff --git a/litellm/main.py b/litellm/main.py index e8a8b504d96..f4f27eb5841 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2141,49 +2141,6 @@ def completion( # type: ignore # noqa: PLR0915 logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements client=client, ) - elif custom_llm_provider == "gigachat": - # GigaChat - Sber AI's LLM (Russia) - api_key = ( - api_key - or litellm.api_key - or litellm.gigachat_key - or get_secret("GIGACHAT_API_KEY") - or get_secret("GIGACHAT_CREDENTIALS") - ) - - headers = headers or litellm.headers or {} - - ## COMPLETION CALL - try: - 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=_get_encoding(), - stream=stream, - provider_config=provider_config, - ) - except Exception as e: - ## LOGGING - log the original exception returned - logging.post_call( - input=messages, - api_key=api_key, - original_response=str(e), - additional_args={"headers": headers}, - ) - raise e - elif custom_llm_provider == "sap": headers = headers or litellm.headers ## LOAD CONFIG - if set @@ -5267,28 +5224,6 @@ def embedding( # noqa: PLR0915 aembedding=aembedding, litellm_params={}, ) - elif custom_llm_provider == "gigachat": - api_key = ( - api_key - or litellm.api_key - or litellm.gigachat_key - or get_secret_str("GIGACHAT_CREDENTIALS") - or get_secret_str("GIGACHAT_API_KEY") - ) - response = base_llm_http_handler.embedding( - model=model, - input=input, - custom_llm_provider=custom_llm_provider, - api_base=api_base, - api_key=api_key, - logging_obj=logging, - timeout=timeout, - model_response=EmbeddingResponse(), - optional_params=optional_params, - client=client, - aembedding=aembedding, - litellm_params={"ssl_verify": kwargs.get("ssl_verify", None)}, - ) else: raise LiteLLMUnknownProvider( model=model, custom_llm_provider=custom_llm_provider diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c7a2f60856d..e823dd5dc6b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -15831,68 +15831,6 @@ "max_tokens": 8191, "mode": "embedding" }, - "gigachat/GigaChat-2-Lite": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 0.0, - "supports_function_calling": true, - "supports_system_messages": true - }, - "gigachat/GigaChat-2-Max": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 0.0, - "supports_function_calling": true, - "supports_system_messages": true, - "supports_vision": true - }, - "gigachat/GigaChat-2-Pro": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 0.0, - "supports_function_calling": true, - "supports_system_messages": true, - "supports_vision": true - }, - "gigachat/Embeddings": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 512, - "max_tokens": 512, - "mode": "embedding", - "output_cost_per_token": 0.0, - "output_vector_size": 1024 - }, - "gigachat/Embeddings-2": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 512, - "max_tokens": 512, - "mode": "embedding", - "output_cost_per_token": 0.0, - "output_vector_size": 1024 - }, - "gigachat/EmbeddingsGigaR": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 4096, - "max_tokens": 4096, - "mode": "embedding", - "output_cost_per_token": 0.0, - "output_vector_size": 2560 - }, "google.gemma-3-12b-it": { "input_cost_per_token": 9e-08, "litellm_provider": "bedrock_converse", @@ -32154,4 +32092,3 @@ "mode": "chat" } } - diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ded591a8f53..ffa17a5b7c4 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -15,7 +15,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( ) from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.mcp_server.mcp_server_manager import MCPServer -from litellm.proxy.utils import get_server_root_path router = APIRouter( tags=["mcp"], @@ -382,30 +381,13 @@ async def callback(code: str, state: str): # ------------------------------ # Optional .well-known endpoints for MCP + OAuth discovery # ------------------------------ -""" - Per SEP-985, the client MUST: - 1. Try resource_metadata from WWW-Authenticate header (if present) - 2. Fall back to path-based well-known URI: /.well-known/oauth-protected-resource/{path} - ( - If the resource identifier value contains a path or query component, any terminating slash (/) - following the host component MUST be removed before inserting /.well-known/ and the well-known - URI path suffix between the host component and the path(include root path) and/or query components. - https://datatracker.ietf.org/doc/html/rfc9728#section-3.1) - 3. Fall back to root-based well-known URI: /.well-known/oauth-protected-resource -""" -@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp") +@router.get("/.well-known/oauth-protected-resource/{mcp_server_name}/mcp") @router.get("/.well-known/oauth-protected-resource") async def oauth_protected_resource_mcp( request: Request, mcp_server_name: Optional[str] = None ): - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) # Get the correct base URL considering X-Forwarded-* headers request_base_url = get_request_base_url(request) - mcp_server: Optional[MCPServer] = None - if mcp_server_name: - mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name) return { "authorization_servers": [ ( @@ -419,25 +401,14 @@ async def oauth_protected_resource_mcp( if mcp_server_name else f"{request_base_url}/mcp" ), # this is what Claude will call - "scopes_supported": mcp_server.scopes if mcp_server else [], } -""" - https://datatracker.ietf.org/doc/html/rfc8414#section-3.1 - RFC 8414: Path-aware OAuth discovery - If the issuer identifier value contains a path component, any - terminating "/" MUST be removed before inserting "/.well-known/" and - the well-known URI suffix between the host component and the path(include root path) - component. -""" -@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}") + +@router.get("/.well-known/oauth-authorization-server/{mcp_server_name}") @router.get("/.well-known/oauth-authorization-server") async def oauth_authorization_server_mcp( request: Request, mcp_server_name: Optional[str] = None ): - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) # Get the correct base URL considering X-Forwarded-* headers request_base_url = get_request_base_url(request) @@ -452,21 +423,16 @@ async def oauth_authorization_server_mcp( else f"{request_base_url}/token" ) - mcp_server: Optional[MCPServer] = None - if mcp_server_name: - mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name) - return { "issuer": request_base_url, # point to your proxy "authorization_endpoint": authorization_endpoint, "token_endpoint": token_endpoint, "response_types_supported": ["code"], - "scopes_supported": mcp_server.scopes if mcp_server else [], - "grant_types_supported": ["authorization_code", "refresh_token"], + "grant_types_supported": ["authorization_code"], "code_challenge_methods_supported": ["S256"], "token_endpoint_auth_methods_supported": ["client_secret_post"], # Claude expects a registration endpoint, even if we just fake it - "registration_endpoint": f"{request_base_url}/{mcp_server_name}/register" if mcp_server_name else f"{request_base_url}/register", + "registration_endpoint": f"{request_base_url}/{mcp_server_name}/register", } diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3a548e203c5..a5ac966062e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -660,14 +660,14 @@ class MCPServerManager: """ allowed_mcp_servers = await self.get_allowed_mcp_servers(user_api_key_auth) + list_tools_result: List[MCPTool] = [] verbose_logger.debug("SERVER MANAGER LISTING TOOLS") - async def _fetch_server_tools(server_id: str) -> List[MCPTool]: - """Fetch tools from a single server with error handling.""" + for server_id in allowed_mcp_servers: server = self.get_mcp_server_by_id(server_id) if server is None: verbose_logger.warning(f"MCP Server {server_id} not found") - return [] + continue # Get server-specific auth header if available server_auth_header = None @@ -685,21 +685,15 @@ class MCPServerManager: server=server, mcp_auth_header=server_auth_header, ) - return tools + list_tools_result.extend(tools) + verbose_logger.info( + f"Successfully fetched {len(tools)} tools from server {server.name}" + ) except Exception as e: verbose_logger.warning( f"Failed to list tools from server {server.name}: {str(e)}. Continuing with other servers." ) - return [] - - # Fetch tools from all servers in parallel - tasks = [_fetch_server_tools(server_id) for server_id in allowed_mcp_servers] - results = await asyncio.gather(*tasks) - - # Flatten results into single list - list_tools_result: List[MCPTool] = [ - tool for tools in results for tool in tools - ] + # Continue with other servers instead of failing completely verbose_logger.info( f"Successfully fetched {len(list_tools_result)} tools total from all servers" @@ -2009,9 +2003,6 @@ class MCPServerManager: Note: This now handles prefixed tool names """ for server in self.get_registry().values(): - if server.auth_type == MCPAuth.oauth2: - # Skip OAuth2 servers for now as they may require user-specific tokens - continue tools = await self._get_tools_from_server(server) for tool in tools: # The tool.name here is already prefixed from _get_tools_from_server @@ -2293,7 +2284,14 @@ class MCPServerManager: # Check all accessible servers target_server_ids = allowed_server_ids - return await self._run_health_checks(target_server_ids) + # Run health checks concurrently + tasks = [self.health_check_server(server_id) for server_id in target_server_ids] + results = await asyncio.gather(*tasks) + + # Filter out None results (servers that were not found) + list_mcp_servers = [server for server in results if server is not None] + + return list_mcp_servers async def get_all_allowed_mcp_servers( self, @@ -2308,6 +2306,8 @@ class MCPServerManager: Returns: List of MCP server objects without health status """ + from datetime import datetime + # Get allowed server IDs allowed_server_ids = await self.get_allowed_mcp_servers(user_api_key_auth) @@ -2319,56 +2319,40 @@ class MCPServerManager: verbose_logger.warning(f"MCP Server {server_id} not found in registry") continue - mcp_server_table = self._build_mcp_server_table(server) + # Build LiteLLM_MCPServerTable without health check + mcp_server_table = LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.server_name, + alias=server.alias, + description=( + server.mcp_info.get("description") if server.mcp_info else None + ), + url=server.url, + transport=server.transport, + auth_type=server.auth_type, + created_at=datetime.now(), + updated_at=datetime.now(), + teams=[], + mcp_access_groups=server.access_groups or [], + allowed_tools=server.allowed_tools or [], + extra_headers=server.extra_headers or [], + mcp_info=server.mcp_info, + static_headers=server.static_headers, + status=None, # No health check performed + last_health_check=None, # No health check performed + health_check_error=None, + command=getattr(server, "command", None), + args=getattr(server, "args", None) or [], + env=getattr(server, "env", None) or {}, + authorization_url=server.authorization_url, + token_url=server.token_url, + registration_url=server.registration_url, + allow_all_keys=server.allow_all_keys, + ) list_mcp_servers.append(mcp_server_table) return list_mcp_servers - def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: - from datetime import datetime - - return LiteLLM_MCPServerTable( - server_id=server.server_id, - server_name=server.server_name, - alias=server.alias, - description=( - server.mcp_info.get("description") if server.mcp_info else None - ), - url=server.url, - transport=server.transport, - auth_type=server.auth_type, - created_at=datetime.now(), - updated_at=datetime.now(), - teams=[], - mcp_access_groups=server.access_groups or [], - allowed_tools=server.allowed_tools or [], - extra_headers=server.extra_headers or [], - mcp_info=server.mcp_info, - static_headers=server.static_headers, - status=None, # No health check performed - last_health_check=None, # No health check performed - health_check_error=None, - command=getattr(server, "command", None), - args=getattr(server, "args", None) or [], - env=getattr(server, "env", None) or {}, - authorization_url=server.authorization_url, - token_url=server.token_url, - registration_url=server.registration_url, - allow_all_keys=server.allow_all_keys, - ) - - async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]: - """Return all MCP servers from registry without applying access controls.""" - - registry = self.get_registry() - if not registry: - return [] - - servers: List[LiteLLM_MCPServerTable] = [] - for server in registry.values(): - servers.append(self._build_mcp_server_table(server)) - return servers - async def reload_servers_from_database(self): """ Public method to reload all MCP servers from database into registry. @@ -2376,34 +2360,5 @@ class MCPServerManager: """ await self._add_mcp_servers_from_db_to_in_memory_registry() - async def get_all_mcp_servers_with_health_unfiltered( - self, server_ids: Optional[List[str]] = None - ) -> List[LiteLLM_MCPServerTable]: - """Return health info for all servers in registry regardless of user access.""" - - registry = self.get_registry() - if not registry: - return [] - - if server_ids: - target_server_ids = [sid for sid in server_ids if sid in registry] - else: - target_server_ids = list(registry.keys()) - - if not target_server_ids: - return [] - - return await self._run_health_checks(target_server_ids) - - async def _run_health_checks( - self, target_server_ids: List[str] - ) -> List[LiteLLM_MCPServerTable]: - if not target_server_ids: - return [] - - tasks = [self.health_check_server(server_id) for server_id in target_server_ids] - results = await asyncio.gather(*tasks) - return [server for server in results if server is not None] - global_mcp_server_manager: MCPServerManager = MCPServerManager() diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9c7001266f0..e00fdbfb930 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -709,8 +709,7 @@ if MCP_AVAILABLE: extra_headers: Optional[Dict[str, str]] = None if server.auth_type == MCPAuth.oauth2: - # Copy to avoid mutating the original dict (important for parallel fetching) - extra_headers = oauth2_headers.copy() if oauth2_headers else None + extra_headers = oauth2_headers if server.extra_headers and raw_headers: if extra_headers is None: @@ -756,10 +755,11 @@ if MCP_AVAILABLE: # Decide whether to add prefix based on number of allowed servers add_prefix = not (len(allowed_mcp_servers) == 1) - async def _fetch_and_filter_server_tools(server: MCPServer) -> List[MCPTool]: - """Fetch and filter tools from a single server with error handling.""" + # Get tools from each allowed server + all_tools = [] + for server in allowed_mcp_servers: if server is None: - return [] + continue server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, @@ -786,24 +786,16 @@ if MCP_AVAILABLE: user_api_key_auth=user_api_key_auth, ) + all_tools.extend(filtered_tools) + verbose_logger.debug( f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering" ) - return filtered_tools except Exception as e: verbose_logger.exception( f"Error getting tools from server {server.name}: {str(e)}" ) - return [] - - # Fetch tools from all servers in parallel - tasks = [ - _fetch_and_filter_server_tools(server) for server in allowed_mcp_servers - ] - results = await asyncio.gather(*tasks) - - # Flatten results into single list - all_tools: List[MCPTool] = [tool for tools in results for tool in tools] + # Continue with other servers instead of failing completely verbose_logger.info( f"Successfully fetched {len(all_tools)} tools total from all MCP servers" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 954c26e2cb2..b77cc40d6dc 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1908,9 +1908,6 @@ class UserHeaderMapping(LiteLLMPydanticObjectBase): } -UserMCPManagementMode = Literal["restricted", "view_all"] - - class ConfigGeneralSettings(LiteLLMPydanticObjectBase): """ Documents all the fields supported by `general_settings` in config.yaml @@ -2028,10 +2025,6 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).", ) - user_mcp_management_mode: Optional[UserMCPManagementMode] = Field( - None, - description="Controls how non-admin users interact with MCP servers in the dashboard. 'restricted' shows only accessible servers, 'view_all' lists every server in read-only mode.", - ) class ConfigYAML(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index ae7cb612c9c..7947eb9b045 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -323,7 +323,6 @@ class ProxyBaseLLMRequestProcessing: "aget_responses", "adelete_responses", "acancel_responses", - "acompact_responses", "acreate_batch", "aretrieve_batch", "alist_batches", @@ -485,7 +484,6 @@ class ProxyBaseLLMRequestProcessing: "aget_responses", "adelete_responses", "acancel_responses", - "acompact_responses", "atext_completion", "aimage_edit", "alist_input_items", diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index fe27af78d58..406ddceabf5 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -154,11 +154,7 @@ class PrismaManager: prisma_dir = PrismaManager._get_prisma_dir() - from litellm.proxy.proxy_server import redis_usage_cache - - return ProxyExtrasDBManager.setup_database( - use_migrate=use_migrate, redis_cache=redis_usage_cache - ) + return ProxyExtrasDBManager.setup_database(use_migrate=use_migrate) else: # Use prisma db push with increased timeout subprocess.run( diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index 5850103132c..ea8f1b0a97f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -118,7 +118,7 @@ class LassoGuardrail(CustomGuardrail): Falls back to UUID if ULID library is not available. """ if ULID_AVAILABLE and ulid is not None: - return str(ulid.ULID()) # type: ignore + return str(ulid.new()) # type: ignore else: verbose_proxy_logger.debug("ULID library not available, using UUID") return str(uuid.uuid4()) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 47793c8fc8e..a871a6637a2 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -32,8 +32,8 @@ from fastapi import ( from fastapi.responses import JSONResponse import litellm -from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._uuid import uuid +from litellm._logging import verbose_logger, verbose_proxy_logger from litellm.constants import LITELLM_PROXY_ADMIN_NAME from litellm.proxy._experimental.mcp_server.utils import ( validate_and_normalize_mcp_server_payload, @@ -67,6 +67,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.ui_session_utils import ( build_effective_auth_contexts, ) + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.proxy._types import ( LiteLLM_MCPServerTable, LitellmUserRoles, @@ -75,10 +76,8 @@ if MCP_AVAILABLE: SpecialMCPServerName, UpdateMCPServerRequest, UserAPIKeyAuth, - UserMCPManagementMode, ) 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 from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.types.mcp import MCPCredentials @@ -303,20 +302,6 @@ if MCP_AVAILABLE: return {"access_groups": access_groups_list} ## FastAPI Routes - def _get_user_mcp_management_mode() -> UserMCPManagementMode: - proxy_general_settings: dict = {} - try: - from litellm.proxy.proxy_server import ( - general_settings as proxy_general_settings, - ) - except Exception: - pass - - mode = proxy_general_settings.get("user_mcp_management_mode") - if mode == "view_all": - return "view_all" - return "restricted" - @router.get( "/server", description="Returns the mcp server list with associated teams", @@ -334,26 +319,18 @@ if MCP_AVAILABLE: ``` """ - user_mcp_management_mode = _get_user_mcp_management_mode() + auth_contexts = await build_effective_auth_contexts(user_api_key_dict) - if user_mcp_management_mode == "view_all": - servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered() - redacted_mcp_servers = _redact_mcp_credentials_list(servers) - else: - auth_contexts = await build_effective_auth_contexts(user_api_key_dict) - - aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {} - for auth_context in auth_contexts: - servers = await global_mcp_server_manager.get_all_allowed_mcp_servers( - user_api_key_auth=auth_context - ) - for server in servers: - if server.server_id not in aggregated_servers: - aggregated_servers[server.server_id] = server - - redacted_mcp_servers = _redact_mcp_credentials_list( - aggregated_servers.values() + aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {} + for auth_context in auth_contexts: + servers = await global_mcp_server_manager.get_all_allowed_mcp_servers( + user_api_key_auth=auth_context ) + for server in servers: + if server.server_id not in aggregated_servers: + aggregated_servers[server.server_id] = server + + redacted_mcp_servers = _redact_mcp_credentials_list(aggregated_servers.values()) # augment the mcp servers with public status if litellm.public_mcp_servers is not None: @@ -395,17 +372,6 @@ if MCP_AVAILABLE: --header 'Authorization: Bearer your_api_key_here' ``` """ - user_mcp_management_mode = _get_user_mcp_management_mode() - - if user_mcp_management_mode == "view_all": - servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered( - server_ids=server_ids - ) - return [ - {"server_id": server.server_id, "status": server.status} - for server in servers - ] - auth_contexts = await build_effective_auth_contexts(user_api_key_dict) server_status_map: Dict[ diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index ec1bc5497bd..623e8408862 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -698,88 +698,6 @@ async def get_response_input_items( ) -@router.post( - "/v1/responses/compact", - dependencies=[Depends(user_api_key_auth)], - tags=["responses"], -) -@router.post( - "/responses/compact", - dependencies=[Depends(user_api_key_auth)], - tags=["responses"], -) -@router.post( - "/openai/v1/responses/compact", - dependencies=[Depends(user_api_key_auth)], - tags=["responses"], -) -async def compact_response( - request: Request, - fastapi_response: Response, - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), -): - """ - Compact a response by running a compaction pass over a conversation. - - Returns encrypted, opaque items that can be used to reduce context size. - - Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/compact - - ```bash - curl -X POST http://localhost:4000/v1/responses/compact \ - -H "Content-Type: application/json" \ - -H "Authorization: Bearer sk-1234" \ - -d '{ - "model": "gpt-4o", - "input": [{"role": "user", "content": "Hello"}] - }' - ``` - """ - from litellm.proxy.proxy_server import ( - _read_request_body, - general_settings, - llm_router, - proxy_config, - proxy_logging_obj, - select_data_generator, - user_api_base, - user_max_tokens, - user_model, - user_request_timeout, - user_temperature, - version, - ) - - data = await _read_request_body(request=request) - processor = ProxyBaseLLMRequestProcessing(data=data) - try: - return await processor.base_process_llm_request( - request=request, - fastapi_response=fastapi_response, - user_api_key_dict=user_api_key_dict, - route_type="acompact_responses", - proxy_logging_obj=proxy_logging_obj, - llm_router=llm_router, - general_settings=general_settings, - proxy_config=proxy_config, - select_data_generator=select_data_generator, - model=None, - user_model=user_model, - user_temperature=user_temperature, - user_request_timeout=user_request_timeout, - user_max_tokens=user_max_tokens, - user_api_base=user_api_base, - version=version, - ) - except Exception as e: - raise await processor._handle_llm_api_exception( - e=e, - user_api_key_dict=user_api_key_dict, - proxy_logging_obj=proxy_logging_obj, - version=version, - ) - - @router.post( "/v1/responses/{response_id}/cancel", dependencies=[Depends(user_api_key_auth)], diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index a321e25a9a5..fd00cfc1c0a 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -25,7 +25,6 @@ ROUTE_ENDPOINT_MAPPING = { "alist_input_items": "/responses/{response_id}/input_items", "aimage_edit": "/images/edits", "acancel_responses": "/responses/{response_id}/cancel", - "acompact_responses": "/responses/compact", "aocr": "/ocr", "asearch": "/search", "avideo_generation": "/videos", @@ -117,7 +116,6 @@ async def route_request( "aget_responses", "adelete_responses", "acancel_responses", - "acompact_responses", "acreate_response_reply", "alist_input_items", "_arealtime", # private function for realtime API diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 8177b177fe6..e837346df23 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -1361,205 +1361,3 @@ def cancel_responses( completion_kwargs=local_vars, extra_kwargs=kwargs, ) - - -@client -async def acompact_responses( - input: Union[str, ResponseInputParam], - model: str, - instructions: Optional[str] = None, - previous_response_id: Optional[str] = None, - # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. - # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, - **kwargs, -) -> ResponsesAPIResponse: - """ - Async version of the POST Compact Responses API - - POST /v1/responses/compact endpoint in the responses API - - Runs a compaction pass over a conversation, returning encrypted, opaque items. - """ - local_vars = locals() - try: - loop = asyncio.get_event_loop() - kwargs["acompact_responses"] = True - - # get custom llm provider so we can use this for mapping exceptions - if custom_llm_provider is None: - _, custom_llm_provider, _, _ = litellm.get_llm_provider( - model=model, api_base=local_vars.get("base_url", None) - ) - - func = partial( - compact_responses, - input=input, - model=model, - instructions=instructions, - previous_response_id=previous_response_id, - extra_headers=extra_headers, - extra_query=extra_query, - extra_body=extra_body, - timeout=timeout, - custom_llm_provider=custom_llm_provider, - **kwargs, - ) - - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - init_response = await loop.run_in_executor(None, func_with_context) - - if asyncio.iscoroutine(init_response): - response = await init_response - else: - response = init_response - - # Update the responses_api_response_id with the model_id - if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( - responses_api_response=response, - litellm_metadata=kwargs.get("litellm_metadata", {}), - custom_llm_provider=custom_llm_provider, - ) - - return response - except Exception as e: - raise litellm.exception_type( - model=model, - custom_llm_provider=custom_llm_provider, - original_exception=e, - completion_kwargs=local_vars, - extra_kwargs=kwargs, - ) - - -@client -def compact_responses( - input: Union[str, ResponseInputParam], - model: str, - instructions: Optional[str] = None, - previous_response_id: Optional[str] = None, - # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. - # The extra values given here take precedence over values defined on the client or passed to this method. - extra_headers: Optional[Dict[str, Any]] = None, - extra_query: Optional[Dict[str, Any]] = None, - extra_body: Optional[Dict[str, Any]] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - # LiteLLM specific params, - custom_llm_provider: Optional[str] = None, - **kwargs, -) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: - """ - Synchronous version of the POST Compact Responses API - - POST /v1/responses/compact endpoint in the responses API - - Runs a compaction pass over a conversation, returning encrypted, opaque items. - """ - local_vars = locals() - try: - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) - _is_async = kwargs.pop("acompact_responses", False) is True - - # get llm provider logic - litellm_params = GenericLiteLLMParams(**kwargs) - - ( - model, - custom_llm_provider, - dynamic_api_key, - dynamic_api_base, - ) = litellm.get_llm_provider( - model=model, - custom_llm_provider=custom_llm_provider, - api_base=litellm_params.api_base, - api_key=litellm_params.api_key, - ) - - if custom_llm_provider is None: - raise ValueError("custom_llm_provider is required but passed as None") - - # get provider config - responses_api_provider_config: Optional[BaseResponsesAPIConfig] = ( - ProviderConfigManager.get_provider_responses_api_config( - model=model, - provider=litellm.LlmProviders(custom_llm_provider), - ) - ) - - if responses_api_provider_config is None: - raise ValueError( - f"COMPACT responses is not supported for {custom_llm_provider}" - ) - - local_vars.update(kwargs) - - # Build optional params for compact endpoint - response_api_optional_params: ResponsesAPIOptionalRequestParams = ( - ResponsesAPIRequestUtils.get_requested_response_api_optional_param( - local_vars - ) - ) - - # Get optional parameters for the responses API - responses_api_request_params: Dict = ( - ResponsesAPIRequestUtils.get_optional_params_responses_api( - model=model, - responses_api_provider_config=responses_api_provider_config, - response_api_optional_params=response_api_optional_params, - allowed_openai_params=None, - ) - ) - - # Pre Call logging - litellm_logging_obj.update_environment_variables( - model=model, - optional_params=dict(responses_api_request_params), - litellm_params={ - **responses_api_request_params, - "litellm_call_id": litellm_call_id, - }, - custom_llm_provider=custom_llm_provider, - ) - - # Call the handler with _is_async flag instead of directly calling the async handler - response = base_llm_http_handler.compact_response_api_handler( - model=model, - input=input, - responses_api_provider_config=responses_api_provider_config, - response_api_optional_request_params=responses_api_request_params, - litellm_params=litellm_params, - logging_obj=litellm_logging_obj, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - extra_body=extra_body, - timeout=timeout or request_timeout, - _is_async=_is_async, - client=kwargs.get("client"), - shared_session=kwargs.get("shared_session"), - ) - - # Update the responses_api_response_id with the model_id - if isinstance(response, ResponsesAPIResponse): - response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( - responses_api_response=response, - litellm_metadata=kwargs.get("litellm_metadata", {}), - custom_llm_provider=custom_llm_provider, - ) - - return response - except Exception as e: - raise litellm.exception_type( - model=model, - custom_llm_provider=custom_llm_provider, - original_exception=e, - completion_kwargs=local_vars, - extra_kwargs=kwargs, - ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 0b838f916e2..0407776029d 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1,6 +1,5 @@ import asyncio import json -import traceback from datetime import datetime from typing import Any, Dict, Optional @@ -12,9 +11,6 @@ from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base -from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( - update_response_metadata, -) from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.utils import ResponsesAPIRequestUtils @@ -26,8 +22,7 @@ from litellm.types.llms.openai import ( ResponsesAPIStreamEvents, ResponsesAPIStreamingResponse, ) -from litellm.types.utils import CallTypes -from litellm.utils import CustomStreamWrapper, async_post_call_success_deployment_hook +from litellm.utils import CustomStreamWrapper class BaseResponsesAPIStreamingIterator: @@ -45,8 +40,6 @@ class BaseResponsesAPIStreamingIterator: logging_obj: LiteLLMLoggingObj, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, ): self.response = response self.model = model @@ -54,25 +47,21 @@ class BaseResponsesAPIStreamingIterator: self.finished = False self.responses_api_provider_config = responses_api_provider_config self.completed_response: Optional[ResponsesAPIStreamingResponse] = None - self.start_time = getattr(logging_obj, "start_time", datetime.now()) + self.start_time = datetime.now() - # track request context for hooks + # set request kwargs self.litellm_metadata = litellm_metadata self.custom_llm_provider = custom_llm_provider - self.request_data: Dict[str, Any] = request_data or {} - self.call_type: Optional[str] = call_type # set hidden params for response headers (e.g., x-litellm-model-id) - # This matches the stream wrapper in litellm/litellm_core_utils/streaming_handler.py + # This matches ths stream wrapper in litellm/litellm_core_utils/streaming_handler.py _api_base = get_api_base( model=model or "", optional_params=self.logging_obj.model_call_details.get( "litellm_params", {} ), ) - _model_info: Dict = ( - litellm_metadata.get("model_info", {}) if litellm_metadata else {} - ) + _model_info: Dict = litellm_metadata.get("model_info", {}) if litellm_metadata else {} self._hidden_params = { "model_id": _model_info.get("id", None), "api_base": _api_base, @@ -113,21 +102,13 @@ class BaseResponsesAPIStreamingIterator: # if "response" in parsed_chunk, then encode litellm specific information like custom_llm_provider response_object = getattr(openai_responses_api_chunk, "response", None) if response_object: - response = ( - ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( - responses_api_response=response_object, - litellm_metadata=self.litellm_metadata, - custom_llm_provider=self.custom_llm_provider, - ) + response = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( + responses_api_response=response_object, + litellm_metadata=self.litellm_metadata, + custom_llm_provider=self.custom_llm_provider, ) setattr(openai_responses_api_chunk, "response", response) - # Allow callbacks to modify chunk before returning - openai_responses_api_chunk = run_async_function( - async_function=self._call_post_streaming_deployment_hook, - chunk=openai_responses_api_chunk, - ) - # Store the completed response if ( openai_responses_api_chunk @@ -168,159 +149,11 @@ class BaseResponsesAPIStreamingIterator: except json.JSONDecodeError: # If we can't parse the chunk, continue return None - except Exception as e: - # Ensure failures trigger failure hooks - self._handle_failure(e) - raise def _handle_logging_completed_response(self): """Base implementation - should be overridden by subclasses""" pass - async def _call_post_streaming_deployment_hook(self, chunk): - """ - Allow callbacks to modify streaming chunks before returning (parity with chat). - """ - try: - # Align with chat pipeline: use logging_obj model_call_details + call_type - typed_call_type: Optional[CallTypes] = None - if self.call_type is not None: - try: - typed_call_type = CallTypes(self.call_type) - except ValueError: - typed_call_type = None - if typed_call_type is None: - try: - typed_call_type = CallTypes(getattr(self.logging_obj, "call_type", None)) - except Exception: - typed_call_type = None - - request_data = self.request_data or getattr( - self.logging_obj, "model_call_details", {} - ) - callbacks = getattr(litellm, "callbacks", None) or [] - hooks_ran = False - for callback in callbacks: - if hasattr(callback, "async_post_call_streaming_deployment_hook"): - hooks_ran = True - result = await callback.async_post_call_streaming_deployment_hook( - request_data=request_data, - response_chunk=chunk, - call_type=typed_call_type, - ) - if result is not None: - chunk = result - if hooks_ran: - setattr(chunk, "_post_streaming_hooks_ran", True) - return chunk - except Exception: - return chunk - - async def call_post_streaming_hooks_for_testing(self, chunk): - """ - Helper to invoke streaming deployment hooks explicitly (used in tests). - """ - return await self._call_post_streaming_deployment_hook(chunk) - - def _run_post_success_hooks(self, end_time: datetime): - """ - Run post-call deployment hooks and update metadata similar to chat pipeline. - """ - if self.completed_response is None: - return - - request_payload: Dict[str, Any] = {} - if isinstance(self.request_data, dict): - request_payload.update(self.request_data) - try: - if hasattr(self.logging_obj, "model_call_details"): - request_payload.update(self.logging_obj.model_call_details) - except Exception: - pass - if "litellm_params" not in request_payload: - try: - request_payload["litellm_params"] = getattr( - self.logging_obj, "model_call_details", {} - ).get("litellm_params", {}) - except Exception: - request_payload["litellm_params"] = {} - - try: - update_response_metadata( - result=self.completed_response, - logging_obj=self.logging_obj, - model=self.model, - kwargs=request_payload, - start_time=self.start_time, - end_time=end_time, - ) - except Exception: - # Non-blocking - pass - - try: - typed_call_type: Optional[CallTypes] = None - if self.call_type is not None: - try: - typed_call_type = CallTypes(self.call_type) - except ValueError: - typed_call_type = None - except Exception: - typed_call_type = None - if typed_call_type is None: - try: - typed_call_type = CallTypes.responses - except Exception: - typed_call_type = None - - try: - # Call synchronously; async hook will be executed via asyncio.run in a new loop - run_async_function( - async_function=async_post_call_success_deployment_hook, - request_data=request_payload, - response=self.completed_response, - call_type=typed_call_type, - ) - except Exception: - pass - - def _handle_failure(self, exception: Exception): - """ - Trigger failure handlers before bubbling the exception. - """ - traceback_exception = traceback.format_exc() - try: - run_async_function( - async_function=self.logging_obj.async_failure_handler, - exception=exception, - traceback_exception=traceback_exception, - start_time=self.start_time, - end_time=datetime.now(), - ) - except Exception: - pass - - try: - executor.submit( - self.logging_obj.failure_handler, - exception, - traceback_exception, - self.start_time, - datetime.now(), - ) - except Exception: - pass - - -async def call_post_streaming_hooks_for_testing(iterator, chunk): - """ - Module-level helper for tests to ensure hooks can be invoked even if the iterator is wrapped. - """ - hook_fn = getattr(iterator, "_call_post_streaming_deployment_hook", None) - if hook_fn is None: - return chunk - return await hook_fn(chunk) - class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): """ @@ -335,8 +168,6 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): logging_obj: LiteLLMLoggingObj, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, ): super().__init__( response, @@ -345,8 +176,6 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): logging_obj, litellm_metadata, custom_llm_provider, - request_data, - call_type, ) self.stream_iterator = response.aiter_lines() @@ -374,21 +203,16 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): except httpx.HTTPError as e: # Handle HTTP errors self.finished = True - self._handle_failure(e) - raise e - except Exception as e: - self.finished = True - self._handle_failure(e) raise e def _handle_logging_completed_response(self): """Handle logging for completed responses in async context""" # Create a deep copy for logging to avoid modifying the response object that will be returned to the user - # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) + # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) # to chat completion format (prompt_tokens/completion_tokens) for internal logging import copy logging_response = copy.deepcopy(self.completed_response) - + asyncio.create_task( self.logging_obj.async_success_handler( result=logging_response, @@ -405,7 +229,6 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): start_time=self.start_time, end_time=datetime.now(), ) - self._run_post_success_hooks(end_time=datetime.now()) class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -421,8 +244,6 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): logging_obj: LiteLLMLoggingObj, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, ): super().__init__( response, @@ -431,8 +252,6 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): logging_obj, litellm_metadata, custom_llm_provider, - request_data, - call_type, ) self.stream_iterator = response.iter_lines() @@ -460,21 +279,16 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): except httpx.HTTPError as e: # Handle HTTP errors self.finished = True - self._handle_failure(e) - raise e - except Exception as e: - self.finished = True - self._handle_failure(e) raise e def _handle_logging_completed_response(self): """Handle logging for completed responses in sync context""" # Create a deep copy for logging to avoid modifying the response object that will be returned to the user - # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) + # The logging handlers may transform usage from Responses API format (input_tokens/output_tokens) # to chat completion format (prompt_tokens/completion_tokens) for internal logging import copy logging_response = copy.deepcopy(self.completed_response) - + run_async_function( async_function=self.logging_obj.async_success_handler, result=logging_response, @@ -490,7 +304,6 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): start_time=self.start_time, end_time=datetime.now(), ) - self._run_post_success_hooks(end_time=datetime.now()) class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -511,8 +324,6 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): logging_obj: LiteLLMLoggingObj, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, - request_data: Optional[Dict[str, Any]] = None, - call_type: Optional[str] = None, ): super().__init__( response=response, @@ -521,8 +332,6 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): logging_obj=logging_obj, litellm_metadata=litellm_metadata, custom_llm_provider=custom_llm_provider, - request_data=request_data, - call_type=call_type, ) # one-time transform diff --git a/litellm/router.py b/litellm/router.py index d980b5f74d8..6821ab9e6c6 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -713,23 +713,6 @@ class Router: self, routing_strategy: Union[RoutingStrategy, str], routing_strategy_args: dict ): verbose_router_logger.info(f"Routing strategy: {routing_strategy}") - - # Validate routing_strategy value to fail fast with helpful error - # See: https://github.com/BerriAI/litellm/issues/11330 - # Derive valid strategies from RoutingStrategy enum + "simple-shuffle" (default, not in enum) - valid_strategy_strings = ["simple-shuffle"] + [s.value for s in RoutingStrategy] - - if routing_strategy is not None: - is_valid_string = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings - is_valid_enum = isinstance(routing_strategy, RoutingStrategy) - if not is_valid_string and not is_valid_enum: - raise ValueError( - f"Invalid routing_strategy: '{routing_strategy}'. " - f"Valid options: {valid_strategy_strings}. " - f"Check 'router_settings.routing_strategy' in your config.yaml " - f"or the 'routing_strategy' parameter if using the Router SDK directly." - ) - if ( routing_strategy == RoutingStrategy.LEAST_BUSY.value or routing_strategy == RoutingStrategy.LEAST_BUSY @@ -829,9 +812,6 @@ class Router: self.acancel_responses = self.factory_function( litellm.acancel_responses, call_type="acancel_responses" ) - self.acompact_responses = self.factory_function( - litellm.acompact_responses, call_type="acompact_responses" - ) self.adelete_responses = self.factory_function( litellm.adelete_responses, call_type="adelete_responses" ) @@ -3944,7 +3924,6 @@ class Router: "anthropic_messages", "aresponses", "acancel_responses", - "acompact_responses", "responses", "aget_responses", "adelete_responses", @@ -4173,7 +4152,6 @@ class Router: elif call_type in ( "aget_responses", "acancel_responses", - "acompact_responses", "adelete_responses", "alist_input_items", ): diff --git a/litellm/types/integrations/langsmith.py b/litellm/types/integrations/langsmith.py index 9c026a117fd..23f760ecf32 100644 --- a/litellm/types/integrations/langsmith.py +++ b/litellm/types/integrations/langsmith.py @@ -31,7 +31,6 @@ class LangsmithCredentialsObject(TypedDict): LANGSMITH_API_KEY: Optional[str] LANGSMITH_PROJECT: Optional[str] LANGSMITH_BASE_URL: str - LANGSMITH_TENANT_ID: Optional[str] class LangsmithQueueObject(TypedDict): @@ -53,7 +52,6 @@ class CredentialsKey(NamedTuple): api_key: str project: str base_url: str - tenant_id: Optional[str] @dataclass diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 784c8403c3f..3eec67d9d26 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2677,7 +2677,6 @@ class StandardCallbackDynamicParams(TypedDict, total=False): langsmith_project: Optional[str] langsmith_base_url: Optional[str] langsmith_sampling_rate: Optional[float] - langsmith_tenant_id: Optional[str] # Humanloop dynamic params humanloop_api_key: Optional[str] @@ -2947,7 +2946,6 @@ class LlmProviders(str, Enum): MISTRAL = "mistral" MILVUS = "milvus" GROQ = "groq" - GIGACHAT = "gigachat" NVIDIA_NIM = "nvidia_nim" CEREBRAS = "cerebras" AI21_CHAT = "ai21_chat" diff --git a/litellm/utils.py b/litellm/utils.py index fbbaa94f7a1..6b9aeca0934 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7521,8 +7521,6 @@ class ProviderConfigManager: return litellm.CompactifAIChatConfig() elif litellm.LlmProviders.GITHUB_COPILOT == provider: return litellm.GithubCopilotConfig() - elif litellm.LlmProviders.GIGACHAT == provider: - return litellm.GigaChatConfig() elif litellm.LlmProviders.RAGFLOW == provider: return litellm.RAGFlowConfig() elif ( @@ -7718,8 +7716,6 @@ class ProviderConfigManager: return litellm.CometAPIEmbeddingConfig() elif litellm.LlmProviders.GITHUB_COPILOT == provider: return litellm.GithubCopilotEmbeddingConfig() - elif litellm.LlmProviders.GIGACHAT == provider: - return litellm.GigaChatEmbeddingConfig() elif litellm.LlmProviders.SAGEMAKER == provider: from litellm.llms.sagemaker.embedding.transformation import ( SagemakerEmbeddingConfig, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c7a2f60856d..e823dd5dc6b 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -15831,68 +15831,6 @@ "max_tokens": 8191, "mode": "embedding" }, - "gigachat/GigaChat-2-Lite": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 0.0, - "supports_function_calling": true, - "supports_system_messages": true - }, - "gigachat/GigaChat-2-Max": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 0.0, - "supports_function_calling": true, - "supports_system_messages": true, - "supports_vision": true - }, - "gigachat/GigaChat-2-Pro": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 0.0, - "supports_function_calling": true, - "supports_system_messages": true, - "supports_vision": true - }, - "gigachat/Embeddings": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 512, - "max_tokens": 512, - "mode": "embedding", - "output_cost_per_token": 0.0, - "output_vector_size": 1024 - }, - "gigachat/Embeddings-2": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 512, - "max_tokens": 512, - "mode": "embedding", - "output_cost_per_token": 0.0, - "output_vector_size": 1024 - }, - "gigachat/EmbeddingsGigaR": { - "input_cost_per_token": 0.0, - "litellm_provider": "gigachat", - "max_input_tokens": 4096, - "max_tokens": 4096, - "mode": "embedding", - "output_cost_per_token": 0.0, - "output_vector_size": 2560 - }, "google.gemma-3-12b-it": { "input_cost_per_token": 9e-08, "litellm_provider": "bedrock_converse", @@ -32154,4 +32092,3 @@ "mode": "chat" } } - diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index bc5dea7b97c..45ee47c01bc 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -28,8 +28,7 @@ "list_container_files": "Supports GET /containers/{id}/files endpoint", "retrieve_container_file": "Supports GET /containers/{id}/files/{file_id} endpoint", "retrieve_container_file_content": "Supports GET /containers/{id}/files/{file_id}/content endpoint", - "delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint", - "compact": "Supports /responses/compact endpoint" + "delete_container_file": "Supports DELETE /containers/{id}/files/{file_id} endpoint" } } }, @@ -1520,7 +1519,6 @@ "retrieve_container_file": true, "retrieve_container_file_content": true, "delete_container_file": true, - "compact": true, "a2a": true, "interactions": true } diff --git a/requirements.txt b/requirements.txt index 249b899b86b..06a7c17336c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -20,7 +20,7 @@ google-cloud-aiplatform==1.47.0 # for vertex ai calls google-cloud-iam==2.19.1 # for GCP IAM Redis authentication google-genai==1.22.0 anthropic[vertex]==0.54.0 -mcp==1.25.0 ; python_version >= "3.10" # for MCP server +mcp==1.23.0 ; python_version >= "3.10" # for MCP server google-generativeai==0.5.0 # for vertex ai calls async_generator==1.10.0 # for async ollama calls langfuse==2.59.7 # for langfuse self-hosted logging diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py index 60d4f479733..9a96919da87 100644 --- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py @@ -61,19 +61,13 @@ async def test_bedrock_apply_guardrail_blocked(): guardrailVersion="DRAFT", ) - # Mock the make_bedrock_api_request method to raise an exception for blocked content + # Mock the make_bedrock_api_request method with patch.object( - guardrail, "make_bedrock_api_request", new_callable=AsyncMock + guardrail, "make_bedrock_api_request", new_callable=AsyncMock ) as mock_api_request: - # Mock the method to raise an HTTPException as it would for blocked content - from fastapi import HTTPException - mock_api_request.side_effect = HTTPException( - status_code=400, - detail={ - "error": "Violated guardrail policy", - "bedrock_guardrail_response": "", - }, - ) + # Mock a blocked response from Bedrock + mock_response = {"action": "BLOCKED", "reason": "Content violates policy"} + mock_api_request.return_value = mock_response # Test the apply_guardrail method should raise an exception with pytest.raises(Exception) as exc_info: @@ -83,9 +77,8 @@ async def test_bedrock_apply_guardrail_blocked(): input_type="request", ) - # The apply_guardrail method wraps the original exception in a generic Exception - assert "Bedrock guardrail failed:" in str(exc_info.value) - assert "Violated guardrail policy" in str(exc_info.value) + assert "Content blocked by Bedrock guardrail" in str(exc_info.value) + assert "Content violates policy" in str(exc_info.value) @pytest.mark.asyncio @@ -260,15 +253,7 @@ async def test_bedrock_apply_guardrail_filters_request_messages_when_flag_enable with patch.object( guardrail, "make_bedrock_api_request", new_callable=AsyncMock ) as mock_api: - # Mock the method to raise an HTTPException as it would for blocked content - from fastapi import HTTPException - mock_api.side_effect = HTTPException( - status_code=400, - detail={ - "error": "Violated guardrail policy", - "bedrock_guardrail_response": "policy", - }, - ) + mock_api.return_value = {"action": "BLOCKED", "reason": "policy"} with pytest.raises(Exception, match="policy") as exc_info: await guardrail.apply_guardrail( @@ -280,8 +265,7 @@ async def test_bedrock_apply_guardrail_filters_request_messages_when_flag_enable assert mock_api.called _, kwargs = mock_api.call_args assert kwargs["messages"] == [request_messages[-1]] - # The apply_guardrail method wraps the original exception in a generic Exception - assert "Bedrock guardrail failed:" in str(exc_info.value) + assert "Content blocked by Bedrock guardrail" in str(exc_info.value) def test_bedrock_guardrail_filters_latest_user_message_when_enabled(): diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py index c981ccd8dac..5714cd5c487 100644 --- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py +++ b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py @@ -1,14 +1,11 @@ import os import sys -import time -from unittest.mock import Mock, patch sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -from litellm_proxy_extras.utils import ProxyExtrasDBManager, MigrationLockManager - +from litellm_proxy_extras.utils import ProxyExtrasDBManager def test_custom_prisma_dir(monkeypatch): @@ -30,279 +27,101 @@ def test_custom_prisma_dir(monkeypatch): assert os.path.exists(migrations_dir) -class TestMigrationLockManager: - """Test cases for MigrationLockManager""" +class TestPermissionErrorDetection: + """Test cases for permission error detection in Prisma migrations""" - def test_acquire_lock_without_redis(self): - """Test lock acquisition when Redis is not available""" - lock_manager = MigrationLockManager() - result = lock_manager.acquire_lock() - assert result is True # Should return True when Redis is not available - assert lock_manager.lock_acquired is True # Redis 없을 때도 lock_acquired는 True + def test_is_permission_error_postgres_42501(self): + """Test detection of PostgreSQL 42501 error code (insufficient privilege)""" + error_message = "Database error code: 42501 - permission denied for table users" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True - def test_acquire_lock_with_redis_success(self): - """Test successful lock acquisition with Redis""" - mock_redis = Mock() - mock_redis.set_cache.return_value = True - lock_manager = MigrationLockManager(mock_redis) + def test_is_permission_error_must_be_owner(self): + """Test detection of 'must be owner of table' error""" + error_message = "ERROR: must be owner of table my_table" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True - result = lock_manager.acquire_lock() + def test_is_permission_error_permission_denied_schema(self): + """Test detection of 'permission denied for schema' error""" + error_message = "permission denied for schema public" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True - assert result is True - assert lock_manager.lock_acquired is True - mock_redis.set_cache.assert_called_once() + def test_is_permission_error_permission_denied_table(self): + """Test detection of 'permission denied for table' error""" + error_message = "permission denied for table my_table" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True - def test_acquire_lock_with_redis_failure(self): - """Test failed lock acquisition with Redis""" - mock_redis = Mock() - mock_redis.set_cache.return_value = False - lock_manager = MigrationLockManager(mock_redis) + def test_is_permission_error_must_be_owner_schema(self): + """Test detection of 'must be owner of schema' error""" + error_message = "must be owner of schema public" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True - result = lock_manager.acquire_lock() + def test_is_permission_error_case_insensitive(self): + """Test that permission error detection is case insensitive""" + error_message = "PERMISSION DENIED FOR TABLE my_table" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True - assert result is False - assert lock_manager.lock_acquired is False - mock_redis.set_cache.assert_called_once() - - def test_acquire_lock_with_redis_exception(self): - """Test lock acquisition with Redis exception""" - mock_redis = Mock() - mock_redis.set_cache.side_effect = Exception("Redis error") - lock_manager = MigrationLockManager(mock_redis) - - result = lock_manager.acquire_lock() - - assert result is False - assert lock_manager.lock_acquired is False - - def test_wait_for_lock_release_success(self): - """Test successful waiting for lock release""" - mock_redis = Mock() - # First call returns False (lock held), second call returns True (lock acquired) - mock_redis.set_cache.side_effect = [False, True] - lock_manager = MigrationLockManager(mock_redis) - - result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=1) - - assert result is True - assert lock_manager.lock_acquired is True - assert mock_redis.set_cache.call_count == 2 - - def test_wait_for_lock_release_timeout(self): - """Test timeout while waiting for lock release""" - mock_redis = Mock() - mock_redis.set_cache.return_value = False # Lock always held - lock_manager = MigrationLockManager(mock_redis) - - result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=0.2) - - assert result is False - assert lock_manager.lock_acquired is False - - def test_release_lock_not_acquired(self): - """Test releasing lock when not acquired""" - mock_redis = Mock() - lock_manager = MigrationLockManager(mock_redis) - - lock_manager.release_lock() - - mock_redis.get_cache.assert_not_called() - mock_redis.delete_cache.assert_not_called() - - def test_release_lock_success(self): - """Test successful lock release""" - mock_redis = Mock() - mock_redis.get_cache.return_value = "pod_123_456" - lock_manager = MigrationLockManager(mock_redis) - lock_manager.pod_id = "pod_123_456" - lock_manager.lock_acquired = True - - lock_manager.release_lock() - - mock_redis.get_cache.assert_called_once() - mock_redis.delete_cache.assert_called_once() - assert lock_manager.lock_acquired is False - - def test_release_lock_wrong_owner(self): - """Test releasing lock when not the owner""" - mock_redis = Mock() - mock_redis.get_cache.return_value = "pod_999_999" # Different pod - lock_manager = MigrationLockManager(mock_redis) - lock_manager.pod_id = "pod_123_456" - lock_manager.lock_acquired = True - - lock_manager.release_lock() - - mock_redis.get_cache.assert_called_once() - mock_redis.delete_cache.assert_not_called() - assert lock_manager.lock_acquired is False - - def test_context_manager(self): - """Test MigrationLockManager as context manager""" - mock_redis = Mock() - mock_redis.set_cache.return_value = True - # Mock get_cache to return the same pod_id for successful release - mock_redis.get_cache.return_value = "pod_123_456" - - lock_manager = MigrationLockManager(mock_redis) - lock_manager.pod_id = "pod_123_456" # Set consistent pod_id - - with lock_manager: - assert lock_manager.lock_acquired is True - - # Should call release_lock when exiting context - mock_redis.get_cache.assert_called_once() - mock_redis.delete_cache.assert_called_once() + def test_is_permission_error_negative(self): + """Test that non-permission errors are not detected as permission errors""" + error_message = "column 'id' already exists" + assert ProxyExtrasDBManager._is_permission_error(error_message) is False -class TestProxyExtrasDBManagerMigrationLock: - """Test cases for ProxyExtrasDBManager with migration locking""" +class TestIdempotentErrorDetection: + """Test cases for idempotent error detection in Prisma migrations""" - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations") - @patch("litellm_proxy_extras.utils.subprocess.run") - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir") - @patch("os.chdir") - def test_setup_database_with_redis_lock_success( - self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations - ): - """Test successful database setup with Redis lock""" - # Setup mocks - mock_get_prisma_dir.return_value = "/test/prisma" - mock_subprocess.return_value = Mock(stdout="Migration completed", stderr="") - mock_resolve_migrations.return_value = None + def test_is_idempotent_error_already_exists(self): + """Test detection of generic 'already exists' error""" + error_message = "object already exists" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True - # Mock Redis cache - mock_redis = Mock() - mock_redis.set_cache.return_value = True # Lock acquired successfully - mock_redis.get_cache.return_value = "pod_123_456" + def test_is_idempotent_error_column_already_exists(self): + """Test detection of 'column already exists' error""" + error_message = "column 'email' already exists" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True - # Set DATABASE_URL - with patch.dict( - os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"} - ): - result = ProxyExtrasDBManager.setup_database( - use_migrate=True, redis_cache=mock_redis - ) + def test_is_idempotent_error_duplicate_key(self): + """Test detection of duplicate key violation error""" + error_message = "duplicate key value violates unique constraint" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True - assert result is True - # set_cache is called once in acquire_lock (__enter__ calls acquire_lock) - assert mock_redis.set_cache.call_count == 1 - mock_subprocess.assert_called_once() + def test_is_idempotent_error_relation_already_exists(self): + """Test detection of 'relation already exists' error""" + error_message = "relation 'users_pkey' already exists" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations") - @patch("litellm_proxy_extras.utils.subprocess.run") - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir") - @patch("os.chdir") - def test_setup_database_with_redis_lock_wait_and_skip( - self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations - ): - """Test database setup when lock is held by another pod, then acquired after waiting""" - # Setup mocks - mock_get_prisma_dir.return_value = "/test/prisma" - mock_resolve_migrations.return_value = None + def test_is_idempotent_error_constraint_already_exists(self): + """Test detection of 'constraint already exists' error""" + error_message = "constraint 'fk_user_id' already exists" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True - # Mock Redis cache - first call fails, second call succeeds - mock_redis = Mock() - mock_redis.set_cache.side_effect = [False, True] # First fails, then succeeds - mock_redis.get_cache.return_value = "pod_123_456" + def test_is_idempotent_error_case_insensitive(self): + """Test that idempotent error detection is case insensitive""" + error_message = "COLUMN 'ID' ALREADY EXISTS" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True - # Set DATABASE_URL - with patch.dict( - os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"} - ): - result = ProxyExtrasDBManager.setup_database( - use_migrate=True, redis_cache=mock_redis - ) + def test_is_idempotent_error_negative(self): + """Test that non-idempotent errors are not detected as idempotent errors""" + error_message = "Database error code: 42501 - permission denied" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False - assert result is True # Should return True after waiting and acquiring lock - # set_cache is called 2 times: once in __enter__, once in wait_for_lock_release - assert mock_redis.set_cache.call_count == 2 - # Proceed for case handling in case of migration failure - mock_subprocess.assert_called_once() - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._resolve_all_migrations") - @patch("litellm_proxy_extras.utils.subprocess.run") - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir") - @patch("os.chdir") - def test_setup_database_without_redis( - self, mock_chdir, mock_get_prisma_dir, mock_subprocess, mock_resolve_migrations - ): - """Test database setup without Redis cache""" - # Setup mocks - mock_get_prisma_dir.return_value = "/test/prisma" - mock_subprocess.return_value = Mock(stdout="Migration completed", stderr="") - mock_resolve_migrations.return_value = None +class TestErrorClassificationPriority: + """Test cases to ensure errors are correctly classified""" - # Set DATABASE_URL - with patch.dict( - os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"} - ): - result = ProxyExtrasDBManager.setup_database( - use_migrate=True, redis_cache=None - ) + def test_permission_error_not_classified_as_idempotent(self): + """Ensure permission errors are not mistakenly classified as idempotent""" + error_message = "Database error code: 42501 - must be owner of table users" + assert ProxyExtrasDBManager._is_permission_error(error_message) is True + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False - assert result is True - # Redis가 없을 때는 락 보호 없이 마이그레이션을 실행해야 함 - mock_subprocess.assert_called_once() + def test_idempotent_error_not_classified_as_permission(self): + """Ensure idempotent errors are not mistakenly classified as permission errors""" + error_message = "column 'created_at' already exists" + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True + assert ProxyExtrasDBManager._is_permission_error(error_message) is False - def test_setup_database_no_database_url(self): - """Test database setup without DATABASE_URL""" - with patch.dict(os.environ, {}, clear=True): - result = ProxyExtrasDBManager.setup_database( - use_migrate=True, redis_cache=None - ) - - assert result is False - - @patch("litellm_proxy_extras.utils.subprocess.run") - @patch("litellm_proxy_extras.utils.ProxyExtrasDBManager._get_prisma_dir") - @patch("os.chdir") - @patch.object( - MigrationLockManager, "LOCK_TTL_SECONDS", 1 - ) # Set short TTL for testing - def test_setup_database_lock_timeout( - self, mock_chdir, mock_get_prisma_dir, mock_subprocess - ): - """Test database setup when lock acquisition times out""" - # Setup mocks - mock_get_prisma_dir.return_value = "/test/prisma" - - # Mock Redis cache - always fails to acquire lock - mock_redis = Mock() - mock_redis.set_cache.return_value = False # Always fails - - # Set DATABASE_URL - with patch.dict( - os.environ, {"DATABASE_URL": "postgresql://test:test@localhost/test"} - ): - # Patch the wait_for_lock_release method to use shorter timeout - with patch.object( - MigrationLockManager, "wait_for_lock_release" - ) as mock_wait: - mock_wait.return_value = False # Simulate timeout - - result = ProxyExtrasDBManager.setup_database( - use_migrate=True, redis_cache=mock_redis - ) - - assert result is False # Should return False after timeout - mock_subprocess.assert_not_called() # Should not run migration - # Verify that wait_for_lock_release was called with default parameters - mock_wait.assert_called_once_with() - - def test_wait_for_lock_release_actual_timeout(self): - """Test actual timeout behavior of wait_for_lock_release with real timing""" - mock_redis = Mock() - mock_redis.set_cache.return_value = False # Always fails to acquire lock - lock_manager = MigrationLockManager(mock_redis) - - # Test with very short timeout to verify actual timeout behavior - start_time = time.time() - result = lock_manager.wait_for_lock_release(check_interval=0.1, max_wait=0.5) - end_time = time.time() - - assert result is False # Should timeout - assert end_time - start_time >= 0.5 # Should wait at least the max_wait time - assert end_time - start_time < 1.0 # But not too much longer - # Should have called set_cache multiple times during the wait - assert mock_redis.set_cache.call_count > 1 + def test_unknown_error_classified_as_neither(self): + """Ensure unknown errors are classified as neither permission nor idempotent""" + error_message = "connection timeout" + assert ProxyExtrasDBManager._is_permission_error(error_message) is False + assert ProxyExtrasDBManager._is_idempotent_error(error_message) is False 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 5f35d6837c0..7553c670774 100644 --- a/tests/llm_responses_api_testing/test_openai_responses_api.py +++ b/tests/llm_responses_api_testing/test_openai_responses_api.py @@ -1814,49 +1814,3 @@ async def test_extra_body_merges_with_request_data(extra_body_mock_response_data assert "temperature" in request_body assert "custom_field" in request_body assert request_body["custom_field"] == "custom_value" - - -@pytest.mark.asyncio -@pytest.mark.parametrize("sync_mode", [True, False]) -async def test_openai_compact_responses_api(sync_mode): - """ - Test the compact_responses API for OpenAI. - - This test verifies that the compact_responses endpoint works correctly - for compressing conversation history. - """ - litellm._turn_on_debug() - litellm.set_verbose = True - - input_messages = [ - {"role": "user", "content": "Hello, how are you?"}, - {"role": "assistant", "content": "I'm doing well, thank you for asking!"}, - {"role": "user", "content": "What is the weather like today?"}, - ] - - try: - if sync_mode: - response = litellm.compact_responses( - model="openai/gpt-4o", - input=input_messages, - instructions="Be helpful and concise", - ) - else: - response = await litellm.acompact_responses( - model="openai/gpt-4o", - input=input_messages, - instructions="Be helpful and concise", - ) - except litellm.InternalServerError: - pytest.skip("Skipping test due to InternalServerError") - except litellm.BadRequestError as e: - # compact_responses may not be available for all models/accounts - pytest.skip(f"Skipping test due to BadRequestError: {e}") - - print("compact_responses response=", json.dumps(response, indent=4, default=str)) - - # Validate response structure - assert response is not None - assert "id" in response, "Response should have an 'id' field" - assert "output" in response, "Response should have an 'output' field" - assert isinstance(response["output"], list), "Output should be a list" diff --git a/tests/llm_responses_api_testing/test_responses_hooks.py b/tests/llm_responses_api_testing/test_responses_hooks.py deleted file mode 100644 index 8c0f7dab2af..00000000000 --- a/tests/llm_responses_api_testing/test_responses_hooks.py +++ /dev/null @@ -1,165 +0,0 @@ -import asyncio -from datetime import datetime -from types import SimpleNamespace - -import httpx -import pytest - -import litellm -from litellm.integrations.custom_logger import CustomLogger -from litellm.responses import streaming_iterator as streaming_module -from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator -from litellm.types.llms.openai import ResponsesAPIStreamEvents -from litellm.types.utils import CallTypes - - -class _FakeLoggingObj: - def __init__(self): - self.success_calls = 0 - self.async_success_calls = 0 - self.failure_calls = 0 - self.async_failure_calls = 0 - self.start_time = datetime.now() - self.model_call_details = {"litellm_params": {}} - - # Signature alignment with Logging handlers - def success_handler(self, *args, **kwargs): - self.success_calls += 1 - - async def async_success_handler(self, *args, **kwargs): - self.async_success_calls += 1 - - def failure_handler(self, *args, **kwargs): - self.failure_calls += 1 - - async def async_failure_handler(self, *args, **kwargs): - self.async_failure_calls += 1 - - -@pytest.mark.asyncio -async def test_responses_streaming_triggers_hooks(monkeypatch): - """ - Ensure streaming iterator fires success + post-call hooks for responses API. - """ - hook_calls = {"post_call": 0, "metadata": 0} - seen = {} - - async def fake_post_call(request_data, response, call_type): - hook_calls["post_call"] += 1 - seen["request_data"] = request_data - seen["call_type"] = call_type - - def fake_update_metadata(**kwargs): - hook_calls["metadata"] += 1 - - monkeypatch.setattr( - streaming_module, - "async_post_call_success_deployment_hook", - fake_post_call, - ) - monkeypatch.setattr( - streaming_module, - "update_response_metadata", - fake_update_metadata, - ) - - logging_obj = _FakeLoggingObj() - - iterator = ResponsesAPIStreamingIterator( - response=httpx.Response(200), - model="test-model", - responses_api_provider_config=SimpleNamespace(), # not used in this test - logging_obj=logging_obj, - request_data={"foo": "bar", "litellm_params": {}}, - call_type=CallTypes.responses.value, - ) - - # Simulate completed streaming event - iterator.completed_response = SimpleNamespace( - type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, response=SimpleNamespace() - ) - - iterator._handle_logging_completed_response() - await asyncio.sleep(0.2) # allow async tasks to run - - assert logging_obj.success_calls == 1 - assert logging_obj.async_success_calls == 1 - assert hook_calls["post_call"] == 1 - assert hook_calls["metadata"] == 1 - assert seen["request_data"]["foo"] == "bar" - assert seen["request_data"].get("litellm_params") is not None - assert seen["call_type"] == CallTypes.responses - - -@pytest.mark.asyncio -async def test_responses_streaming_calls_post_streaming_deployment_hook(monkeypatch): - """ - Ensure per-chunk streaming deployment hook can modify chunks. - """ - - class _HookLogger(CustomLogger): - async def async_post_call_streaming_deployment_hook( - self, request_data, response_chunk, call_type - ): - response_chunk.tagged = True - return response_chunk - - # Set callbacks to our fake hook - original_callbacks = litellm.callbacks - litellm.callbacks = [_HookLogger()] - - logging_obj = _FakeLoggingObj() - - class _StubConfig: - def transform_streaming_response(self, **kwargs): - return SimpleNamespace( - type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None - ) - - iterator = ResponsesAPIStreamingIterator( - response=httpx.Response(200), - model="test-model", - responses_api_provider_config=_StubConfig(), - logging_obj=logging_obj, - request_data={"foo": "bar"}, - call_type=CallTypes.responses.value, - ) - - # Call hook helper directly to verify chunk is modified/flagged - chunk = SimpleNamespace(type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, response=None) - chunk = await streaming_module.call_post_streaming_hooks_for_testing(iterator, chunk) - assert getattr(chunk, "_post_streaming_hooks_ran", False) is True - assert getattr(chunk, "tagged", False) is True - - # reset callbacks - litellm.callbacks = original_callbacks - - -@pytest.mark.asyncio -async def test_responses_streaming_failure_triggers_failure_handlers(): - """ - If transform raises, failure handlers should be called. - """ - - class _FailConfig: - def transform_streaming_response(self, **kwargs): - raise ValueError("boom") - - logging_obj = _FakeLoggingObj() - - iterator = ResponsesAPIStreamingIterator( - response=httpx.Response(200), - model="test-model", - responses_api_provider_config=_FailConfig(), - logging_obj=logging_obj, - request_data={"foo": "bar"}, - call_type=CallTypes.responses.value, - ) - - with pytest.raises(ValueError): - iterator._process_chunk('{"delta": "chunk"}') - - # allow failure callbacks to run - await asyncio.sleep(0.2) - assert logging_obj.failure_calls >= 1 - assert logging_obj.async_failure_calls >= 1 diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index ab5709cd72d..7c849650bf6 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -385,7 +385,7 @@ def test_anthropic_tool_use(tool_type, tool_config, message_content): "computer_tool_used, prompt_caching_set, expected_beta_header", [ (True, False, True), - (False, True, False), + (False, True, True), (True, True, True), (False, False, False), ], diff --git a/tests/llm_translation/test_databricks.py b/tests/llm_translation/test_databricks.py index 3013d00288f..40fc712f2b7 100644 --- a/tests/llm_translation/test_databricks.py +++ b/tests/llm_translation/test_databricks.py @@ -15,7 +15,6 @@ import litellm from litellm.exceptions import BadRequestError from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import CustomStreamWrapper -from litellm._version import version from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest try: @@ -726,7 +725,6 @@ def test_embeddings_with_sync_http_handler(monkeypatch): headers={ "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", - "User-Agent": f"litellm/{version}", }, data=json.dumps( { @@ -769,7 +767,6 @@ def test_embeddings_with_async_http_handler(monkeypatch): headers={ "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", - "User-Agent": f"litellm/{version}", }, data=json.dumps( { @@ -826,7 +823,6 @@ def test_embeddings_uses_databricks_sdk_if_api_key_and_base_not_specified(monkey headers={ "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", - "User-Agent": f"litellm/{version}", }, data=json.dumps( { @@ -899,7 +895,6 @@ async def test_databricks_embeddings(sync_mode, monkeypatch): headers={ "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", - "User-Agent": f"litellm/{version}", }, data=json.dumps( { @@ -928,7 +923,6 @@ async def test_databricks_embeddings(sync_mode, monkeypatch): headers={ "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", - "User-Agent": f"litellm/{version}", }, data=json.dumps( { diff --git a/tests/llm_translation/test_gigachat.py b/tests/llm_translation/test_gigachat.py deleted file mode 100644 index 80bf51b4646..00000000000 --- a/tests/llm_translation/test_gigachat.py +++ /dev/null @@ -1,349 +0,0 @@ -""" -Tests for GigaChat LiteLLM Provider - -Tests message transformation, parameter handling, and response transformation. -Run with: pytest tests/llm_translation/test_gigachat.py -v -""" - -import json -import pytest -from unittest.mock import Mock, MagicMock - - -class TestGigaChatMessageTransformation: - """Tests for message transformation (OpenAI -> GigaChat format)""" - - @pytest.fixture - def config(self): - from litellm.llms.gigachat.chat.transformation import GigaChatConfig - return GigaChatConfig() - - def test_simple_user_message(self, config): - """Basic user message should pass through""" - messages = [{"role": "user", "content": "Hello"}] - result = config._transform_messages(messages) - - assert len(result) == 1 - assert result[0]["role"] == "user" - assert result[0]["content"] == "Hello" - - def test_developer_role_to_system(self, config): - """Developer role should be converted to system""" - messages = [{"role": "developer", "content": "You are helpful"}] - result = config._transform_messages(messages) - - assert result[0]["role"] == "system" - - def test_system_after_first_becomes_user(self, config): - """System message after first position should become user""" - messages = [ - {"role": "assistant", "content": "Response"}, - {"role": "system", "content": "Additional instruction"}, - ] - result = config._transform_messages(messages) - - assert result[0]["role"] == "assistant" - assert result[1]["role"] == "user" # system after first becomes user - - def test_tool_role_to_function(self, config): - """Tool role should be converted to function""" - messages = [{"role": "tool", "content": "result data"}] - result = config._transform_messages(messages) - - assert result[0]["role"] == "function" - - def test_tool_calls_to_function_call(self, config): - """tool_calls should be converted to function_call""" - messages = [{ - "role": "assistant", - "content": "", - "tool_calls": [{ - "id": "call_123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"city": "Moscow"}' - } - }] - }] - result = config._transform_messages(messages) - - assert "function_call" in result[0] - assert result[0]["function_call"]["name"] == "get_weather" - assert result[0]["function_call"]["arguments"] == {"city": "Moscow"} - assert "tool_calls" not in result[0] - - def test_none_content_becomes_empty_string(self, config): - """None content should become empty string""" - messages = [{"role": "assistant", "content": None}] - result = config._transform_messages(messages) - - assert result[0]["content"] == "" - - def test_name_field_removed(self, config): - """name field should be removed (not supported by GigaChat)""" - messages = [{"role": "user", "content": "Hi", "name": "John"}] - result = config._transform_messages(messages) - - assert "name" not in result[0] - - -class TestGigaChatCollapseUserMessages: - """Tests for collapsing consecutive user messages""" - - @pytest.fixture - def config(self): - from litellm.llms.gigachat.chat.transformation import GigaChatConfig - return GigaChatConfig() - - def test_no_collapse_single_message(self, config): - """Single message should not be changed""" - messages = [{"role": "user", "content": "Hello"}] - result = config._collapse_user_messages(messages) - - assert len(result) == 1 - assert result[0]["content"] == "Hello" - - def test_collapse_consecutive_user_messages(self, config): - """Consecutive user messages should be collapsed""" - messages = [ - {"role": "user", "content": "First"}, - {"role": "user", "content": "Second"}, - {"role": "user", "content": "Third"}, - ] - result = config._collapse_user_messages(messages) - - assert len(result) == 1 - assert "First" in result[0]["content"] - assert "Second" in result[0]["content"] - assert "Third" in result[0]["content"] - - def test_no_collapse_with_assistant_between(self, config): - """Messages with assistant between should not be collapsed""" - messages = [ - {"role": "user", "content": "First"}, - {"role": "assistant", "content": "Response"}, - {"role": "user", "content": "Second"}, - ] - result = config._collapse_user_messages(messages) - - assert len(result) == 3 - - -class TestGigaChatToolsTransformation: - """Tests for tools -> functions conversion""" - - @pytest.fixture - def config(self): - from litellm.llms.gigachat.chat.transformation import GigaChatConfig - return GigaChatConfig() - - def test_single_tool_conversion(self, config): - """Single tool should be converted correctly""" - tools = [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a city", - "parameters": { - "type": "object", - "properties": { - "city": {"type": "string"} - } - } - } - }] - result = config._convert_tools_to_functions(tools) - - assert len(result) == 1 - assert result[0]["name"] == "get_weather" - assert result[0]["description"] == "Get weather for a city" - - def test_multiple_tools_conversion(self, config): - """Multiple tools should all be converted""" - tools = [ - {"type": "function", "function": {"name": "func1", "description": "First", "parameters": {"type": "object", "properties": {}}}}, - {"type": "function", "function": {"name": "func2", "description": "Second", "parameters": {"type": "object", "properties": {}}}}, - ] - result = config._convert_tools_to_functions(tools) - - assert len(result) == 2 - assert result[0]["name"] == "func1" - assert result[1]["name"] == "func2" - - -class TestGigaChatParamsTransformation: - """Tests for parameter transformation""" - - @pytest.fixture - def config(self): - from litellm.llms.gigachat.chat.transformation import GigaChatConfig - return GigaChatConfig() - - def test_temperature_zero_becomes_top_p_zero(self, config): - """temperature=0 should become top_p=0""" - params = {"temperature": 0} - result = config.map_openai_params( - non_default_params=params, - optional_params={}, - model="GigaChat", - drop_params=False, - ) - - assert "top_p" in result - assert result["top_p"] == 0 - assert "temperature" not in result - - def test_temperature_nonzero_preserved(self, config): - """Non-zero temperature should be preserved""" - params = {"temperature": 0.7} - result = config.map_openai_params( - non_default_params=params, - optional_params={}, - model="GigaChat", - drop_params=False, - ) - - assert result["temperature"] == 0.7 - - def test_max_completion_tokens_to_max_tokens(self, config): - """max_completion_tokens should become max_tokens""" - params = {"max_completion_tokens": 100} - result = config.map_openai_params( - non_default_params=params, - optional_params={}, - model="GigaChat", - drop_params=False, - ) - - assert result["max_tokens"] == 100 - - def test_structured_output_via_json_schema(self, config): - """json_schema response_format should trigger structured output mode""" - params = { - "response_format": { - "type": "json_schema", - "json_schema": { - "name": "person", - "schema": { - "type": "object", - "properties": { - "name": {"type": "string"}, - "age": {"type": "integer"} - } - } - } - } - } - result = config.map_openai_params( - non_default_params=params, - optional_params={}, - model="GigaChat", - drop_params=False, - ) - - assert "_structured_output" in result - assert result["_structured_output"] is True - assert "function_call" in result - assert result["function_call"]["name"] == "person" - - -class TestGigaChatProviderRegistration: - """Tests for provider registration in LiteLLM""" - - def test_gigachat_in_provider_list(self): - """GigaChat should be in provider list""" - from litellm.types.utils import LlmProviders - - assert hasattr(LlmProviders, "GIGACHAT") - assert LlmProviders.GIGACHAT.value == "gigachat" - - def test_gigachat_in_chat_providers(self): - """GigaChat should be in LITELLM_CHAT_PROVIDERS""" - from litellm.constants import LITELLM_CHAT_PROVIDERS - - assert "gigachat" in LITELLM_CHAT_PROVIDERS - - def test_gigachat_key_exists(self): - """gigachat_key should be available""" - import litellm - - assert hasattr(litellm, "gigachat_key") - - def test_gigachat_config_exists(self): - """GigaChatConfig should be available""" - import litellm - - assert hasattr(litellm, "GigaChatConfig") - - -class TestGigaChatTransformRequest: - """Tests for request transformation""" - - @pytest.fixture - def config(self): - from litellm.llms.gigachat.chat.transformation import GigaChatConfig - return GigaChatConfig() - - def test_basic_request(self, config): - """Basic request should be transformed correctly""" - messages = [{"role": "user", "content": "Hello"}] - result = config.transform_request( - model="gigachat/GigaChat", - messages=messages, - optional_params={}, - litellm_params={}, - headers={}, - ) - - assert result["model"] == "GigaChat" - assert len(result["messages"]) == 1 - assert result["messages"][0]["role"] == "user" - - def test_request_with_temperature(self, config): - """Request with temperature should include it""" - messages = [{"role": "user", "content": "Hello"}] - result = config.transform_request( - model="gigachat/GigaChat", - messages=messages, - optional_params={"temperature": 0.7}, - litellm_params={}, - headers={}, - ) - - assert result["temperature"] == 0.7 - - def test_request_with_functions(self, config): - """Request with functions should include them""" - messages = [{"role": "user", "content": "Hello"}] - functions = [{"name": "test", "description": "Test", "parameters": {}}] - result = config.transform_request( - model="gigachat/GigaChat", - messages=messages, - optional_params={"functions": functions}, - litellm_params={}, - headers={}, - ) - - assert "functions" in result - assert len(result["functions"]) == 1 - - -class TestGigaChatSupportedParams: - """Tests for supported parameters""" - - @pytest.fixture - def config(self): - from litellm.llms.gigachat.chat.transformation import GigaChatConfig - return GigaChatConfig() - - def test_supported_params(self, config): - """Check supported parameters list""" - supported = config.get_supported_openai_params("GigaChat") - - assert "temperature" in supported - assert "max_tokens" in supported - assert "max_completion_tokens" in supported - assert "tools" in supported - assert "response_format" in supported - assert "stream" in supported diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index 5e92c10fbdc..2b01b4c2a12 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -286,7 +286,7 @@ def test_completion_claude_3_empty_response(): }, ] try: - response = litellm.completion(model="claude-3-7-sonnet-20250219", messages=messages) + response = litellm.completion(model="claude-3-opus-20240229", messages=messages) print(response) except litellm.InternalServerError as e: pytest.skip(f"InternalServerError - {str(e)}") @@ -313,7 +313,7 @@ def test_completion_claude_3(): try: # test without max tokens response = completion( - model="anthropic/claude-3-7-sonnet-20250219", + model="anthropic/claude-3-opus-20240229", messages=messages, ) # Add any assertions, here to check response args @@ -326,7 +326,7 @@ def test_completion_claude_3(): @pytest.mark.parametrize( "model", - ["anthropic/claude-3-7-sonnet-20250219", "anthropic.claude-3-sonnet-20240229-v1:0"], + ["anthropic/claude-3-opus-20240229", "anthropic.claude-3-sonnet-20240229-v1:0"], ) def test_completion_claude_3_function_call(model): litellm.set_verbose = True @@ -411,7 +411,7 @@ def test_completion_claude_3_function_call(model): "model, api_key, api_base", [ ("gpt-3.5-turbo", None, None), - ("claude-3-7-sonnet-20250219", None, None), + ("claude-3-opus-20240229", None, None), ("anthropic.claude-3-sonnet-20240229-v1:0", None, None), # ( # "azure_ai/command-r-plus", @@ -512,7 +512,7 @@ async def test_anthropic_no_content_error(): try: litellm.drop_params = True response = await litellm.acompletion( - model="anthropic/claude-3-7-sonnet-20250219", + model="anthropic/claude-3-opus-20240229", api_key=os.getenv("ANTHROPIC_API_KEY"), messages=[ { @@ -630,7 +630,7 @@ def test_completion_claude_3_multi_turn_conversations(): ] try: response = completion( - model="anthropic/claude-3-7-sonnet-20250219", + model="anthropic/claude-3-opus-20240229", messages=messages, ) print(response) @@ -644,7 +644,7 @@ def test_completion_claude_3_stream(): try: # test without max tokens response = completion( - model="anthropic/claude-3-7-sonnet-20250219", + model="anthropic/claude-3-opus-20240229", messages=messages, max_tokens=10, stream=True, @@ -669,7 +669,7 @@ def encode_image(image_path): [ "gpt-4o", "azure/gpt-4.1-mini", - "anthropic/claude-3-7-sonnet-20250219", + "anthropic/claude-3-opus-20240229", ], ) # def test_completion_base64(model): diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py index b9b5d0fdb07..0d9f84a301c 100644 --- a/tests/local_testing/test_streaming.py +++ b/tests/local_testing/test_streaming.py @@ -1418,7 +1418,7 @@ def test_bedrock_claude_3_streaming(): @pytest.mark.parametrize( "model", [ - "claude-3-7-sonnet-20250219", + "claude-3-opus-20240229", "cohere.command-r-plus-v1:0", # bedrock "gpt-3.5-turbo", ], @@ -2914,7 +2914,7 @@ def test_completion_claude_3_function_call_with_streaming(): try: # test without max tokens response = completion( - model="claude-3-7-sonnet-20250219", + model="claude-3-opus-20240229", messages=messages, tools=tools, tool_choice="required", @@ -2946,7 +2946,7 @@ def test_completion_claude_3_function_call_with_streaming(): "model", [ "gemini/gemini-2.5-flash-lite", - ], + ], # "claude-3-opus-20240229" ) # @pytest.mark.asyncio async def test_acompletion_function_call_with_streaming(model): diff --git a/tests/logging_callback_tests/test_langsmith_unit_test.py b/tests/logging_callback_tests/test_langsmith_unit_test.py index bde2b944579..e63ce9f8b38 100644 --- a/tests/logging_callback_tests/test_langsmith_unit_test.py +++ b/tests/logging_callback_tests/test_langsmith_unit_test.py @@ -47,19 +47,6 @@ async def test_get_credentials_from_env(): credentials = logger.get_credentials_from_env() assert credentials["LANGSMITH_BASE_URL"] == "https://api.smith.langchain.com" - # Test with tenant_id - credentials = logger.get_credentials_from_env( - langsmith_tenant_id="test-tenant-id" - ) - assert credentials["LANGSMITH_TENANT_ID"] == "test-tenant-id" - - # Test tenant_id from environment variable - import os - os.environ["LANGSMITH_TENANT_ID"] = "env-tenant-id" - credentials = logger.get_credentials_from_env() - assert credentials["LANGSMITH_TENANT_ID"] == "env-tenant-id" - del os.environ["LANGSMITH_TENANT_ID"] - @pytest.mark.asyncio async def test_group_batches_by_credentials(): @@ -73,7 +60,6 @@ async def test_group_batches_by_credentials(): "LANGSMITH_API_KEY": "key1", "LANGSMITH_PROJECT": "proj1", "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": None, }, ) @@ -83,7 +69,6 @@ async def test_group_batches_by_credentials(): "LANGSMITH_API_KEY": "key1", "LANGSMITH_PROJECT": "proj1", "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": None, }, ) @@ -110,7 +95,6 @@ async def test_group_batches_by_credentials_multiple_credentials(): "LANGSMITH_API_KEY": "key1", "LANGSMITH_PROJECT": "proj1", "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": None, }, ) @@ -120,7 +104,6 @@ async def test_group_batches_by_credentials_multiple_credentials(): "LANGSMITH_API_KEY": "key2", # Different API key "LANGSMITH_PROJECT": "proj1", "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": None, }, ) @@ -130,7 +113,6 @@ async def test_group_batches_by_credentials_multiple_credentials(): "LANGSMITH_API_KEY": "key1", "LANGSMITH_PROJECT": "proj2", # Different project "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": None, }, ) @@ -145,57 +127,6 @@ async def test_group_batches_by_credentials_multiple_credentials(): assert len(batch_group.queue_objects) == 1 # Each group should have one object -@pytest.mark.asyncio -async def test_group_batches_by_credentials_with_tenant_id(): - - # Test that different tenant_ids create separate groups - logger = LangsmithLogger(langsmith_api_key="test-key") - - queue_obj1 = LangsmithQueueObject( - data={"test": "data1"}, - credentials={ - "LANGSMITH_API_KEY": "key1", - "LANGSMITH_PROJECT": "proj1", - "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": "tenant1", - }, - ) - - queue_obj2 = LangsmithQueueObject( - data={"test": "data2"}, - credentials={ - "LANGSMITH_API_KEY": "key1", - "LANGSMITH_PROJECT": "proj1", - "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": "tenant2", # Different tenant_id - }, - ) - - queue_obj3 = LangsmithQueueObject( - data={"test": "data3"}, - credentials={ - "LANGSMITH_API_KEY": "key1", - "LANGSMITH_PROJECT": "proj1", - "LANGSMITH_BASE_URL": "url1", - "LANGSMITH_TENANT_ID": "tenant1", # Same as queue_obj1 - }, - ) - - logger.log_queue = [queue_obj1, queue_obj2, queue_obj3] - - grouped = logger._group_batches_by_credentials() - - # Should have two groups: one for tenant1 (queue_obj1 and queue_obj3), one for tenant2 (queue_obj2) - assert len(grouped) == 2 - for key, batch_group in grouped.items(): - assert isinstance(key, CredentialsKey) - assert key.tenant_id in ["tenant1", "tenant2"] - if key.tenant_id == "tenant1": - assert len(batch_group.queue_objects) == 2 - else: - assert len(batch_group.queue_objects) == 1 - - # Test make_dot_order @pytest.mark.asyncio async def test_make_dot_order(): @@ -270,43 +201,10 @@ async def test_async_send_batch(): call_args = logger.async_httpx_client.post.call_args assert "runs/batch" in call_args[1]["url"] assert "x-api-key" in call_args[1]["headers"] - # tenant_id should not be in headers if not provided - assert "x-tenant-id" not in call_args[1]["headers"] @pytest.mark.asyncio -async def test_async_send_batch_with_tenant_id(): - logger = LangsmithLogger( - langsmith_api_key="test-key", - langsmith_tenant_id="test-tenant-id" - ) - - # Mock the httpx client - mock_response = AsyncMock() - mock_response.status_code = 200 - logger.async_httpx_client = AsyncMock() - logger.async_httpx_client.post.return_value = mock_response - - # Add test data to queue - logger.log_queue = [ - LangsmithQueueObject( - data={"test": "data"}, credentials=logger.default_credentials - ) - ] - - await logger.async_send_batch() - - # Verify the API call includes tenant_id header - logger.async_httpx_client.post.assert_called_once() - call_args = logger.async_httpx_client.post.call_args - assert "runs/batch" in call_args[1]["url"] - assert "x-api-key" in call_args[1]["headers"] - assert "x-tenant-id" in call_args[1]["headers"] - assert call_args[1]["headers"]["x-tenant-id"] == "test-tenant-id" - - -@pytest.mark.asyncio -async def test_langsmith_key_based_logging(): +async def test_langsmith_key_based_logging(mocker): """ In key based logging langsmith_api_key and langsmith_project are passed directly to litellm.acompletion """ @@ -321,11 +219,10 @@ async def test_langsmith_key_based_logging(): mock_response.text = "" mock_async_httpx_handler.post = AsyncMock(return_value=mock_response) - mock_get_client = patch( + mock_get_client = mocker.patch( "litellm.integrations.langsmith.get_async_httpx_client", return_value=mock_async_httpx_handler ) - mock_get_client.start() litellm.set_verbose = True litellm.DEFAULT_FLUSH_INTERVAL_SECONDS = 1 @@ -356,8 +253,6 @@ async def test_langsmith_key_based_logging(): # Check headers contain the correct API key assert call_args[1]["headers"]["x-api-key"] == "fake_key_project2" - # tenant_id should not be in headers if not provided - assert "x-tenant-id" not in call_args[1]["headers"] # Verify the request body contains the expected data request_body = call_args[1]["json"] @@ -449,8 +344,6 @@ async def test_langsmith_key_based_logging(): actual_body["post"][0]["session_name"] == expected_body["post"][0]["session_name"] ) - - mock_get_client.stop() except Exception as e: pytest.fail(f"Error occurred: {e}") diff --git a/tests/logging_callback_tests/test_opentelemetry_unit_tests.py b/tests/logging_callback_tests/test_opentelemetry_unit_tests.py index 04f8abe64de..3d0682d9033 100644 --- a/tests/logging_callback_tests/test_opentelemetry_unit_tests.py +++ b/tests/logging_callback_tests/test_opentelemetry_unit_tests.py @@ -65,6 +65,31 @@ class TestOpentelemetryUnitTests(BaseLoggingCallbackTest): # External spans should only be closed by their creators parent_otel_span.end.assert_not_called() + def test_init_tracing_respects_existing_tracer_provider(self): + """ + Unit test: _init_tracing() should respect existing TracerProvider. + + When a TracerProvider already exists (e.g., set by Langfuse SDK), + LiteLLM should use it instead of creating a new one. + """ + from opentelemetry import trace + from opentelemetry.sdk.trace import TracerProvider + from litellm.integrations.opentelemetry import OpenTelemetry + + # Setup: Create and set an existing TracerProvider + tracer_provider = TracerProvider() + trace.set_tracer_provider(tracer_provider) + existing_provider = trace.get_tracer_provider() + + # Act: Initialize OpenTelemetry integration (should detect existing provider) + otel_integration = OpenTelemetry() + + # Assert: The existing provider should still be active + current_provider = trace.get_tracer_provider() + assert current_provider is existing_provider, ( + "Existing TracerProvider should be respected and not overridden" + ) + def test_get_span_context_detects_active_span(self): """ Unit test: _get_span_context() should auto-detect active spans from global context. diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index 45aae3b9aee..40e223ffe07 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -79,73 +79,6 @@ def test_routing_strategy_init(model_list): ) -def test_routing_strategy_init_invalid_strategy(model_list): - """Test that invalid routing_strategy raises ValueError with helpful message. - - See: https://github.com/BerriAI/litellm/issues/11330 - Invalid strategies like 'simple' (without '-shuffle') should fail fast - with a clear error, not silently cause 'No deployments available' errors. - """ - router = Router(model_list=model_list) - - # Test common mistake: "simple" instead of "simple-shuffle" - with pytest.raises(ValueError) as exc_info: - router.routing_strategy_init( - routing_strategy="simple", - routing_strategy_args={} - ) - - # Verify error message is helpful - error_msg = str(exc_info.value) - assert "Invalid routing_strategy" in error_msg - assert "simple" in error_msg - assert "simple-shuffle" in error_msg # Suggests the correct option - # Verify error message tells user WHERE to fix it - assert "config.yaml" in error_msg - assert "router_settings.routing_strategy" in error_msg - assert "Router SDK" in error_msg - - # Test completely invalid strategy - with pytest.raises(ValueError) as exc_info: - router.routing_strategy_init( - routing_strategy="not-a-real-strategy", - routing_strategy_args={} - ) - assert "Invalid routing_strategy" in str(exc_info.value) - - -def test_routing_strategy_init_valid_string_strategies(model_list): - """Test that all valid string routing strategies work without error. - - Valid strategies are derived from RoutingStrategy enum values plus 'simple-shuffle'. - """ - from litellm.types.router import RoutingStrategy - - router = Router(model_list=model_list) - - # All strategies from enum + simple-shuffle (default, not in enum) - valid_strategies = ["simple-shuffle"] + [s.value for s in RoutingStrategy] - - for strategy in valid_strategies: - # Should not raise - router.routing_strategy_init( - routing_strategy=strategy, routing_strategy_args={} - ) - - -def test_routing_strategy_init_valid_enum_strategies(model_list): - """Test that RoutingStrategy enum values work without error.""" - from litellm.types.router import RoutingStrategy - - router = Router(model_list=model_list) - - for strategy in RoutingStrategy: - # Should not raise when passing enum directly - router.routing_strategy_init( - routing_strategy=strategy, routing_strategy_args={} - ) - - def test_print_deployment(model_list): """Test if the api key is masked correctly""" diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 596398e639f..6490352c39b 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -1011,87 +1011,39 @@ def test_multiple_tool_calls_in_single_choice(): def test_map_reasoning_effort_adds_summary_detailed(): """ - Test that _map_reasoning_effort behavior with reasoning_auto_summary flag. + Test that _map_reasoning_effort adds summary="detailed" when user provides reasoning_effort as a string. - By default (flag=False), summary should NOT be added to avoid: - 1. Breaking for users without verified OpenAI orgs (400 errors) - 2. Making requests more expensive by including summary reasoning tokens - - When flag is enabled (flag=True or env var), summary="detailed" is added. + This ensures that when users pass reasoning_effort in the completions API for OpenAI responses/models, + the transformation automatically includes summary="detailed" in the reasoning parameter. """ - import os - - import litellm from litellm.completion_extras.litellm_responses_transformation.transformation import ( LiteLLMResponsesTransformationHandler, ) handler = LiteLLMResponsesTransformationHandler() - # Test all string effort levels - DEFAULT BEHAVIOR (no summary) + # Test all string effort levels effort_levels = ["none", "low", "medium", "high", "xhigh", "minimal"] - # Save original flag value - original_flag = litellm.reasoning_auto_summary - original_env = os.environ.get("LITELLM_REASONING_AUTO_SUMMARY") + for effort in effort_levels: + result = handler._map_reasoning_effort(effort) + + assert result is not None, f"Result should not be None for effort={effort}" + assert result["effort"] == effort, f"Effort should be {effort}" + assert result["summary"] == "detailed", f"Summary should be 'detailed' for effort={effort}" + + print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed'") - try: - # Test 1: Default behavior (flag=False, no env var) - NO summary - litellm.reasoning_auto_summary = False - if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ: - del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] - - for effort in effort_levels: - result = handler._map_reasoning_effort(effort) - - assert result is not None, f"Result should not be None for effort={effort}" - assert result["effort"] == effort, f"Effort should be {effort}" - assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}" - - print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)") - - # Test 2: With flag enabled - summary IS added - litellm.reasoning_auto_summary = True - - for effort in effort_levels: - result = handler._map_reasoning_effort(effort) - - assert result is not None, f"Result should not be None for effort={effort}" - assert result["effort"] == effort, f"Effort should be {effort}" - assert result["summary"] == "detailed", f"Summary should be 'detailed' when flag is enabled for effort={effort}" - - print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)") - - # Test 3: With env var enabled (flag disabled) - summary IS added - litellm.reasoning_auto_summary = False - os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true" - - result = handler._map_reasoning_effort("high") - assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled" - print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly") - - # Test 4: Dict input is passed through as-is (no modification) - litellm.reasoning_auto_summary = False - if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ: - del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] - - dict_input = {"effort": "high", "summary": "custom_summary"} - result_dict = handler._map_reasoning_effort(dict_input) - assert result_dict["effort"] == "high" - assert result_dict["summary"] == "custom_summary" - print("✓ Dict input is passed through without modification") - - # Test 5: None/unknown values return None - result_unknown = handler._map_reasoning_effort("unknown_value") - assert result_unknown is None - print("✓ Unknown reasoning_effort values return None") - - print("✓ All reasoning_effort behaviors work correctly with flag/env var control") + # Test that dict input is passed through as-is (no modification) + dict_input = {"effort": "high", "summary": "custom_summary"} + result_dict = handler._map_reasoning_effort(dict_input) + assert result_dict["effort"] == "high" + assert result_dict["summary"] == "custom_summary" + print("✓ Dict input is passed through without modification") - finally: - # Restore original values - litellm.reasoning_auto_summary = original_flag - if original_env is not None: - os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = original_env - elif "LITELLM_REASONING_AUTO_SUMMARY" in os.environ: - del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] + # Test that None/unknown values return None + result_unknown = handler._map_reasoning_effort("unknown_value") + assert result_unknown is None + print("✓ Unknown reasoning_effort values return None") + + print("✓ All reasoning_effort string values correctly map to summary='detailed'") diff --git a/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py b/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py deleted file mode 100644 index e717840ec95..00000000000 --- a/tests/test_litellm/integrations/langfuse/test_gemini_cached_tokens.py +++ /dev/null @@ -1,90 +0,0 @@ -""" -Test for Langfuse integration with Gemini cached_tokens bug -https://github.com/BerriAI/litellm/issues/18520 -""" -import pytest -from litellm.types.utils import PromptTokensDetailsWrapper, Usage - - -def test_cached_tokens_extraction(): - """ - Test that we can extract cached_tokens from prompt_tokens_details. - This is the core logic fix for https://github.com/BerriAI/litellm/issues/18520 - """ - # Create usage object like Gemini returns - usage = Usage( - prompt_tokens=20209, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=20203, - text_tokens=6, - ), - completion_tokens=541, - ) - - # Simulate the logic from langfuse.py lines 745-757 (after the fix) - cache_read_input_tokens = 0 # Default value - - # Check prompt_tokens_details.cached_tokens (the fix) - if hasattr(usage, "prompt_tokens_details"): - prompt_tokens_details = getattr(usage, "prompt_tokens_details", None) - if ( - prompt_tokens_details is not None - and hasattr(prompt_tokens_details, "cached_tokens") - ): - cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None) - if cached_tokens is not None and cached_tokens > 0: - cache_read_input_tokens = cached_tokens - - # Verify the fix works - assert cache_read_input_tokens == 20203, f"Expected 20203, got {cache_read_input_tokens}" - - -def test_cached_tokens_not_present(): - """Test backward compatibility when cached_tokens is not present""" - # Usage without prompt_tokens_details - usage = Usage( - prompt_tokens=100, - completion_tokens=50, - ) - - cache_read_input_tokens = 0 - - if hasattr(usage, "prompt_tokens_details"): - prompt_tokens_details = getattr(usage, "prompt_tokens_details", None) - if ( - prompt_tokens_details is not None - and hasattr(prompt_tokens_details, "cached_tokens") - ): - cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None) - if cached_tokens is not None and cached_tokens > 0: - cache_read_input_tokens = cached_tokens - - # Should remain 0 - assert cache_read_input_tokens == 0 - - -def test_cached_tokens_is_zero(): - """Test when cached_tokens is explicitly 0""" - usage = Usage( - prompt_tokens=100, - prompt_tokens_details=PromptTokensDetailsWrapper( - cached_tokens=0, - text_tokens=100, - ), - completion_tokens=50, - ) - - cache_read_input_tokens = 0 - - if hasattr(usage, "prompt_tokens_details"): - prompt_tokens_details = getattr(usage, "prompt_tokens_details", None) - if ( - prompt_tokens_details is not None - and hasattr(prompt_tokens_details, "cached_tokens") - ): - cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None) - if cached_tokens is not None and cached_tokens > 0: - cache_read_input_tokens = cached_tokens - - # Should remain 0 when cached_tokens is 0 - assert cache_read_input_tokens == 0 diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 6c17570e135..5d648f601f6 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -172,86 +172,6 @@ class TestOpenTelemetryCostBreakdown(unittest.TestCase): assert ("gen_ai.cost.original_cost", 0.004) not in call_args_list -class TestOpenTelemetryProviderInitialization(unittest.TestCase): - """Test suite for verifying provider initialization respects existing providers""" - - def test_init_tracing_respects_existing_tracer_provider(self): - """ - Unit test: _init_tracing() should respect existing TracerProvider. - - When a TracerProvider already exists (e.g., set by Langfuse SDK), - LiteLLM should use it instead of creating a new one. - """ - from opentelemetry import trace - from opentelemetry.sdk.trace import TracerProvider - - # Setup: Create and set an existing TracerProvider - tracer_provider = TracerProvider() - trace.set_tracer_provider(tracer_provider) - existing_provider = trace.get_tracer_provider() - - # Act: Initialize OpenTelemetry integration (should detect existing provider) - otel_integration = OpenTelemetry() - - # Assert: The existing provider should still be active - current_provider = trace.get_tracer_provider() - assert current_provider is existing_provider, ( - "Existing TracerProvider should be respected and not overridden" - ) - - @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True) - def test_init_metrics_respects_existing_meter_provider(self): - """ - Unit test: _init_metrics() should respect existing MeterProvider. - - When a MeterProvider already exists (e.g., set by Langfuse SDK), - LiteLLM should use it instead of creating a new one. - """ - from opentelemetry import metrics - from opentelemetry.sdk.metrics import MeterProvider - - # Create and set an existing MeterProvider - meter_provider = MeterProvider() - metrics.set_meter_provider(meter_provider) - existing_provider = metrics.get_meter_provider() - - # Act: Initialize OpenTelemetry integration (should detect existing provider) - config = OpenTelemetryConfig.from_env() - otel_integration = OpenTelemetry(config=config) - - # Assert: The existing provider should still be active - current_provider = metrics.get_meter_provider() - assert current_provider is existing_provider, ( - "Existing MeterProvider should be respected and not overridden" - ) - - @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS": "true"}, clear=True) - def test_init_logs_respects_existing_logger_provider(self): - """ - Unit test: _init_logs() should respect existing LoggerProvider. - - When a LoggerProvider already exists (e.g., set by Langfuse SDK), - LiteLLM should use it instead of creating a new one. - """ - from opentelemetry._logs import get_logger_provider, set_logger_provider - from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider - - # Create and set an existing LoggerProvider - logger_provider = OTLoggerProvider() - set_logger_provider(logger_provider) - existing_provider = get_logger_provider() - - # Act: Initialize OpenTelemetry integration (should detect existing provider) - config = OpenTelemetryConfig.from_env() - otel_integration = OpenTelemetry(config=config) - - # Assert: The existing provider should still be active - current_provider = get_logger_provider() - assert current_provider is existing_provider, ( - "Existing LoggerProvider should be respected and not overridden" - ) - - class TestOpenTelemetry(unittest.TestCase): POLL_INTERVAL = 0.05 POLL_TIMEOUT = 2.0 @@ -700,6 +620,7 @@ class TestOpenTelemetry(unittest.TestCase): self.assertEqual(attributes.get("extra.attr"), "extra-value") + def test_handle_success_spans_only(self): # make sure neither events nor metrics is on os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None) @@ -766,8 +687,11 @@ class TestOpenTelemetry(unittest.TestCase): logs = log_exporter.get_finished_logs() self.assertFalse(logs, "Did not expect any logs") - @patch.dict(os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True) def test_handle_success_spans_and_metrics(self): + # only metrics on + os.environ.pop("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", None) + os.environ["LITELLM_OTEL_INTEGRATION_ENABLE_METRICS"] = "true" + # ─── build in‐memory OTEL providers/exporters ───────────────────────────── span_exporter = InMemorySpanExporter() tracer_provider = TracerProvider() @@ -1396,23 +1320,6 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): ) self.assertEqual(normalized, "http://collector:4317/v1/logs") - def test_get_metric_reader_uses_http_exporter_for_http_protobuf(self): - """Test that http/protobuf protocol uses OTLPMetricExporterHTTP""" - from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( - OTLPMetricExporter, - ) - from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader - - config = OpenTelemetryConfig( - exporter="http/protobuf", endpoint="http://collector:4318" - ) - otel = OpenTelemetry(config=config) - - reader = otel._get_metric_reader() - - self.assertIsInstance(reader, PeriodicExportingMetricReader) - self.assertIsInstance(reader._exporter, OTLPMetricExporter) - class TestOpenTelemetryExternalSpan(unittest.TestCase): """ diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index ec2f528a35d..6a528fef8f0 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -691,29 +691,6 @@ async def test_streaming_completion_start_time(logging_obj: Logging): ) -@pytest.mark.asyncio -async def test_vertex_streaming_bad_request_not_midstream(logging_obj: Logging): - """Ensure Vertex bad request errors surface as 400, not mid-stream fallbacks.""" - from litellm.llms.vertex_ai.common_utils import VertexAIError - - async def _raise_bad_request(**kwargs): - raise VertexAIError(status_code=400, message="invalid maxOutputTokens", headers=None) - - response = CustomStreamWrapper( - completion_stream=None, - model="gemini-3-pro-preview", - logging_obj=logging_obj, - custom_llm_provider="vertex_ai_beta", - make_call=_raise_bad_request, - ) - - with pytest.raises(litellm.BadRequestError) as excinfo: - await response.__anext__() - - assert getattr(excinfo.value, "status_code", None) == 400 - assert "invalid maxOutputTokens" in str(excinfo.value) - - def test_streaming_handler_with_created_time_propagation( initialized_custom_stream_wrapper: CustomStreamWrapper, logging_obj: Logging ): diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py b/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py deleted file mode 100644 index 0f369fbb8b9..00000000000 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py +++ /dev/null @@ -1,355 +0,0 @@ -""" -Test cases for functionCall args serialization in Vertex AI Gemini. - -This test file specifically tests the edge cases where Vertex AI might return -functionCall args in unexpected formats that could lead to invalid JSON strings -like: {"x":"x"}{"a":"a"} -""" -import json -from typing import List, Optional - -import pytest - -from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, -) -from litellm.types.llms.vertex_ai import HttpxPartType - - -class TestFunctionCallArgsSerialization: - """Test cases for functionCall args serialization edge cases.""" - - def test_normal_dict_args(self): - """Test normal case: args is a dict.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": {"location": "Boston", "unit": "celsius"}, - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - assert tools[0]["function"]["name"] == "get_weather" - - # Verify arguments is a valid JSON string - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - # Should be valid JSON - parsed = json.loads(arguments) - assert parsed == {"location": "Boston", "unit": "celsius"} - - def test_none_args(self): - """Test case: args is None.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": None, - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - # Should serialize None to "null" or empty dict - assert isinstance(arguments, str) - parsed = json.loads(arguments) - # json.dumps(None) returns "null" - assert parsed is None or parsed == {} - - def test_args_as_string_valid_json(self): - """Test case: args is already a valid JSON string.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": '{"location": "Boston"}', # String, not dict - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - # If args is a string, json.dumps will double-encode it - # This would result in: "{\"location\": \"Boston\"}" - assert isinstance(arguments, str) - # This is the problematic case - string gets double-encoded - # The result would be a JSON string containing a JSON string - parsed = json.loads(arguments) - # If it's double-encoded, parsed would be a string, not a dict - if isinstance(parsed, str): - # Double-encoded case - inner_parsed = json.loads(parsed) - assert inner_parsed == {"location": "Boston"} - else: - # Normal case (shouldn't happen if args is string) - assert parsed == {"location": "Boston"} - - def test_args_as_string_invalid_json_concatenated(self): - """Test case: args is a string with concatenated JSON objects (the bug case). - - When args is a string like '{"x":"x"}{"a":"a"}', json.dumps() will serialize it - as a JSON string, resulting in: "{\"x\":\"x\"}{\"a\":\"a\"}" - This is a valid JSON string (the outer quotes), but the content inside is invalid JSON. - When you try to parse the inner content, it fails. - """ - # This simulates the case where Vertex might return something like: - # args = '{"x":"x"}{"a":"a"}' # Two JSON objects concatenated - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": '{"x":"x"}{"a":"a"}', # Invalid concatenated JSON - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - - # json.dumps() on a string will escape it, so we get: - # arguments = '"{\\"x\\":\\"x\\"}{\\"a\\":\\"a\\"}"' - # This is a valid JSON string (the outer quotes), but the inner content is invalid - parsed_outer = json.loads(arguments) - assert isinstance(parsed_outer, str) - - # The inner string is invalid JSON (two objects concatenated) - # This is the bug: the inner content cannot be parsed as valid JSON - with pytest.raises(json.JSONDecodeError): - json.loads(parsed_outer) - - # The arguments string would be: "{\"x\":\"x\"}{\"a\":\"a\"}" - # Which when parsed gives: '{"x":"x"}{"a":"a"}' (invalid JSON) - - def test_args_as_array(self): - """Test case: args is an array (unexpected but possible).""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": [{"x": "x"}, {"a": "a"}], # Array of objects - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - # Should serialize array correctly - parsed = json.loads(arguments) - assert parsed == [{"x": "x"}, {"a": "a"}] - - def test_args_missing_key(self): - """Test case: args key is missing from functionCall. - - This will raise a KeyError because the code directly accesses part["functionCall"]["args"] - without checking if the key exists. This is a bug that should be fixed. - """ - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - # args key missing - } - } - ] - - # This should raise KeyError because args key is missing - with pytest.raises(KeyError): - VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - def test_multiple_function_calls(self): - """Test case: multiple function calls in parts.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": {"location": "Boston"}, - } - }, - { - "functionCall": { - "name": "get_time", - "args": {"timezone": "EST"}, - } - }, - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 2 - assert tools[0]["function"]["name"] == "get_weather" - assert tools[1]["function"]["name"] == "get_time" - - # Both should have valid JSON arguments - args1 = json.loads(tools[0]["function"]["arguments"]) - args2 = json.loads(tools[1]["function"]["arguments"]) - assert args1 == {"location": "Boston"} - assert args2 == {"timezone": "EST"} - - def test_args_with_vertex_protobuf_format(self): - """Test case: args in Vertex protobuf format with string_value, etc.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": { - "location": {"string_value": "Boston, MA"}, - "unit": {"string_value": "celsius"}, - }, - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - # Should serialize the nested structure correctly - parsed = json.loads(arguments) - assert "location" in parsed - assert "unit" in parsed - - def test_args_as_empty_dict(self): - """Test case: args is an empty dict.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": {}, - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - parsed = json.loads(arguments) - assert parsed == {} - - def test_args_with_special_characters(self): - """Test case: args contains special characters that need escaping.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": { - "location": 'Boston, MA "downtown"', - "note": "Line 1\nLine 2", - }, - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - # Should handle special characters correctly - parsed = json.loads(arguments) - assert parsed["location"] == 'Boston, MA "downtown"' - assert parsed["note"] == "Line 1\nLine 2" - - def test_args_as_list_of_strings_that_look_like_json(self): - """Test case: args is a list containing strings that look like JSON objects.""" - # This could potentially cause issues if not handled correctly - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "get_weather", - "args": ['{"x":"x"}', '{"a":"a"}'], # List of JSON strings - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - # Should serialize list correctly - parsed = json.loads(arguments) - assert isinstance(parsed, list) - assert parsed == ['{"x":"x"}', '{"a":"a"}'] - - def test_args_as_dict_with_nested_structures(self): - """Test case: args contains nested dicts and lists.""" - parts: List[HttpxPartType] = [ - { - "functionCall": { - "name": "complex_function", - "args": { - "nested": {"key": "value"}, - "list": [1, 2, 3], - "mixed": [{"a": 1}, {"b": 2}], - }, - } - } - ] - - function, tools, idx = VertexGeminiConfig._transform_parts( - parts=parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - assert tools is not None - assert len(tools) == 1 - arguments = tools[0]["function"]["arguments"] - assert isinstance(arguments, str) - parsed = json.loads(arguments) - assert parsed["nested"] == {"key": "value"} - assert parsed["list"] == [1, 2, 3] - assert parsed["mixed"] == [{"a": 1}, {"b": 2}] - - -if __name__ == "__main__": - pytest.main([__file__, "-v"]) - diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index d09de3a0f26..91a28ee6ec9 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -10,13 +10,11 @@ from pydantic import BaseModel import litellm from litellm import ModelResponse, completion -from litellm.llms.vertex_ai.common_utils import VertexAIError from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) from litellm.types.llms.vertex_ai import UsageMetadata from litellm.types.utils import ChoiceLogprobs, Usage -from litellm.utils import CustomStreamWrapper def test_top_logprobs(): @@ -1607,39 +1605,6 @@ def test_vertex_ai_annotation_streaming_events(): assert "Weather information" in annotation["url_citation"]["title"] -@pytest.mark.asyncio -async def test_vertex_ai_streaming_bad_request_is_not_wrapped(): - class DummyLogging: - def __init__(self): - self.model_call_details = {"litellm_params": {}} - self.optional_params = {} - self.messages = [] - self.completion_start_time = None - self.stream_options = None - - def failure_handler(self, *args, **kwargs): - return None - - async def async_failure_handler(self, *args, **kwargs): - return None - - async def failing_make_call(client=None, **kwargs): - raise VertexAIError(status_code=400, message="bad input", headers={}) - - stream = CustomStreamWrapper( - completion_stream=None, - make_call=failing_make_call, - model="gemini-3-pro-preview", - logging_obj=DummyLogging(), - custom_llm_provider="vertex_ai_beta", - ) - - with pytest.raises(litellm.BadRequestError) as exc_info: - await stream.__anext__() - - assert getattr(exc_info.value, "status_code", None) == 400 - - def test_vertex_ai_annotation_conversion(): """ Test the conversion of Vertex AI grounding metadata to OpenAI annotations. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4c5723b8284..6df9abd3fee 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -354,7 +354,7 @@ async def test_register_client_remote_registration_success(): request_payload = { "client_name": "Litellm Proxy", - "grant_types": ["authorization_code", "refresh_token"], + "grant_types": ["authorization_code"], "response_types": ["code"], "token_endpoint_auth_method": "client_secret_post", } @@ -556,33 +556,9 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( oauth_protected_resource_mcp, ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - from litellm.types.mcp import MCPAuth - from litellm.types.mcp_server.mcp_server_manager import MCPServer - from litellm.proxy._types import MCPTransport from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") - # Clear registry - global_mcp_server_manager.registry.clear() - - # Create mock OAuth2 server - oauth2_server = MCPServer( - server_id="test_oauth_server", - name="test_oauth", - server_name="test_oauth", - alias="test_oauth", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - client_id="test_client_id", - client_secret="test_client_secret", - authorization_url="https://provider.com/oauth/authorize", - token_url="https://provider.com/oauth/token", - scopes=["read", "write"], - ) - global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server # Mock request with http base_url but X-Forwarded-Proto: https mock_request = MagicMock(spec=Request) @@ -592,14 +568,13 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): # Call the endpoint response = await oauth_protected_resource_mcp( request=mock_request, - mcp_server_name="test_oauth", + mcp_server_name="test_server", ) # Verify response uses HTTPS URLs assert response["authorization_servers"][0].startswith( "https://litellm.example.com/" ) - assert response["scopes_supported"] == oauth2_server.scopes @pytest.mark.asyncio @@ -609,33 +584,9 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( oauth_authorization_server_mcp, ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - from litellm.types.mcp import MCPAuth - from litellm.types.mcp_server.mcp_server_manager import MCPServer - from litellm.proxy._types import MCPTransport from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") - # Clear registry - global_mcp_server_manager.registry.clear() - - # Create mock OAuth2 server - oauth2_server = MCPServer( - server_id="test_oauth_server", - name="test_oauth", - server_name="test_oauth", - alias="test_oauth", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - client_id="test_client_id", - client_secret="test_client_secret", - authorization_url="https://provider.com/oauth/authorize", - token_url="https://provider.com/oauth/token", - scopes=["read", "write"], - ) - global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server # Mock request with http base_url but X-Forwarded-Proto: https mock_request = MagicMock(spec=Request) @@ -645,15 +596,13 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): # Call the endpoint response = await oauth_authorization_server_mcp( request=mock_request, - mcp_server_name="test_oauth", + mcp_server_name="test_server", ) # Verify response uses HTTPS URLs assert response["authorization_endpoint"].startswith("https://litellm.example.com/") assert response["token_endpoint"].startswith("https://litellm.example.com/") assert response["registration_endpoint"].startswith("https://litellm.example.com/") - assert response["grant_types_supported"] == ["authorization_code", "refresh_token"] - assert response["scopes_supported"] == oauth2_server.scopes @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index d59b3f04ef5..6491e11024a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -594,26 +594,7 @@ class TestMCPServerManager: assert ( server.registration_url == "https://discovered.example.com/register" ) - @pytest.mark.asyncio - async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self): - manager = MCPServerManager() - config = { - "example": { - "url": "https://example.com/mcp", - "transport": MCPTransport.http, - "auth_type": MCPAuth.oauth2, - "scopes": ["config"], - "authorization_url": "https://config.example.com/auth", - } - } - - await manager.load_servers_from_config(config) - - # Initialize the tool mapping - await manager._initialize_tool_name_to_mcp_server_name_mapping() - assert manager.tool_name_to_mcp_server_name_mapping == {} - @pytest.mark.asyncio async def test_list_tools_handles_missing_server_alias(self): """Test that list_tools handles servers without alias gracefully""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py index 35ed49a84ed..6d6129b17bb 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_qualifire.py @@ -2,7 +2,6 @@ Unit tests for Qualifire guardrail integration. """ -import sys from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -140,98 +139,76 @@ class TestQualifireGuardrailEvaluateKwargs: @pytest.mark.asyncio async def test_evaluate_called_with_prompt_injections(self): """Test that evaluate is called with prompt_injections enabled.""" - # Mock the qualifire module and its types - mock_qualifire_types = MagicMock() - mock_llm_message = MagicMock() - mock_llm_tool_call = MagicMock() - mock_message_instance = MagicMock() - mock_llm_message.return_value = mock_message_instance - - mock_qualifire_types.LLMMessage = mock_llm_message - mock_qualifire_types.LLMToolCall = mock_llm_tool_call - - with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}): - from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( - QualifireGuardrail, - ) + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) - guardrail = QualifireGuardrail( - api_key="test_key", - prompt_injections=True, - guardrail_name="test_guardrail", - ) + guardrail = QualifireGuardrail( + api_key="test_key", + prompt_injections=True, + guardrail_name="test_guardrail", + ) - # Mock the client - mock_client = MagicMock() - mock_result = MagicMock() - mock_result.score = 100 - mock_result.status = "completed" - mock_result.evaluationResults = [] - mock_client.evaluate.return_value = mock_result - guardrail._client = mock_client + # Mock the client + mock_client = MagicMock() + mock_result = MagicMock() + mock_result.score = 100 + mock_result.status = "completed" + mock_result.evaluationResults = [] + mock_client.evaluate.return_value = mock_result + guardrail._client = mock_client - messages = [{"role": "user", "content": "Hello, world!"}] + messages = [{"role": "user", "content": "Hello, world!"}] - await guardrail._run_qualifire_check( - messages=messages, output=None, dynamic_params={} - ) + await guardrail._run_qualifire_check( + messages=messages, output=None, dynamic_params={} + ) - # Verify evaluate was called with correct kwargs - mock_client.evaluate.assert_called_once() - call_kwargs = mock_client.evaluate.call_args[1] - assert call_kwargs["prompt_injections"] is True - assert "messages" in call_kwargs + # Verify evaluate was called with correct kwargs + mock_client.evaluate.assert_called_once() + call_kwargs = mock_client.evaluate.call_args[1] + assert call_kwargs["prompt_injections"] is True + assert "messages" in call_kwargs @pytest.mark.asyncio async def test_evaluate_called_with_multiple_checks(self): """Test that evaluate is called with multiple checks enabled.""" - # Mock the qualifire module and its types - mock_qualifire_types = MagicMock() - mock_llm_message = MagicMock() - mock_llm_tool_call = MagicMock() - mock_message_instance = MagicMock() - mock_llm_message.return_value = mock_message_instance - - mock_qualifire_types.LLMMessage = mock_llm_message - mock_qualifire_types.LLMToolCall = mock_llm_tool_call - - with patch.dict('sys.modules', {'qualifire': MagicMock(), 'qualifire.types': mock_qualifire_types}): - from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( - QualifireGuardrail, - ) + from litellm.proxy.guardrails.guardrail_hooks.qualifire.qualifire import ( + QualifireGuardrail, + ) - guardrail = QualifireGuardrail( - api_key="test_key", - prompt_injections=True, - pii_check=True, - hallucinations_check=True, - assertions=["Output must be valid JSON"], - guardrail_name="test_guardrail", - ) + guardrail = QualifireGuardrail( + api_key="test_key", + prompt_injections=True, + pii_check=True, + hallucinations_check=True, + assertions=["Output must be valid JSON"], + guardrail_name="test_guardrail", + ) - # Mock the client - mock_client = MagicMock() - mock_result = MagicMock() - mock_result.score = 100 - mock_result.status = "completed" - mock_result.evaluationResults = [] - mock_client.evaluate.return_value = mock_result - guardrail._client = mock_client + # Mock the client + mock_client = MagicMock() + mock_result = MagicMock() + mock_result.score = 100 + mock_result.status = "completed" + mock_result.evaluationResults = [] + mock_client.evaluate.return_value = mock_result + guardrail._client = mock_client - messages = [{"role": "user", "content": "Hello, world!"}] + messages = [{"role": "user", "content": "Hello, world!"}] - await guardrail._run_qualifire_check( - messages=messages, output="Test output", dynamic_params={} - ) + await guardrail._run_qualifire_check( + messages=messages, output="Test output", dynamic_params={} + ) - # Verify evaluate was called with correct kwargs - mock_client.evaluate.assert_called_once() - call_kwargs = mock_client.evaluate.call_args[1] - assert call_kwargs["prompt_injections"] is True - assert call_kwargs["pii_check"] is True - assert call_kwargs["hallucinations_check"] is True - assert call_kwargs["assertions"] == ["Output must be valid JSON"] - assert call_kwargs["output"] == "Test output" + # Verify evaluate was called with correct kwargs + mock_client.evaluate.assert_called_once() + call_kwargs = mock_client.evaluate.call_args[1] + assert call_kwargs["prompt_injections"] is True + assert call_kwargs["pii_check"] is True + assert call_kwargs["hallucinations_check"] is True + assert call_kwargs["assertions"] == ["Output must be valid JSON"] + assert call_kwargs["output"] == "Test output" class TestQualifireGuardrailCheckIfFlagged: diff --git a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py index 97c1733a935..011031c1e4f 100644 --- a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py +++ b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py @@ -40,7 +40,6 @@ class TestKeyManagementEventHooksIndependentOperations: mock_data = MagicMock() mock_data.key_alias = "test-key-alias" mock_data.team_id = None - mock_data.send_invite_email = True mock_response = MagicMock() mock_response.model_dump.return_value = {"key": "sk-test", "token": "test-token"} @@ -60,10 +59,6 @@ class TestKeyManagementEventHooksIndependentOperations: KeyManagementEventHooks, "_store_virtual_key_in_secret_manager", side_effect=mock_store_secret, - ), patch.object( - KeyManagementEventHooks, - "_is_email_sending_enabled", - return_value=True, ), patch( "litellm.store_audit_logs", False ), patch( @@ -101,7 +96,6 @@ class TestKeyManagementEventHooksIndependentOperations: mock_data = MagicMock() mock_data.key_alias = "test-key-alias" mock_data.team_id = None - mock_data.send_invite_email = True mock_response = MagicMock() mock_response.model_dump.return_value = {"key": "sk-test", "token": "test-token"} @@ -121,10 +115,6 @@ class TestKeyManagementEventHooksIndependentOperations: KeyManagementEventHooks, "_store_virtual_key_in_secret_manager", side_effect=mock_store_secret_raises, - ), patch.object( - KeyManagementEventHooks, - "_is_email_sending_enabled", - return_value=True, ), patch( "litellm.store_audit_logs", False ), patch( diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index f2bae2cb14a..cc268ab9925 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -231,40 +231,6 @@ class TestListMCPServers: assert server.url == "https://mcp.deepwiki.com/mcp" assert server.transport == "http" - @pytest.mark.asyncio - async def test_list_mcp_servers_view_all_mode(self): - """Users should see all MCP servers when view_all mode is enabled.""" - - mock_user_auth = generate_mock_user_api_key_auth( - user_role=LitellmUserRoles.INTERNAL_USER - ) - - mock_servers = [ - generate_mock_mcp_server_db_record(server_id="server-1", alias="One"), - generate_mock_mcp_server_db_record(server_id="server-2", alias="Two"), - ] - - mock_manager = MagicMock() - mock_manager.get_all_mcp_servers_unfiltered = AsyncMock( - return_value=mock_servers - ) - - with patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode", - return_value="view_all", - ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", - mock_manager, - ): - from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - fetch_all_mcp_servers, - ) - - result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth) - - assert len(result) == 2 - assert {server.server_id for server in result} == {"server-1", "server-2"} - @pytest.mark.asyncio async def test_list_mcp_servers_combined_config_and_db(self): """ @@ -1130,51 +1096,6 @@ class TestHealthCheckServers: assert result[0]["server_id"] == "server-1" assert result[0]["status"] == "healthy" - @pytest.mark.asyncio - async def test_health_check_view_all_mode(self): - """view_all mode should return health info for all MCP servers.""" - - from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - health_check_servers, - ) - - mock_user_auth = generate_mock_user_api_key_auth( - user_role=LitellmUserRoles.INTERNAL_USER - ) - - health_result_one = generate_mock_mcp_server_db_record( - server_id="server-1", alias="One" - ) - health_result_one.status = "healthy" - - health_result_two = generate_mock_mcp_server_db_record( - server_id="server-2", alias="Two" - ) - health_result_two.status = "unhealthy" - - mock_manager = MagicMock() - mock_manager.get_all_mcp_servers_with_health_unfiltered = AsyncMock( - return_value=[health_result_one, health_result_two] - ) - - with patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode", - return_value="view_all", - ), patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", - mock_manager, - ): - result = await health_check_servers( - server_ids=None, - user_api_key_dict=mock_user_auth, - ) - - assert len(result) == 2 - assert result[0]["server_id"] == "server-1" - assert result[0]["status"] == "healthy" - assert result[1]["server_id"] == "server-2" - assert result[1]["status"] == "unhealthy" - @pytest.mark.asyncio async def test_health_check_unauthorized_servers(self): """ diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index ad4f53dac4b..8fdfd6897a8 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -742,16 +742,18 @@ class TestProxySettingEndpoints: ): """Test updating UI settings with an allowlisted field""" from unittest.mock import AsyncMock, MagicMock - from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - from litellm.proxy._types import UserAPIKeyAuth - # Override the FastAPI dependency with a proper mock - mock_user_auth = UserAPIKeyAuth( - user_id="test-user-123", - user_role=LitellmUserRoles.PROXY_ADMIN, + class MockUser: + def __init__(self, user_role): + self.user_role = user_role + + async def mock_admin_auth(): + return MockUser(LitellmUserRoles.PROXY_ADMIN) + + monkeypatch.setattr( + "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.user_api_key_auth", + mock_admin_auth, ) - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth - monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() @@ -759,11 +761,7 @@ class TestProxySettingEndpoints: payload = {"disable_model_add_for_internal_users": True} - try: - response = client.patch("/update/ui_settings", json=payload) - finally: - # Clean up the dependency override - app.dependency_overrides.clear() + response = client.patch("/update/ui_settings", json=payload) assert response.status_code == 200 data = response.json() @@ -782,16 +780,18 @@ class TestProxySettingEndpoints: ): """Test non-allowlisted UI settings are ignored on update""" from unittest.mock import AsyncMock, MagicMock - from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - from litellm.proxy._types import UserAPIKeyAuth - # Override the FastAPI dependency with a proper mock - mock_user_auth = UserAPIKeyAuth( - user_id="test-user-123", - user_role=LitellmUserRoles.PROXY_ADMIN, + class MockUser: + def __init__(self, user_role): + self.user_role = user_role + + async def mock_admin_auth(): + return MockUser(LitellmUserRoles.PROXY_ADMIN) + + monkeypatch.setattr( + "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.user_api_key_auth", + mock_admin_auth, ) - app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth - monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() @@ -802,11 +802,7 @@ class TestProxySettingEndpoints: "unsupported_flag": True, } - try: - response = client.patch("/update/ui_settings", json=payload) - finally: - # Clean up the dependency override - app.dependency_overrides.clear() + response = client.patch("/update/ui_settings", json=payload) assert response.status_code == 200 data = response.json() diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts index f03f3977115..3e09c3c2ca8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts @@ -1,7 +1,7 @@ -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { getSSOSettings } from "@/components/networking"; import { useQuery, UseQueryResult } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; +import { getSSOSettings } from "@/components/networking"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; export interface SSOFieldSchema { description: string; @@ -27,15 +27,13 @@ export interface SSOSettingsValues { proxy_base_url: string | null; user_email: string | null; ui_access_mode: string | null; - role_mappings: RoleMappings; -} - -export interface RoleMappings { - provider: string; - group_claim: string; - default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer"; - roles: { - [key: string]: string[]; + role_mappings: { + provider: string; + group_claim: string; + default_role: "internal_user" | "internal_user_viewer" | "proxy_admin" | "proxy_admin_viewer"; + roles: { + [key: string]: string[]; + }; }; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index a8b1d2cddc9..65418179595 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -18,7 +18,6 @@ import { useQueryClient } from "@tanstack/react-query"; import { Col, Grid, Icon, Tab, TabGroup, TabList, TabPanel, TabPanels, Text } from "@tremor/react"; import type { UploadProps } from "antd"; import { Form, Typography } from "antd"; -import { PlusCircleOutlined } from "@ant-design/icons"; import React, { useEffect, useMemo, useState } from "react"; import AddModelTab from "../../../components/add_model/add_model_tab"; import HealthCheckComponent from "../../../components/model_dashboard/HealthCheckComponent"; @@ -275,30 +274,6 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te )} - - {/* Missing Provider Banner */} -
-
- -
-
-

Missing a provider?

-

- The LiteLLM engineering team is constantly adding support for new LLM models, providers, endpoints. If you don't see the one you need, let us know and we'll prioritize it. -

-
- - Request Provider - - - - -
{selectedModelId && !isLoading ? ( = ({ form, onFormS - prevValues.use_role_mappings !== currentValues.use_role_mappings || - prevValues.sso_provider !== currentValues.sso_provider - } + shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings} > {({ getFieldValue }) => { const useRoleMappings = getFieldValue("use_role_mappings"); - const provider = getFieldValue("sso_provider"); - const supportsRoleMappings = provider === "okta" || provider === "generic"; - return useRoleMappings && supportsRoleMappings ? ( + return useRoleMappings ? ( = ({ form, onFormS - prevValues.use_role_mappings !== currentValues.use_role_mappings || - prevValues.sso_provider !== currentValues.sso_provider - } + shouldUpdate={(prevValues, currentValues) => prevValues.use_role_mappings !== currentValues.use_role_mappings} > {({ getFieldValue }) => { const useRoleMappings = getFieldValue("use_role_mappings"); - const provider = getFieldValue("sso_provider"); - const supportsRoleMappings = provider === "okta" || provider === "generic"; - return useRoleMappings && supportsRoleMappings ? ( + return useRoleMappings ? ( <>