diff --git a/AGENTS.md b/AGENTS.md index 5a48049ef45..bfd44304d55 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -174,6 +174,8 @@ When opening issues or pull requests, follow these templates: 3. **Rate Limits**: Respect provider rate limits in tests 4. **Memory Usage**: Be mindful of memory usage in streaming scenarios 5. **Dependencies**: Keep dependencies minimal and well-justified +6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections +7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks ## HELPFUL RESOURCES diff --git a/CLAUDE.md b/CLAUDE.md index 3cb67908076..3b597fb8a90 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -97,6 +97,10 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components: - Integration tests for each provider in `tests/llm_translation/` - Proxy tests in `tests/proxy_unit_tests/` - Load tests in `tests/load_tests/` +- **Always add tests when adding new entity types or features** — if the existing test file covers other entity types, add corresponding tests for the new one + +### UI / Backend Consistency +- When wiring a new UI entity type to an existing backend endpoint, verify the backend API contract (single value vs. array, required vs. optional params) and ensure the UI controls match — e.g., use a single-select dropdown when the backend accepts a single value, not a multi-select ### Database Migrations - Prisma handles schema migrations diff --git a/docs/my-website/docs/completion/prompt_caching.md b/docs/my-website/docs/completion/prompt_caching.md index 630c9e58d24..dca5f5c0cff 100644 --- a/docs/my-website/docs/completion/prompt_caching.md +++ b/docs/my-website/docs/completion/prompt_caching.md @@ -63,7 +63,6 @@ for _ in range(2): } ], }, - # marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache. { "role": "user", "content": [ @@ -77,7 +76,6 @@ for _ in range(2): "role": "assistant", "content": "Certainly! the key terms and conditions are the following: the contract is 1 year long for $10/mo", }, - # The final turn is marked with cache-control, for continuing in followups. { "role": "user", "content": [ @@ -112,16 +110,16 @@ model_list: api_key: os.environ/OPENAI_API_KEY ``` -2. Start proxy +2. Start proxy ```bash litellm --config /path/to/config.yaml ``` -3. Test it! +3. Test it! ```python -from openai import OpenAI +from openai import OpenAI import os client = OpenAI( @@ -144,7 +142,6 @@ for _ in range(2): } ], }, - # marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache. { "role": "user", "content": [ @@ -158,7 +155,6 @@ for _ in range(2): "role": "assistant", "content": "Certainly! the key terms and conditions are the following: the contract is 1 year long for $10/mo", }, - # The final turn is marked with cache-control, for continuing in followups. { "role": "user", "content": [ @@ -183,6 +179,78 @@ assert response.usage.prompt_tokens_details.cached_tokens > 0 +### OpenAI `prompt_cache_key` and `prompt_cache_retention` + +OpenAI prompt caching is [**automatic**](https://platform.openai.com/docs/guides/prompt-caching) — no `cache_control` message annotations are needed. Any request with 1024+ prompt tokens is eligible for caching. + +OpenAI also supports two optional parameters for more control over caching behavior: + +- **`prompt_cache_key`** (string) — A routing hint that improves cache hit rates for requests sharing long common prefixes. Requests with the same cache key are routed to the same backend, increasing the likelihood of a cache hit. +- **`prompt_cache_retention`** (`"in_memory"` or `"24h"`) — Controls cache TTL. Default is `"in_memory"` (5–10 min). Set to `"24h"` for extended caching that offloads KV tensors to GPU-local storage. + + + + +```python +from litellm import completion +import os + +os.environ["OPENAI_API_KEY"] = "" + +response = completion( + model="gpt-4o", + messages=[ + { + "role": "system", + "content": "You are an AI assistant tasked with analyzing legal documents. " + + "Here is the full text of a complex legal agreement " * 400, + }, + { + "role": "user", + "content": "What are the key terms and conditions?", + }, + ], + prompt_cache_key="legal-doc-analysis", + prompt_cache_retention="24h", +) +print(response.usage) +``` + + + + +```python +from openai import OpenAI + +client = OpenAI( + api_key="LITELLM_PROXY_KEY", + base_url="LITELLM_PROXY_BASE", +) + +response = client.chat.completions.create( + model="gpt-4o", + messages=[ + { + "role": "system", + "content": "You are an AI assistant tasked with analyzing legal documents. " + + "Here is the full text of a complex legal agreement " * 400, + }, + { + "role": "user", + "content": "What are the key terms and conditions?", + }, + ], + extra_body={ + "prompt_cache_key": "legal-doc-analysis", + "prompt_cache_retention": "24h", + }, +) +print(response.usage) +``` + + + + ### Anthropic Example Anthropic charges for cache writes. diff --git a/docs/my-website/docs/generateContent.md b/docs/my-website/docs/generateContent.md index 4453e5ce06d..bf8e1b6c03b 100644 --- a/docs/my-website/docs/generateContent.md +++ b/docs/my-website/docs/generateContent.md @@ -15,6 +15,7 @@ Use LiteLLM to call Google AI's generateContent endpoints for text generation, m | Streaming | ✅ | | | Fallbacks | ✅ | between supported models | | Loadbalancing | ✅ | between supported models | +| Metadata Tracking | ✅ | passes trace ID, metadata to observability callbacks (e.g. S3, Langfuse) | ## Usage --- diff --git a/docs/my-website/docs/proxy/credential_usage_tracking.md b/docs/my-website/docs/proxy/credential_usage_tracking.md new file mode 100644 index 00000000000..25658144c49 --- /dev/null +++ b/docs/my-website/docs/proxy/credential_usage_tracking.md @@ -0,0 +1,19 @@ +# Credential Usage Tracking + +When a model is attached to a [reusable credential](./ui_credentials.md), LiteLLM automatically injects the credential name as a tag on every request that uses that model. This means credential-level spend and usage are tracked with zero extra configuration. + +## How It Works + +When you attach a model to a reusable credential via `litellm_credential_name`, each request routed through that model is tagged `Credential: ` (for example, `Credential: xAI`). This tag flows into `DailyTagSpend` and appears in the **Tag** view on the Usage page, where you can filter spend and usage by credential. + +If a model has no credential attached, behavior is unchanged—no credential tag is added. + +## Viewing Credential Usage + +In the Admin UI, go to **Usage → Tag** and look for tags with the `Credential: ` prefix. These represent aggregated spend and token usage across all requests that used that credential. + +## Related Documentation + +- [Adding LLM Credentials](./ui_credentials.md) - How to create and attach reusable credentials to models +- [Tag Budgets](./tag_budgets.md) - Setting spend limits on tags +- [Tag Routing](./tag_routing.md) - Routing requests based on tags diff --git a/docs/my-website/docs/proxy/forward_client_headers.md b/docs/my-website/docs/proxy/forward_client_headers.md index 2155a7517be..17f813eabee 100644 --- a/docs/my-website/docs/proxy/forward_client_headers.md +++ b/docs/my-website/docs/proxy/forward_client_headers.md @@ -37,11 +37,11 @@ The following rules determine which headers are forwarded (see [`_get_forwardabl | Rule | Example | Forwarded? | |---|---|---| -| Headers starting with `x-` | `x-trace-id`, `x-custom-header`, `x-request-source` | ✅ Yes | -| `anthropic-beta` header | `anthropic-beta: prompt-caching-2024-07-31` | ✅ Yes | -| Headers starting with `x-stainless-*` | `x-stainless-lang`, `x-stainless-arch` | ❌ No (causes OpenAI SDK issues) | -| Standard HTTP headers | `Authorization`, `Content-Type`, `Host` | ❌ No | -| Other provider headers | `Accept`, `User-Agent` | ❌ No | +| Headers starting with `x-` | `x-trace-id`, `x-custom-header`, `x-request-source` | Yes | +| `anthropic-beta` header | `anthropic-beta: prompt-caching-2024-07-31` | Yes | +| Headers starting with `x-stainless-*` | `x-stainless-lang`, `x-stainless-arch` | No (causes OpenAI SDK issues) | +| Standard HTTP headers | `Authorization`, `Content-Type`, `Host` | No | +| Other provider headers | `Accept`, `User-Agent` | No | ### Additional Header Mechanisms @@ -61,6 +61,125 @@ general_settings: forward_client_headers_to_llm_api: true ``` +## Forward LLM Provider Authentication Headers + +**New in v1.82+**: By default, LiteLLM strips authentication headers like `x-api-key`, `x-goog-api-key`, and `api-key` from client requests for security (these are typically used to authenticate with the proxy itself). However, you can enable forwarding of these LLM provider authentication headers to allow **Bring Your Own Key (BYOK)** scenarios where clients send their own API keys to the LLM provider. + +### Configuration + +Add `forward_llm_provider_auth_headers: true` to your `general_settings`: + +```yaml +general_settings: + forward_client_headers_to_llm_api: true + forward_llm_provider_auth_headers: true # 👈 Enable BYOK +``` + +### Which Headers Are Forwarded + +When `forward_llm_provider_auth_headers: true`, the following LLM provider authentication headers are preserved and forwarded: + +| Header | Provider | Example | +|--------|----------|---------| +| `x-api-key` | Anthropic, Azure AI, Databricks | `x-api-key: sk-ant-api03-...` | +| `x-goog-api-key` | Google AI Studio | `x-goog-api-key: AIza...` | +| `api-key` | Azure OpenAI | `api-key: your-azure-key` | +| `ocp-apim-subscription-key` | Azure APIM | `ocp-apim-subscription-key: your-key` | + +:::warning Important Security Note +The proxy's `Authorization` header (used for proxy authentication) is **never** forwarded to LLM providers, even with this setting enabled. This ensures your proxy authentication remains secure. +::: + +### Use Case: Client-Side API Keys (BYOK) + +This feature enables scenarios where: +1. **Clients bring their own LLM provider API keys** instead of using keys configured in the proxy +2. **Multi-tenant applications** where each tenant has their own Anthropic/OpenAI account +3. **Development environments** where developers use their personal API keys through a shared proxy + +#### Example: Anthropic BYOK + +```yaml +# proxy_config.yaml +model_list: + - model_name: claude-sonnet-4 + litellm_params: + model: anthropic/claude-sonnet-4-20250514 + # No api_key configured! Will use client's key + +general_settings: + forward_client_headers_to_llm_api: true + forward_llm_provider_auth_headers: true # Enable BYOK +``` + +Client request: +```bash +curl -X POST "http://localhost:4000/v1/messages" \ + -H "Authorization: Bearer sk-proxy-auth-123" \ # Proxy authentication (stripped) + -H "x-api-key: sk-ant-api03-YOUR-KEY..." \ # Client's Anthropic key (forwarded!) + -H "Content-Type: application/json" \ + -d '{ + "model": "claude-sonnet-4", + "messages": [{"role": "user", "content": "Hello"}], + "max_tokens": 100 + }' +``` + +#### Example: Google AI Studio BYOK + +```yaml +model_list: + - model_name: gemini-pro + litellm_params: + model: gemini/gemini-1.5-pro + # No api_key configured + +general_settings: + forward_client_headers_to_llm_api: true + forward_llm_provider_auth_headers: true +``` + +Client request: +```bash +curl -X POST "http://localhost:4000/v1/chat/completions" \ + -H "Authorization: Bearer sk-proxy-auth-123" \ + -H "x-goog-api-key: AIza..." \ + -d '{ + "model": "gemini-pro", + "messages": [{"role": "user", "content": "Hello"}] + }' +``` + +### Security Considerations + +**When to Use This Feature:** +- Internal tools where you trust all clients +- Development/testing environments +- Multi-tenant apps with proper client authentication +- Scenarios where you want clients to use their own API keys + +**When NOT to Use:** +- Public APIs where you don't trust all clients +- When you want centralized billing/cost control +- When you need to enforce rate limits at the proxy level + +### Backward Compatibility + +For backward compatibility, if you have `forward_client_headers_to_llm_api: true` but don't explicitly set `forward_llm_provider_auth_headers`, the behavior is: +- **Default**: LLM provider auth headers are **NOT** forwarded (safe default) +- **Explicit `true`**: LLM provider auth headers **ARE** forwarded (BYOK enabled) + +```yaml +# Safe default - auth headers NOT forwarded +general_settings: + forward_client_headers_to_llm_api: true + +# BYOK enabled - auth headers ARE forwarded +general_settings: + forward_client_headers_to_llm_api: true + forward_llm_provider_auth_headers: true # 👈 Opt-in required +``` + ## Enable for a Model Group Add the `forward_client_headers_to_llm_api` setting under `model_group_settings` in your configuration: diff --git a/docs/my-website/docs/proxy/ui_credentials.md b/docs/my-website/docs/proxy/ui_credentials.md index 40db5368596..f10f2631f83 100644 --- a/docs/my-website/docs/proxy/ui_credentials.md +++ b/docs/my-website/docs/proxy/ui_credentials.md @@ -46,6 +46,10 @@ Go to Add Model -> Existing Credentials -> Select your credential in the dropdow +## Usage Tracking + +Models attached to a reusable credential are automatically tracked in the Usage page. Each request is tagged `Credential: ` and appears in the **Tag** view, so you can filter spend and usage by credential without any extra configuration. See [Credential Usage Tracking](./credential_usage_tracking.md) for details. + ## Frequently Asked Questions diff --git a/docs/my-website/docs/tutorials/presidio_pii_masking.md b/docs/my-website/docs/tutorials/presidio_pii_masking.md index ea3761163f2..d6fe1adbd01 100644 --- a/docs/my-website/docs/tutorials/presidio_pii_masking.md +++ b/docs/my-website/docs/tutorials/presidio_pii_masking.md @@ -592,6 +592,21 @@ def test_pii_masking_allows_normal_text(): ## Part 7: Troubleshooting +### Issue: Guardrail failure: non-JSON response from Presidio + +**Symptom:** You receive an error indicating `expected application/json Content-Type but received text/html` or similar. + +**Root cause:** Your ingress controller or reverse proxy might be routing the `/analyze` or `/anonymize` POST request to a health endpoint (like `/health` or `/presidio-analyzer/health`) which returns plain text instead of JSON. + +**Fix:** Ensure your `PRESIDIO_ANALYZER_API_BASE` and `PRESIDIO_ANONYMIZER_API_BASE` are correctly pointing directly to the Presidio API endpoints, or that your ingress routes the path correctly without stripping it and inadvertently forwarding to a plain-text health check endpoint. + +**Verification:** You can verify your endpoints using `curl`. It should return a JSON array, not `text/html`: +```bash +curl -sv -X POST http://your-analyzer-endpoint/analyze \ + -H "Content-Type: application/json" \ + -d '{"text":"test","language":"en"}' +``` + ### Issue: Presidio Not Detecting PII **Check 1: Language Configuration** diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.48-py3-none-any.whl b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.48-py3-none-any.whl new file mode 100644 index 00000000000..8dc2d8e136d Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.48-py3-none-any.whl differ diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.48.tar.gz b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.48.tar.gz new file mode 100644 index 00000000000..65bf8c3718e Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.48.tar.gz differ diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260224201417_spend_logs_request_duration/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260224201417_spend_logs_request_duration/migration.sql new file mode 100644 index 00000000000..892aa59e9f8 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260224201417_spend_logs_request_duration/migration.sql @@ -0,0 +1,3 @@ +-- AlterTable +ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN "request_duration_ms" INTEGER; + diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260224203854_add_agent_object_permissions_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260224203854_add_agent_object_permissions_table/migration.sql new file mode 100644 index 00000000000..78e364d5478 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260224203854_add_agent_object_permissions_table/migration.sql @@ -0,0 +1,40 @@ +-- AlterTable +ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN "object_permission_id" TEXT; + +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" DROP COLUMN "spec_path"; + +-- AlterTable +ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "agent_id" TEXT; + +-- CreateTable +CREATE TABLE "LiteLLM_ToolTable" ( + "tool_id" TEXT NOT NULL, + "tool_name" TEXT NOT NULL, + "origin" TEXT, + "call_policy" TEXT NOT NULL DEFAULT 'untrusted', + "call_count" INTEGER NOT NULL DEFAULT 0, + "assignments" JSONB DEFAULT '{}', + "key_hash" TEXT, + "team_id" TEXT, + "key_alias" TEXT, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT, + + CONSTRAINT "LiteLLM_ToolTable_pkey" PRIMARY KEY ("tool_id") +); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_ToolTable_tool_name_key" ON "LiteLLM_ToolTable"("tool_name"); + +-- CreateIndex +CREATE INDEX "LiteLLM_ToolTable_call_policy_idx" ON "LiteLLM_ToolTable"("call_policy"); + +-- CreateIndex +CREATE INDEX "LiteLLM_ToolTable_team_id_idx" ON "LiteLLM_ToolTable"("team_id"); + +-- AddForeignKey +ALTER TABLE "LiteLLM_AgentsTable" ADD CONSTRAINT "LiteLLM_AgentsTable_object_permission_id_fkey" FOREIGN KEY ("object_permission_id") REFERENCES "LiteLLM_ObjectPermissionTable"("object_permission_id") ON DELETE SET NULL ON UPDATE CASCADE; + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 4af7484148c..440c9c1d829 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -64,6 +64,8 @@ model LiteLLM_AgentsTable { litellm_params Json? agent_card_params Json agent_access_groups String[] @default([]) + object_permission_id String? + object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) created_at DateTime @default(now()) @map("created_at") created_by String updated_at DateTime @default(now()) @updatedAt @map("updated_at") @@ -264,6 +266,7 @@ model LiteLLM_ObjectPermissionTable { organizations LiteLLM_OrganizationTable[] users LiteLLM_UserTable[] end_users LiteLLM_EndUserTable[] + agents_table LiteLLM_AgentsTable[] } // Holds the MCP server configuration @@ -273,7 +276,6 @@ model LiteLLM_MCPServerTable { alias String? description String? url String? - spec_path String? transport String @default("sse") auth_type String? credentials Json? @default("{}") @@ -315,6 +317,7 @@ model LiteLLM_VerificationToken { router_settings Json? @default("{}") user_id String? team_id String? + agent_id String? project_id String? permissions Json @default("{}") max_parallel_requests Int? @@ -477,6 +480,7 @@ model LiteLLM_SpendLogs { completion_tokens Int @default(0) startTime DateTime // Assuming start_time is a DateTime field endTime DateTime // Assuming end_time is a DateTime field + request_duration_ms Int? completionStartTime DateTime? // Assuming completionStartTime is a DateTime field model String @default("") model_id String? @default("") // the model id stored in proxy model db @@ -1052,6 +1056,26 @@ model LiteLLM_PolicyAttachmentTable { updated_by String? } +// Global tool registry - auto-discovered from LLM responses; admins set call_policy here +model LiteLLM_ToolTable { + tool_id String @id @default(uuid()) + tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space" + origin String? // MCP server name or "user_defined" + call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked" + call_count Int @default(0) // cumulative number of times this tool was seen + assignments Json? @default("{}") + key_hash String? // hash of the virtual key that first called this tool + team_id String? // team that first called this tool + key_alias String? // human-readable alias of the virtual key + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@index([call_policy]) + @@index([team_id]) +} + //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 32425adb82f..bd57b248cdb 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm-proxy-extras" -version = "0.4.47" +version = "0.4.48" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." authors = ["BerriAI"] readme = "README.md" @@ -22,7 +22,7 @@ requires = ["poetry-core"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "0.4.47" +version = "0.4.48" version_files = [ "pyproject.toml:version", "../requirements.txt:litellm-proxy-extras==", diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 25f6e284bcd..9553d2c5246 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -37,7 +37,9 @@ from litellm.types.llms.openai import ( ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( + LIST_BATCHES_SUPPORTED_PROVIDERS, OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, + ListBatchesSupportedProvider, LiteLLMBatch, LlmProviders, ) @@ -674,7 +676,7 @@ def retrieve_batch( async def alist_batches( after: Optional[str] = None, limit: Optional[int] = None, - custom_llm_provider: Literal["openai", "azure", "hosted_vllm", "vertex_ai"] = "openai", + custom_llm_provider: ListBatchesSupportedProvider = "openai", metadata: Optional[Dict[str, str]] = None, extra_headers: Optional[Dict[str, str]] = None, extra_body: Optional[Dict[str, str]] = None, @@ -717,7 +719,7 @@ async def alist_batches( def list_batches( after: Optional[str] = None, limit: Optional[int] = None, - custom_llm_provider: Literal["openai", "azure", "hosted_vllm", "vertex_ai"] = "openai", + custom_llm_provider: ListBatchesSupportedProvider = "openai", extra_headers: Optional[Dict[str, str]] = None, extra_body: Optional[Dict[str, str]] = None, **kwargs, @@ -843,8 +845,9 @@ def list_batches( ) else: raise litellm.exceptions.BadRequestError( - message="LiteLLM doesn't support {} for 'list_batch'. Supported providers: openai, azure, vertex_ai.".format( - custom_llm_provider + message="LiteLLM doesn't support {} for 'list_batch'. Supported providers: {}.".format( + custom_llm_provider, + ", ".join(sorted(LIST_BATCHES_SUPPORTED_PROVIDERS)), ), model="n/a", llm_provider=custom_llm_provider, diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index dcc2df5f91c..fa9b94bc2ac 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -22,7 +22,11 @@ from litellm._logging import print_verbose, verbose_logger from litellm.constants import DEFAULT_REDIS_MAJOR_VERSION from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs from litellm.litellm_core_utils.coroutine_checker import coroutine_checker -from litellm.types.caching import RedisPipelineIncrementOperation +from litellm.types.caching import ( + RedisPipelineIncrementOperation, + RedisPipelineLpopOperation, + RedisPipelineRpushOperation, +) from litellm.types.services import ServiceTypes from .base_cache import BaseCache @@ -1320,6 +1324,75 @@ class RedisCache(BaseCache): ) raise e + async def _pipeline_rpush_helper( + self, + pipe: pipeline, + rpush_list: List[RedisPipelineRpushOperation], + ) -> List[int]: + """Helper function for pipeline rpush operations""" + for rpush_op in rpush_list: + pipe.rpush(rpush_op["key"], *rpush_op["values"]) + results = await pipe.execute() + # Preserve positional correspondence — raise on per-command errors + for r in results: + if isinstance(r, Exception): + raise r + return results + + async def async_rpush_pipeline( + self, + rpush_list: List[RedisPipelineRpushOperation], + ) -> List[int]: + """ + Use Redis Pipelines for bulk RPUSH operations + + Args: + rpush_list: List of RedisPipelineRpushOperation dicts containing: + - key: str + - values: List[Any] + + Returns: + List[int]: List lengths after each push + """ + if len(rpush_list) == 0: + return [] + + _redis_client: Any = self.init_async_client() + start_time = time.time() + + try: + async with _redis_client.pipeline(transaction=False) as pipe: + results = await self._pipeline_rpush_helper(pipe, rpush_list) + + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=_duration, + call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}", + ) + ) + return results + except Exception as e: + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=_duration, + error=e, + call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}", + ) + ) + verbose_logger.error( + "LiteLLM Redis Caching: async_rpush_pipeline() - Got exception from REDIS %s", + str(e), + ) + raise e + async def handle_lpop_count_for_older_redis_versions( self, pipe: pipeline, key: str, count: int ) -> List[bytes]: @@ -1400,3 +1473,120 @@ class RedisCache(BaseCache): f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {str(e)}" ) raise e + + async def _pipeline_lpop_helper( + self, + pipe: pipeline, + lpop_list: List[RedisPipelineLpopOperation], + ) -> List[Optional[List[str]]]: + """Helper function for pipeline lpop operations. + + For Redis >= 7, queues one LPOP(key, count) per operation. + For Redis < 7, queues `count` individual LPOP(key) commands per operation. + """ + major_version = self._parse_redis_major_version() + + if major_version >= 7: + for lpop_op in lpop_list: + pipe.lpop(lpop_op["key"], lpop_op["count"]) + raw_results = await pipe.execute() + else: + # For Redis < 7, LPOP doesn't support count param. + # Issue `count` individual LPOP commands per key, all in one pipeline. + counts: List[int] = [] + for lpop_op in lpop_list: + count = lpop_op["count"] or 1 + counts.append(count) + for _ in range(count): + pipe.lpop(lpop_op["key"]) + flat_results = await pipe.execute() + + # Re-group the flat results back into per-key lists + raw_results = [] + offset = 0 + for count in counts: + key_results = [ + r for r in flat_results[offset : offset + count] if r is not None + ] + raw_results.append(key_results if key_results else None) + offset += count + + # Raise on per-command errors (matches _pipeline_rpush_helper behavior) + for r in raw_results: + if isinstance(r, Exception): + raise r + + # Decode bytes -> str for each result set + decoded_results: List[Optional[List[str]]] = [] + for r in raw_results: + if r is None: + decoded_results.append(None) + elif isinstance(r, list): + try: + decoded_results.append( + [ + item.decode("utf-8") if isinstance(item, bytes) else item + for item in r + if item is not None + ] + or None + ) + except Exception: + decoded_results.append(r) # type: ignore + else: + decoded_results.append(None) + return decoded_results + + async def async_lpop_pipeline( + self, + lpop_list: List[RedisPipelineLpopOperation], + ) -> List[Optional[List[str]]]: + """ + Use Redis Pipelines for bulk LPOP operations + + Args: + lpop_list: List of RedisPipelineLpopOperation dicts containing: + - key: str + - count: Optional[int] + + Returns: + List[Optional[List[str]]]: Decoded results per key, None if key was empty + """ + if len(lpop_list) == 0: + return [] + + _redis_client: Any = self.init_async_client() + start_time = time.time() + + try: + async with _redis_client.pipeline(transaction=False) as pipe: + results = await self._pipeline_lpop_helper(pipe, lpop_list) + + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=_duration, + call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}", + ) + ) + return results + except Exception as e: + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=_duration, + error=e, + call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}", + ) + ) + verbose_logger.error( + "LiteLLM Redis Caching: async_lpop_pipeline() - Got exception from REDIS %s", + str(e), + ) + raise e diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index cc3916af069..47a27c8ef5b 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -4,6 +4,8 @@ Helper functions for health check calls. from typing import TYPE_CHECKING, Callable, Dict, Literal, Optional +from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging @@ -82,6 +84,27 @@ class HealthCheckHelpers: "tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME], } + @staticmethod + async def _batch_health_check( + custom_llm_provider: str, + model_params: dict, + filtered_model_params: dict, + ) -> dict: + """ + Health check for batch mode. + + Calls list_batches for providers that support it (openai, hosted_vllm, azure, + vertex_ai). For all other providers (e.g. bedrock) the batch API surface doesn't + include list_batches, so we fall back to acompletion to verify connectivity and + credential validity instead. + """ + import litellm + + if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS: + return await litellm.alist_batches(**filtered_model_params) + else: + return await litellm.acompletion(**model_params) + @staticmethod def get_mode_handlers( model: str, @@ -176,8 +199,10 @@ class HealthCheckHelpers: api_key=model_params.get("api_key", None), api_version=model_params.get("api_version", None), ), - "batch": lambda: litellm.alist_batches( - **_filter_model_params(model_params=model_params), + "batch": lambda: HealthCheckHelpers._batch_health_check( + custom_llm_provider=custom_llm_provider, + model_params=model_params, + filtered_model_params=_filter_model_params(model_params=model_params), ), "responses": lambda: litellm.aresponses( **_filter_model_params(model_params=model_params), diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index aa8471a5973..ab102a69670 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -162,6 +162,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): "service_tier", "safety_identifier", "prompt_cache_key", + "prompt_cache_retention", "store", ] # works across all models diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 60b29b975f7..860569d24cb 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -412,6 +412,26 @@ class MCPRequestHandler: ) return [] + ######################################################### + # Check agent permissions if agent_id is set on the key + ######################################################### + if user_api_key_auth and user_api_key_auth.agent_id: + allowed_mcp_servers_for_agent = ( + await MCPRequestHandler._get_allowed_mcp_servers_for_agent( + user_api_key_auth + ) + ) + if len(allowed_mcp_servers_for_agent) > 0: + # Intersect: agent can only use servers allowed by BOTH key/team AND agent config + allowed_mcp_servers = [ + s + for s in allowed_mcp_servers + if s in allowed_mcp_servers_for_agent + ] + verbose_logger.debug( + f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}" + ) + return list(set(allowed_mcp_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") @@ -513,13 +533,33 @@ class MCPRequestHandler: if team_tools: if key_tools: # Both have restrictions → intersection - return list(set(team_tools) & set(key_tools)) + allowed_tools = list(set(team_tools) & set(key_tools)) else: # Only team has restrictions → inherit from team - return team_tools + allowed_tools = team_tools else: # No team restrictions → use key restrictions - return key_tools + allowed_tools = key_tools + + # Intersect with agent's tool permissions if agent_id is set + if user_api_key_auth.agent_id: + # Pre-fetch agent object_permission once to avoid duplicate DB query + agent_obj_perm = await MCPRequestHandler._get_agent_object_permission( + user_api_key_auth + ) + agent_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( + server_id=server_id, + user_api_key_auth=user_api_key_auth, + agent_object_permission=agent_obj_perm, + ) + if agent_tools is not None: + if allowed_tools is not None: + allowed_tools = list( + set(allowed_tools) & set(agent_tools) + ) + else: + allowed_tools = agent_tools + return allowed_tools except Exception as e: verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}") @@ -715,6 +755,131 @@ class MCPRequestHandler: ) return [] + @staticmethod + async def _get_agent_object_permission( + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + ): + """ + Fetch the agent's object_permission from the DB (single query). + + Returns the object_permission object or None. + """ + from litellm.proxy.proxy_server import prisma_client + + if not user_api_key_auth or not user_api_key_auth.agent_id: + return None + + if prisma_client is None: + verbose_logger.debug("prisma_client is None") + return None + + try: + agent_row = await prisma_client.db.litellm_agentstable.find_unique( + where={"agent_id": user_api_key_auth.agent_id}, + include={"object_permission": True}, + ) + if agent_row is None or agent_row.object_permission is None: + return None + + return agent_row.object_permission + except Exception as e: + verbose_logger.warning( + f"Failed to get agent object permission: {str(e)}" + ) + return None + + @staticmethod + async def _get_allowed_mcp_servers_for_agent( + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + agent_object_permission=None, + ) -> List[str]: + """ + Get allowed MCP servers for an agent (from the agent's object_permission). + + Returns the MCP servers from the agent's object_permission. + If agent has no object_permission, returns [] (no extra restriction). + + Args: + user_api_key_auth: User auth with agent_id + agent_object_permission: Pre-fetched object_permission to avoid duplicate DB query. + If None, will be fetched from DB. + """ + if not user_api_key_auth or not user_api_key_auth.agent_id: + return [] + + try: + obj_perm = agent_object_permission + if obj_perm is None: + obj_perm = await MCPRequestHandler._get_agent_object_permission( + user_api_key_auth + ) + if obj_perm is None: + return [] + + direct_mcp_servers = getattr(obj_perm, "mcp_servers", None) or [] + if isinstance(direct_mcp_servers, str): + direct_mcp_servers = [] + mcp_access_groups = getattr(obj_perm, "mcp_access_groups", None) or [] + if isinstance(mcp_access_groups, str): + mcp_access_groups = [] + + access_group_servers = ( + await MCPRequestHandler._get_mcp_servers_from_access_groups( + mcp_access_groups + ) + ) + all_servers = list(direct_mcp_servers) + access_group_servers + return list(set(all_servers)) + except Exception as e: + verbose_logger.warning( + f"Failed to get allowed MCP servers for agent: {str(e)}" + ) + return [] + + @staticmethod + async def _get_agent_tool_permissions_for_server( + server_id: str, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + agent_object_permission=None, + ) -> Optional[List[str]]: + """ + Get allowed tool names for a server from the agent's object_permission. + Returns None if agent has no tool restrictions for this server. + + Args: + server_id: Server ID to check permissions for + user_api_key_auth: User auth with agent_id + agent_object_permission: Pre-fetched object_permission to avoid duplicate DB query. + If None, will be fetched from DB. + """ + if not user_api_key_auth or not user_api_key_auth.agent_id: + return None + + try: + obj_perm = agent_object_permission + if obj_perm is None: + obj_perm = await MCPRequestHandler._get_agent_object_permission( + user_api_key_auth + ) + if obj_perm is None: + return None + + mcp_tool_permissions = getattr( + obj_perm, "mcp_tool_permissions", None + ) + if not mcp_tool_permissions: + return None + if isinstance(mcp_tool_permissions, dict): + tools = mcp_tool_permissions.get(server_id) + else: + tools = None + return list(tools) if tools else None + except Exception as e: + verbose_logger.warning( + f"Failed to get agent tool permissions for server: {str(e)}" + ) + return None + @staticmethod def _get_config_server_ids_for_access_groups( config_mcp_servers, access_groups: List[str] diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 75b9f91acd9..aeb9950b11c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,60 +1,40 @@ import enum import json from datetime import datetime -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union +from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal, + Optional, Union) import httpx -from pydantic import ( - BaseModel, - ConfigDict, - Field, - Json, - field_validator, - model_validator, -) +from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator, + model_validator) from typing_extensions import Required, TypedDict from litellm._uuid import uuid from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import ( - AllMessageValues, - OpenAIFileObject, - ResponsesAPIResponse, -) -from litellm.types.mcp import ( - MCPAuth, - MCPAuthType, - MCPCredentials, - MCPTransport, - MCPTransportType, -) +from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject, + ResponsesAPIResponse) +from litellm.types.mcp import (MCPAuth, MCPAuthType, MCPCredentials, + MCPTransport, MCPTransportType) from litellm.types.mcp_server.mcp_server_manager import MCPInfo from litellm.types.router import RouterErrors, UpdateRouterConfig from litellm.types.secret_managers.main import KeyManagementSystem -from litellm.types.utils import ( - CallTypes, - CostBreakdown, - EmbeddingResponse, - GenericBudgetConfigType, - ImageResponse, - LiteLLMBatch, - LiteLLMFineTuningJob, - LiteLLMPydanticObjectBase, - ModelResponse, - ProviderField, - StandardCallbackDynamicParams, - StandardLoggingGuardrailInformation, - StandardLoggingMCPToolCall, - StandardLoggingModelInformation, - StandardLoggingPayloadErrorInformation, - StandardLoggingPayloadStatus, - StandardLoggingVectorStoreRequest, - StandardPassThroughResponseObject, - TextCompletionResponse, -) +from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse, + GenericBudgetConfigType, ImageResponse, + LiteLLMBatch, LiteLLMFineTuningJob, + LiteLLMPydanticObjectBase, ModelResponse, + ProviderField, StandardCallbackDynamicParams, + StandardLoggingGuardrailInformation, + StandardLoggingMCPToolCall, + StandardLoggingModelInformation, + StandardLoggingPayloadErrorInformation, + StandardLoggingPayloadStatus, + StandardLoggingVectorStoreRequest, + StandardPassThroughResponseObject, + TextCompletionResponse) from litellm.types.videos.main import VideoObject -from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type +from .types_utils.utils import (get_instance_fn, + validate_custom_validate_return_type) if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -851,6 +831,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): max_budget: Optional[float] = None user_id: Optional[str] = None team_id: Optional[str] = None + agent_id: Optional[str] = None max_parallel_requests: Optional[int] = None metadata: Optional[dict] = {} tpm_limit: Optional[int] = None @@ -2201,6 +2182,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase): config: Dict = {} user_id: Optional[str] = None team_id: Optional[str] = None + agent_id: Optional[str] = None project_id: Optional[str] = None max_parallel_requests: Optional[int] = None metadata: Dict = {} @@ -2378,7 +2360,8 @@ class UserAPIKeyAuth( This is used to track number of requests/spend for health check calls. """ - from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + from litellm.constants import \ + LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME return cls( api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, @@ -2410,7 +2393,8 @@ class UserAPIKeyAuth( This is used to track actions performed by automated system jobs. """ - from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + from litellm.constants import \ + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME return cls( api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, @@ -2801,7 +2785,8 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): @model_validator(mode="after") def mask_api_keys(self): - from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker + from litellm.litellm_core_utils.sensitive_data_masker import \ + SensitiveDataMasker masker = SensitiveDataMasker(sensitive_patterns={"key"}) @@ -3104,6 +3089,7 @@ class SpendLogsPayload(TypedDict): response: Optional[Union[str, list, dict]] proxy_server_request: Optional[str] session_id: Optional[str] + request_duration_ms: Optional[int] status: Literal["success", "failure"] diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 0d2df3856a1..159c9fb93d9 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -5,6 +5,9 @@ from typing import Any, Dict, List, Optional import litellm from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy.management_helpers.object_permission_utils import ( + handle_update_object_permission_common, +) from litellm.proxy.utils import PrismaClient from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest @@ -117,20 +120,39 @@ class AgentRegistry: ) agent_card_params: str = safe_dumps(agent_card_params_dict) + # Handle object_permission (MCP tool access for agent) + object_permission_id: Optional[str] = None + if agent.get("object_permission") is not None: + agent_copy = dict(agent) + object_permission_id = await handle_update_object_permission_common( + agent_copy, None, prisma_client + ) + + create_data: Dict[str, Any] = { + "agent_name": agent_name, + "litellm_params": litellm_params, + "agent_card_params": agent_card_params, + "created_by": created_by, + "updated_by": created_by, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + } + if object_permission_id is not None: + create_data["object_permission_id"] = object_permission_id + # Create agent in DB created_agent = await prisma_client.db.litellm_agentstable.create( - data={ - "agent_name": agent_name, - "litellm_params": litellm_params, - "agent_card_params": agent_card_params, - "created_by": created_by, - "updated_by": created_by, - "created_at": datetime.now(timezone.utc), - "updated_at": datetime.now(timezone.utc), - } + data=create_data, + include={"object_permission": True}, ) - return AgentResponse(**created_agent.model_dump()) # type: ignore + created_agent_dict = created_agent.model_dump() + if created_agent.object_permission is not None: + try: + created_agent_dict["object_permission"] = created_agent.object_permission.model_dump() + except Exception: + created_agent_dict["object_permission"] = created_agent.object_permission.dict() + return AgentResponse(**created_agent_dict) # type: ignore except Exception as e: raise Exception(f"Error adding agent to DB: {str(e)}") @@ -181,7 +203,7 @@ class AgentRegistry: raise Exception(f"Agent with ID {agent_id} not found") augment_agent = {**existing_agent, **agent} - update_data = {} + update_data: Dict[str, Any] = {} if augment_agent.get("agent_name"): update_data["agent_name"] = augment_agent.get("agent_name") if augment_agent.get("litellm_params"): @@ -192,6 +214,20 @@ class AgentRegistry: update_data["agent_card_params"] = safe_dumps( augment_agent.get("agent_card_params") ) + if agent.get("object_permission") is not None: + agent_copy = dict(augment_agent) + existing_object_permission_id = existing_agent.get( + "object_permission_id" + ) + object_permission_id = ( + await handle_update_object_permission_common( + agent_copy, + existing_object_permission_id, + prisma_client, + ) + ) + if object_permission_id is not None: + update_data["object_permission_id"] = object_permission_id # Patch agent in DB patched_agent = await prisma_client.db.litellm_agentstable.update( where={"agent_id": agent_id}, @@ -200,8 +236,15 @@ class AgentRegistry: "updated_by": updated_by, "updated_at": datetime.now(timezone.utc), }, + include={"object_permission": True}, ) - return AgentResponse(**patched_agent.model_dump()) # type: ignore + patched_agent_dict = patched_agent.model_dump() + if patched_agent.object_permission is not None: + try: + patched_agent_dict["object_permission"] = patched_agent.object_permission.model_dump() + except Exception: + patched_agent_dict["object_permission"] = patched_agent.object_permission.dict() + return AgentResponse(**patched_agent_dict) # type: ignore except Exception as e: raise Exception(f"Error patching agent in DB: {str(e)}") @@ -238,19 +281,47 @@ class AgentRegistry: ) agent_card_params: str = safe_dumps(agent_card_params_dict) + update_data: Dict[str, Any] = { + "agent_name": agent_name, + "litellm_params": litellm_params, + "agent_card_params": agent_card_params, + "updated_by": updated_by, + "updated_at": datetime.now(timezone.utc), + } + if agent.get("object_permission") is not None: + existing_agent = await prisma_client.db.litellm_agentstable.find_unique( + where={"agent_id": agent_id} + ) + existing_object_permission_id = ( + existing_agent.object_permission_id + if existing_agent is not None + else None + ) + agent_copy = dict(agent) + object_permission_id = ( + await handle_update_object_permission_common( + agent_copy, + existing_object_permission_id, + prisma_client, + ) + ) + if object_permission_id is not None: + update_data["object_permission_id"] = object_permission_id + # Update agent in DB updated_agent = await prisma_client.db.litellm_agentstable.update( where={"agent_id": agent_id}, - data={ - "agent_name": agent_name, - "litellm_params": litellm_params, - "agent_card_params": agent_card_params, - "updated_by": updated_by, - "updated_at": datetime.now(timezone.utc), - }, + data=update_data, + include={"object_permission": True}, ) - return AgentResponse(**updated_agent.model_dump()) # type: ignore + updated_agent_dict = updated_agent.model_dump() + if updated_agent.object_permission is not None: + try: + updated_agent_dict["object_permission"] = updated_agent.object_permission.model_dump() + except Exception: + updated_agent_dict["object_permission"] = updated_agent.object_permission.dict() + return AgentResponse(**updated_agent_dict) # type: ignore except Exception as e: raise Exception(f"Error updating agent in DB: {str(e)}") @@ -264,11 +335,19 @@ class AgentRegistry: try: agents_from_db = await prisma_client.db.litellm_agentstable.find_many( order={"created_at": "desc"}, + include={"object_permission": True}, ) agents: List[Dict[str, Any]] = [] for agent in agents_from_db: - agents.append(dict(agent)) + agent_dict = dict(agent) + # object_permission is eagerly loaded via include above + if agent.object_permission is not None: + try: + agent_dict["object_permission"] = agent.object_permission.model_dump() + except Exception: + agent_dict["object_permission"] = agent.object_permission.dict() + agents.append(agent_dict) return agents except Exception as e: diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 4a8d615f0b3..b411b81b434 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -16,6 +16,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity from litellm.types.agents import ( AgentConfig, AgentMakePublicResponse, @@ -23,8 +24,6 @@ from litellm.types.agents import ( MakeAgentsPublicRequest, PatchAgentRequest, ) - -from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, ) @@ -233,11 +232,18 @@ async def get_agent_by_id(agent_id: str): try: agent = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id) if agent is None: - agent = await prisma_client.db.litellm_agentstable.find_unique( - where={"agent_id": agent_id} + agent_row = await prisma_client.db.litellm_agentstable.find_unique( + where={"agent_id": agent_id}, + include={"object_permission": True}, ) - if agent is not None: - agent = AgentResponse(**agent.model_dump()) # type: ignore + if agent_row is not None: + agent_dict = agent_row.model_dump() + if agent_row.object_permission is not None: + try: + agent_dict["object_permission"] = agent_row.object_permission.model_dump() + except Exception: + agent_dict["object_permission"] = agent_row.object_permission.dict() + agent = AgentResponse(**agent_dict) # type: ignore if agent is None: raise HTTPException( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 8f17440773a..9240d84db83 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -593,9 +593,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 user_id=user_id, team_id=team_id, team_alias=( - team_object.team_alias - if team_object is not None - else None + team_object.team_alias if team_object is not None else None ), team_metadata=team_object.metadata if team_object is not None @@ -709,12 +707,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if isinstance(api_key, str): return UserAPIKeyAuth( api_key=api_key, - user_role=LitellmUserRoles.PROXY_ADMIN, + user_role=LitellmUserRoles.INTERNAL_USER, parent_otel_span=parent_otel_span, ) else: return UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN, + user_role=LitellmUserRoles.INTERNAL_USER, parent_otel_span=parent_otel_span, ) elif api_key is None: # only require api key if master key is set @@ -846,7 +844,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) valid_token.parent_otel_span = parent_otel_span if _end_user_object is not None: - valid_token.end_user_object_permission = _end_user_object.object_permission + valid_token.end_user_object_permission = ( + _end_user_object.object_permission + ) return valid_token @@ -954,7 +954,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if isinstance( api_key, str ): # if generated token, make sure it starts with sk-. - _masked_key = "{}****{}".format(api_key[:4], api_key[-4:]) if len(api_key) > 8 else "****" + _masked_key = ( + "{}****{}".format(api_key[:4], api_key[-4:]) + if len(api_key) > 8 + else "****" + ) assert api_key.startswith( "sk-" ), "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format( @@ -1304,9 +1308,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if _end_user_object is not None: valid_token_dict.update(end_user_params) - valid_token_dict["end_user_object_permission"] = ( - _end_user_object.object_permission - ) + valid_token_dict[ + "end_user_object_permission" + ] = _end_user_object.object_permission # check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions # sso/login, ui/login, /key functions and /user functions diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 456bf4022d5..37fac8b56ed 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -760,9 +760,16 @@ class DBSpendUpdateWriter: verbose_proxy_logger.debug("acquired lock for spend updates") try: - db_spend_update_transactions = ( - await self.redis_update_buffer.get_all_update_transactions_from_redis_buffer() - ) + ( + db_spend_update_transactions, + daily_spend_update_transactions, + daily_team_spend_update_transactions, + daily_org_spend_update_transactions, + daily_end_user_spend_update_transactions, + daily_agent_spend_update_transactions, + daily_tag_spend_update_transactions, + ) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline() + if db_spend_update_transactions is not None: verbose_proxy_logger.info( "Spend tracking - committing spend updates from Redis to DB: " @@ -782,9 +789,6 @@ class DBSpendUpdateWriter: db_spend_update_transactions=db_spend_update_transactions, ) - daily_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_spend_update_transactions_from_redis_buffer() - ) if daily_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_user_spend( n_retry_times=n_retry_times, @@ -792,9 +796,6 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_spend_update_transactions, ) - daily_team_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_team_spend_update_transactions_from_redis_buffer() - ) if daily_team_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_team_spend( n_retry_times=n_retry_times, @@ -803,9 +804,6 @@ class DBSpendUpdateWriter: daily_spend_transactions=daily_team_spend_update_transactions, ) - daily_org_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_org_spend_update_transactions_from_redis_buffer() - ) if daily_org_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_org_spend( n_retry_times=n_retry_times, @@ -814,9 +812,6 @@ class DBSpendUpdateWriter: daily_spend_transactions=daily_org_spend_update_transactions, ) - daily_tag_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() - ) if daily_tag_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_tag_spend( n_retry_times=n_retry_times, @@ -824,9 +819,6 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_tag_spend_update_transactions, ) - daily_end_user_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_end_user_spend_update_transactions_from_redis_buffer() - ) if daily_end_user_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_end_user_spend( n_retry_times=n_retry_times, @@ -834,9 +826,6 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_end_user_spend_update_transactions, ) - daily_agent_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_agent_spend_update_transactions_from_redis_buffer() - ) if daily_agent_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_agent_spend( n_retry_times=n_retry_times, diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 027f3e639e3..51201f96d77 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -6,7 +6,7 @@ This is to prevent deadlocks and improve reliability import asyncio import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache @@ -36,6 +36,7 @@ from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import ( ) from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue from litellm.secret_managers.main import str_to_bool +from litellm.types.caching import RedisPipelineLpopOperation, RedisPipelineRpushOperation from litellm.types.services import ServiceTypes if TYPE_CHECKING: @@ -209,47 +210,44 @@ class RedisUpdateBuffer: "ALL DAILY SPEND UPDATE TRANSACTIONS: %s", daily_spend_update_transactions ) - await self._store_transactions_in_redis( - transactions=db_spend_update_transactions, - redis_key=REDIS_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_SPEND_UPDATE_QUEUE, + # Build a list of rpush operations, skipping empty/None transaction sets + _queue_configs: List[Tuple[Any, str, ServiceTypes]] = [ + (db_spend_update_transactions, REDIS_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_SPEND_UPDATE_QUEUE), + (daily_spend_update_transactions, REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_SPEND_UPDATE_QUEUE), + (daily_team_spend_update_transactions, REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_TEAM_SPEND_UPDATE_QUEUE), + (daily_org_spend_update_transactions, REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_ORG_SPEND_UPDATE_QUEUE), + (daily_end_user_spend_update_transactions, REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_END_USER_SPEND_UPDATE_QUEUE), + (daily_agent_spend_update_transactions, REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_AGENT_SPEND_UPDATE_QUEUE), + (daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_TAG_SPEND_UPDATE_QUEUE), + ] + + rpush_list: List[RedisPipelineRpushOperation] = [] + service_types: List[ServiceTypes] = [] + for transactions, redis_key, service_type in _queue_configs: + if transactions is None or len(transactions) == 0: + continue + rpush_list.append( + RedisPipelineRpushOperation( + key=redis_key, + values=[safe_dumps(transactions)], + ) + ) + service_types.append(service_type) + + if len(rpush_list) == 0: + return + + result_lengths = await self.redis_cache.async_rpush_pipeline( + rpush_list=rpush_list, ) - await self._store_transactions_in_redis( - transactions=daily_spend_update_transactions, - redis_key=REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_DAILY_SPEND_UPDATE_QUEUE, - ) - - await self._store_transactions_in_redis( - transactions=daily_team_spend_update_transactions, - redis_key=REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_DAILY_TEAM_SPEND_UPDATE_QUEUE, - ) - - await self._store_transactions_in_redis( - transactions=daily_org_spend_update_transactions, - redis_key=REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_DAILY_SPEND_UPDATE_QUEUE, - ) - - await self._store_transactions_in_redis( - transactions=daily_end_user_spend_update_transactions, - redis_key=REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_DAILY_END_USER_SPEND_UPDATE_QUEUE, - ) - - await self._store_transactions_in_redis( - transactions=daily_agent_spend_update_transactions, - redis_key=REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_DAILY_AGENT_SPEND_UPDATE_QUEUE, - ) - - await self._store_transactions_in_redis( - transactions=daily_tag_spend_update_transactions, - redis_key=REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY, - service_type=ServiceTypes.REDIS_DAILY_TAG_SPEND_UPDATE_QUEUE, - ) + # Emit gauge events for each queue + for i, queue_size in enumerate(result_lengths): + if i < len(service_types): + await self._emit_new_item_added_to_redis_buffer_event( + queue_size=queue_size, + service=service_types[i], + ) @staticmethod def _number_of_transactions_to_store_in_redis( @@ -338,6 +336,77 @@ class RedisUpdateBuffer: return combined_transaction + async def get_all_transactions_from_redis_buffer_pipeline( + self, + ) -> Tuple[ + Optional[DBSpendUpdateTransactions], + Optional[Dict[str, DailyUserSpendTransaction]], + Optional[Dict[str, DailyTeamSpendTransaction]], + Optional[Dict[str, DailyOrganizationSpendTransaction]], + Optional[Dict[str, DailyEndUserSpendTransaction]], + Optional[Dict[str, DailyAgentSpendTransaction]], + Optional[Dict[str, DailyTagSpendTransaction]], + ]: + """ + Drains all 7 Redis buffer queues in a single pipeline round-trip. + + Returns a 7-tuple of parsed results in this order: + 0: DBSpendUpdateTransactions + 1: daily user spend + 2: daily team spend + 3: daily org spend + 4: daily end-user spend + 5: daily agent spend + 6: daily tag spend + """ + if self.redis_cache is None: + return None, None, None, None, None, None, None + + lpop_list: List[RedisPipelineLpopOperation] = [ + RedisPipelineLpopOperation(key=REDIS_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + RedisPipelineLpopOperation(key=REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + RedisPipelineLpopOperation(key=REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + RedisPipelineLpopOperation(key=REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + RedisPipelineLpopOperation(key=REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + RedisPipelineLpopOperation(key=REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + RedisPipelineLpopOperation(key=REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT), + ] + + raw_results = await self.redis_cache.async_lpop_pipeline(lpop_list=lpop_list) + + # Pad with None if pipeline returned fewer results than expected + while len(raw_results) < 7: + raw_results.append(None) + + # Slot 0: DBSpendUpdateTransactions + db_spend: Optional[DBSpendUpdateTransactions] = None + if raw_results[0] is not None: + parsed = self._parse_list_of_transactions(raw_results[0]) + if len(parsed) > 0: + db_spend = self._combine_list_of_transactions(parsed) + + # Slots 1-6: daily spend categories + daily_results: List[Optional[Dict[str, Any]]] = [] + for slot in range(1, 7): + if raw_results[slot] is None: + daily_results.append(None) + else: + list_of_daily = [json.loads(t) for t in raw_results[slot]] # type: ignore + aggregated = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( + list_of_daily + ) + daily_results.append(aggregated) + + return ( + db_spend, + cast(Optional[Dict[str, DailyUserSpendTransaction]], daily_results[0]), + cast(Optional[Dict[str, DailyTeamSpendTransaction]], daily_results[1]), + cast(Optional[Dict[str, DailyOrganizationSpendTransaction]], daily_results[2]), + cast(Optional[Dict[str, DailyEndUserSpendTransaction]], daily_results[3]), + cast(Optional[Dict[str, DailyAgentSpendTransaction]], daily_results[4]), + cast(Optional[Dict[str, DailyTagSpendTransaction]], daily_results[5]), + ) + async def get_all_daily_spend_update_transactions_from_redis_buffer( self, ) -> Optional[Dict[str, DailyUserSpendTransaction]]: diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index 5c41b371ca4..9768d93e922 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -1,6 +1,10 @@ +from datetime import datetime + from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import ORJSONResponse, StreamingResponse +import litellm +from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -17,7 +21,8 @@ router = APIRouter( dependencies=[Depends(user_api_key_auth)], ) @router.post( - "/models/{model_name:path}:generateContent", dependencies=[Depends(user_api_key_auth)] + "/models/{model_name:path}:generateContent", + dependencies=[Depends(user_api_key_auth)], ) async def google_generate_content( request: Request, @@ -36,12 +41,12 @@ async def google_generate_content( data = await _read_request_body(request=request) if "model" not in data: data["model"] = model_name - + # Extract generationConfig and pass it as config parameter generation_config = data.pop("generationConfig", None) if generation_config: data["config"] = generation_config - + # Add user authentication metadata for cost tracking data = await add_litellm_data_to_request( data=data, @@ -51,7 +56,19 @@ async def google_generate_content( general_settings=general_settings, version=version, ) - + + # Create logging object with full request metadata so callbacks (e.g. S3) get user/trace_id + data["litellm_call_id"] = request.headers.get( + "x-litellm-call-id", str(uuid.uuid4()) + ) + logging_obj, data = litellm.utils.function_setup( + original_function="agenerate_content", + rules_obj=litellm.utils.Rules(), + start_time=datetime.now(), + **data, + ) + data["litellm_logging_obj"] = logging_obj + # call router if llm_router is None: raise HTTPException(status_code=500, detail="Router not initialized") @@ -103,6 +120,18 @@ async def google_stream_generate_content( version=version, ) + # Create logging object with full request metadata so streaming END callbacks (e.g. S3) get user/trace_id + data["litellm_call_id"] = request.headers.get( + "x-litellm-call-id", str(uuid.uuid4()) + ) + logging_obj, data = litellm.utils.function_setup( + original_function="agenerate_content_stream", + rules_obj=litellm.utils.Rules(), + start_time=datetime.now(), + **data, + ) + data["litellm_logging_obj"] = logging_obj + # call router if llm_router is None: raise HTTPException(status_code=500, detail="Router not initialized") @@ -247,11 +276,11 @@ async def create_interaction( ) data = await _read_request_body(request=request) - + # Default to gemini provider for interactions if "custom_llm_provider" not in data: data["custom_llm_provider"] = "gemini" - + processor = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( @@ -301,7 +330,7 @@ async def get_interaction( ): """ Get an interaction by ID. - + Per OpenAPI spec: GET /{api_version}/interactions/{interaction_id} """ from litellm.proxy.proxy_server import ( @@ -319,7 +348,7 @@ async def get_interaction( ) data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"} - + processor = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( @@ -369,7 +398,7 @@ async def delete_interaction( ): """ Delete an interaction by ID. - + Per OpenAPI spec: DELETE /{api_version}/interactions/{interaction_id} """ from litellm.proxy.proxy_server import ( @@ -387,7 +416,7 @@ async def delete_interaction( ) data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"} - + processor = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( @@ -437,7 +466,7 @@ async def cancel_interaction( ): """ Cancel an interaction by ID. - + Per OpenAPI spec: POST /{api_version}/interactions/{interaction_id}:cancel """ from litellm.proxy.proxy_server import ( @@ -455,7 +484,7 @@ async def cancel_interaction( ) data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"} - + processor = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index c083c60cb4c..20f6e6f1d39 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -2,6 +2,7 @@ CRUD ENDPOINTS FOR GUARDRAILS """ +import concurrent.futures import inspect from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast @@ -11,9 +12,16 @@ from pydantic import BaseModel from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.integrations.custom_guardrail import CustomGuardrail -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry +from litellm.proxy.guardrails.guardrail_hooks.custom_code.code_validator import ( + CustomCodeValidationError, + validate_custom_code, +) +from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import ( + get_custom_code_primitives, +) from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router from litellm.types.guardrails import ( PII_ENTITY_CATEGORIES_MAP, @@ -243,9 +251,11 @@ class CreateGuardrailRequest(BaseModel): @router.post( "/guardrails", tags=["Guardrails"], - dependencies=[Depends(user_api_key_auth)], ) -async def create_guardrail(request: CreateGuardrailRequest): +async def create_guardrail( + request: CreateGuardrailRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): """ Create a new guardrail @@ -296,6 +306,12 @@ async def create_guardrail(request: CreateGuardrailRequest): from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail="Admin access required to manage guardrails", + ) + if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") @@ -332,9 +348,12 @@ class UpdateGuardrailRequest(BaseModel): @router.put( "/guardrails/{guardrail_id}", tags=["Guardrails"], - dependencies=[Depends(user_api_key_auth)], ) -async def update_guardrail(guardrail_id: str, request: UpdateGuardrailRequest): +async def update_guardrail( + guardrail_id: str, + request: UpdateGuardrailRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): """ Update an existing guardrail @@ -385,6 +404,12 @@ async def update_guardrail(guardrail_id: str, request: UpdateGuardrailRequest): from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail="Admin access required to manage guardrails", + ) + if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") @@ -429,9 +454,11 @@ async def update_guardrail(guardrail_id: str, request: UpdateGuardrailRequest): @router.delete( "/guardrails/{guardrail_id}", tags=["Guardrails"], - dependencies=[Depends(user_api_key_auth)], ) -async def delete_guardrail(guardrail_id: str): +async def delete_guardrail( + guardrail_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): """ Delete a guardrail @@ -453,6 +480,12 @@ async def delete_guardrail(guardrail_id: str): from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail="Admin access required to manage guardrails", + ) + if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") @@ -495,9 +528,12 @@ async def delete_guardrail(guardrail_id: str): @router.patch( "/guardrails/{guardrail_id}", tags=["Guardrails"], - dependencies=[Depends(user_api_key_auth)], ) -async def patch_guardrail(guardrail_id: str, request: PatchGuardrailRequest): +async def patch_guardrail( + guardrail_id: str, + request: PatchGuardrailRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): """ Partially update an existing guardrail @@ -546,6 +582,12 @@ async def patch_guardrail(guardrail_id: str, request: PatchGuardrailRequest): from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER from litellm.proxy.proxy_server import prisma_client + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail="Admin access required to manage guardrails", + ) + if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") @@ -1302,9 +1344,9 @@ async def get_provider_specific_params(): lakera_v2_fields = _get_fields_from_model(LakeraV2GuardrailConfigModel) tool_permission_fields = _get_fields_from_model(ToolPermissionGuardrailConfigModel) - tool_permission_fields["ui_friendly_name"] = ( - ToolPermissionGuardrailConfigModel.ui_friendly_name() - ) + tool_permission_fields[ + "ui_friendly_name" + ] = ToolPermissionGuardrailConfigModel.ui_friendly_name() # Return the provider-specific parameters provider_params = { @@ -1364,10 +1406,12 @@ class TestCustomCodeGuardrailResponse(BaseModel): @router.post( "/guardrails/test_custom_code", tags=["Guardrails"], - dependencies=[Depends(user_api_key_auth)], response_model=TestCustomCodeGuardrailResponse, ) -async def test_custom_code_guardrail(request: TestCustomCodeGuardrailRequest): +async def test_custom_code_guardrail( + request: TestCustomCodeGuardrailRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): """ Test custom code guardrail logic without creating a guardrail. @@ -1440,63 +1484,26 @@ async def test_custom_code_guardrail(request: TestCustomCodeGuardrailRequest): } ``` """ - import concurrent.futures - import re - from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import ( - get_custom_code_primitives, - ) - - # Security validation patterns - FORBIDDEN_PATTERNS = [ - # Import statements - (r"\bimport\s+", "import statements are not allowed"), - (r"\bfrom\s+\w+\s+import\b", "from...import statements are not allowed"), - (r"__import__\s*\(", "__import__() is not allowed"), - # Dangerous builtins - (r"\bexec\s*\(", "exec() is not allowed"), - (r"\beval\s*\(", "eval() is not allowed"), - (r"\bcompile\s*\(", "compile() is not allowed"), - (r"\bopen\s*\(", "open() is not allowed"), - (r"\bgetattr\s*\(", "getattr() is not allowed"), - (r"\bsetattr\s*\(", "setattr() is not allowed"), - (r"\bdelattr\s*\(", "delattr() is not allowed"), - (r"\bglobals\s*\(", "globals() is not allowed"), - (r"\blocals\s*\(", "locals() is not allowed"), - (r"\bvars\s*\(", "vars() is not allowed"), - (r"\bdir\s*\(", "dir() is not allowed"), - (r"\bbreakpoint\s*\(", "breakpoint() is not allowed"), - (r"\binput\s*\(", "input() is not allowed"), - # Dangerous dunder access - (r"__builtins__", "__builtins__ access is not allowed"), - (r"__globals__", "__globals__ access is not allowed"), - (r"__code__", "__code__ access is not allowed"), - (r"__subclasses__", "__subclasses__ access is not allowed"), - (r"__bases__", "__bases__ access is not allowed"), - (r"__mro__", "__mro__ access is not allowed"), - (r"__class__", "__class__ access is not allowed"), - (r"__dict__", "__dict__ access is not allowed"), - (r"__getattribute__", "__getattribute__ access is not allowed"), - (r"__reduce__", "__reduce__ access is not allowed"), - (r"__reduce_ex__", "__reduce_ex__ access is not allowed"), - # OS/system access - (r"\bos\.", "os module access is not allowed"), - (r"\bsys\.", "sys module access is not allowed"), - (r"\bsubprocess\.", "subprocess module access is not allowed"), - ] + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail="Admin access required to test custom code guardrails", + ) EXECUTION_TIMEOUT_SECONDS = 5 try: # Step 0: Security validation - check for forbidden patterns - code = request.custom_code - for pattern, error_msg in FORBIDDEN_PATTERNS: - if re.search(pattern, code): - return TestCustomCodeGuardrailResponse( - success=False, - error=f"Security violation: {error_msg}", - error_type="compilation", - ) + + try: + validate_custom_code(request.custom_code) + except CustomCodeValidationError as e: + return TestCustomCodeGuardrailResponse( + success=False, + error=str(e), + error_type="compilation", + ) # Step 1: Compile the custom code with restricted environment exec_globals = get_custom_code_primitives().copy() @@ -1612,10 +1619,10 @@ async def apply_guardrail( from litellm.proxy.utils import handle_exception_on_proxy try: - active_guardrail: Optional[CustomGuardrail] = ( - GUARDRAIL_REGISTRY.get_initialized_guardrail_callback( - guardrail_name=request.guardrail_name - ) + active_guardrail: Optional[ + CustomGuardrail + ] = GUARDRAIL_REGISTRY.get_initialized_guardrail_callback( + guardrail_name=request.guardrail_name ) if active_guardrail is None: raise HTTPException( diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/code_validator.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/code_validator.py new file mode 100644 index 00000000000..6ef59b522a8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/code_validator.py @@ -0,0 +1,63 @@ +import re +from typing import List, Tuple + +# Security validation patterns +FORBIDDEN_PATTERNS: List[Tuple[str, str]] = [ + # Import statements + (r"\bimport\s+", "import statements are not allowed"), + (r"\bfrom\s+\w+\s+import\b", "from...import statements are not allowed"), + (r"__import__\s*\(", "__import__() is not allowed"), + # Dangerous builtins + (r"\bexec\s*\(", "exec() is not allowed"), + (r"\beval\s*\(", "eval() is not allowed"), + (r"\bcompile\s*\(", "compile() is not allowed"), + (r"\bopen\s*\(", "open() is not allowed"), + (r"\bgetattr\s*\(", "getattr() is not allowed"), + (r"\bsetattr\s*\(", "setattr() is not allowed"), + (r"\bdelattr\s*\(", "delattr() is not allowed"), + (r"\bglobals\s*\(", "globals() is not allowed"), + (r"\blocals\s*\(", "locals() is not allowed"), + (r"\bvars\s*\(", "vars() is not allowed"), + (r"\bdir\s*\(", "dir() is not allowed"), + (r"\bbreakpoint\s*\(", "breakpoint() is not allowed"), + (r"\binput\s*\(", "input() is not allowed"), + # Dangerous dunder access + (r"__builtins__", "__builtins__ access is not allowed"), + (r"__globals__", "__globals__ access is not allowed"), + (r"__code__", "__code__ access is not allowed"), + (r"__subclasses__", "__subclasses__ access is not allowed"), + (r"__bases__", "__bases__ access is not allowed"), + (r"__mro__", "__mro__ access is not allowed"), + (r"__class__", "__class__ access is not allowed"), + (r"__dict__", "__dict__ access is not allowed"), + (r"__getattribute__", "__getattribute__ access is not allowed"), + (r"__reduce__", "__reduce__ access is not allowed"), + (r"__reduce_ex__", "__reduce_ex__ access is not allowed"), + # OS/system access + (r"\bos\.", "os module access is not allowed"), + (r"\bsys\.", "sys module access is not allowed"), + (r"\bsubprocess\.", "subprocess module access is not allowed"), + (r"\bshutil\.", "shutil module access is not allowed"), + (r"\bctypes\.", "ctypes module access is not allowed"), + (r"\bsocket\.", "socket module access is not allowed"), + (r"\bpickle\.", "pickle module access is not allowed"), +] + + +class CustomCodeValidationError(Exception): + """Raised when custom code fails security validation.""" + + pass + + +def validate_custom_code(code: str) -> None: + """ + Validate custom code against forbidden patterns. + + Raises CustomCodeValidationError if any forbidden pattern is found. + """ + if not code: + return + for pattern, error_msg in FORBIDDEN_PATTERNS: + if re.search(pattern, code): + raise CustomCodeValidationError(f"Security violation: {error_msg}") diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py index c557a093c4e..0f5a4384d76 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -41,18 +41,19 @@ from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast from fastapi import HTTPException from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_guardrail import (CustomGuardrail, - log_guardrail_information) +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.proxy.guardrails.guardrail_hooks.base import \ - GuardrailConfigModel +from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel from litellm.types.utils import GenericGuardrailAPIInputs +from .code_validator import CustomCodeValidationError, validate_custom_code from .primitives import get_custom_code_primitives if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import \ - Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj class CustomCodeGuardrailError(Exception): @@ -143,6 +144,33 @@ class CustomCodeGuardrail(CustomGuardrail): """Returns the config model for the UI.""" return CustomCodeGuardrailConfigModel + def _do_compile(self) -> None: + """Internal compilation method without lock. Expected to run inside _compile_lock.""" + # Create a restricted execution environment + # Only include our safe primitives + exec_globals = get_custom_code_primitives().copy() + + # CRITICAL: Restrict __builtins__ to prevent sandbox escape + exec_globals["__builtins__"] = {} + + # Execute the user code in the restricted environment + exec(compile(self.custom_code, "", "exec"), exec_globals) + + # Extract the apply_guardrail function + if "apply_guardrail" not in exec_globals: + raise CustomCodeCompilationError( + "Custom code must define an 'apply_guardrail' function. " + "Expected signature: apply_guardrail(inputs, request_data, input_type)" + ) + + apply_fn = exec_globals["apply_guardrail"] + if not callable(apply_fn): + raise CustomCodeCompilationError( + "'apply_guardrail' must be a callable function" + ) + + self._compiled_function = apply_fn + def _compile_custom_code(self) -> None: """ Compile the custom code and extract the apply_guardrail function. @@ -154,27 +182,14 @@ class CustomCodeGuardrail(CustomGuardrail): return try: - # Create a restricted execution environment - # Only include our safe primitives - exec_globals = get_custom_code_primitives().copy() + # Step 1: Security validation — forbidden pattern check + try: + validate_custom_code(self.custom_code) + except CustomCodeValidationError as e: + raise CustomCodeCompilationError(str(e)) from e - # Execute the user code in the restricted environment - exec(compile(self.custom_code, "", "exec"), exec_globals) - - # Extract the apply_guardrail function - if "apply_guardrail" not in exec_globals: - raise CustomCodeCompilationError( - "Custom code must define an 'apply_guardrail' function. " - "Expected signature: apply_guardrail(inputs, request_data, input_type)" - ) - - apply_fn = exec_globals["apply_guardrail"] - if not callable(apply_fn): - raise CustomCodeCompilationError( - "'apply_guardrail' must be a callable function" - ) - - self._compiled_function = apply_fn + # Step 2: Compile logic + self._do_compile() verbose_proxy_logger.debug( f"Custom code guardrail '{self.guardrail_name}' compiled successfully" ) @@ -390,6 +405,12 @@ class CustomCodeGuardrail(CustomGuardrail): Raises: CustomCodeCompilationError: If the new code fails to compile """ + # Validate BEFORE acquiring lock / resetting state + try: + validate_custom_code(new_code) + except CustomCodeValidationError as e: + raise CustomCodeCompilationError(str(e)) from e + with self._compile_lock: # Reset state old_function = self._compiled_function @@ -399,12 +420,24 @@ class CustomCodeGuardrail(CustomGuardrail): try: self.custom_code = new_code - self._compile_custom_code() + self._do_compile() verbose_proxy_logger.info( f"Custom code guardrail '{self.guardrail_name}': Code updated successfully" ) + except SyntaxError as e: + # Rollback on failure + self.custom_code = old_code + self._compiled_function = old_function + self._compile_error = f"Syntax error in custom code: {e}" + raise CustomCodeCompilationError(self._compile_error) from e except CustomCodeCompilationError: # Rollback on failure self.custom_code = old_code self._compiled_function = old_function raise + except Exception as e: + # Rollback on failure + self.custom_code = old_code + self._compiled_function = old_function + self._compile_error = f"Failed to compile custom code: {e}" + raise CustomCodeCompilationError(self._compile_error) from e diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 28ccff0e36c..34fbf47253b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -322,12 +322,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): analyze_payload, ) - async with session.post(analyze_url, json=analyze_payload) as response: - analyze_results = await response.json() - verbose_proxy_logger.debug("analyze_results: %s", analyze_results) - - # Handle error responses from Presidio (e.g., {'error': 'No text provided'}) - # Presidio may return a dict instead of a list when errors occur def _fail_on_invalid_response( reason: str, ) -> List[PresidioAnalyzeResponseItem]: @@ -347,6 +341,36 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): ) return [] + async with session.post( + analyze_url, + json=analyze_payload, + headers={"Accept": "application/json"}, + ) as response: + # Validate HTTP status + if response.status >= 400: + error_body = await response.text() + return _fail_on_invalid_response( + f"HTTP {response.status} from Presidio analyzer: {error_body[:200]}" + ) + + # Validate Content-Type is JSON + content_type = getattr( + response, + "content_type", + response.headers.get("Content-Type", ""), + ) + if "application/json" not in content_type: + error_body = await response.text() + return _fail_on_invalid_response( + f"expected application/json Content-Type but received '{content_type}'; body: '{error_body[:200]}'" + ) + + analyze_results = await response.json() + verbose_proxy_logger.debug("analyze_results: %s", analyze_results) + + # Handle error responses from Presidio (e.g., {'error': 'No text provided'}) + # Presidio may return a dict instead of a list when errors occur + if isinstance(analyze_results, dict): if "error" in analyze_results: return _fail_on_invalid_response( @@ -423,8 +447,29 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): } async with session.post( - anonymize_url, json=anonymize_payload + anonymize_url, + json=anonymize_payload, + headers={"Accept": "application/json"}, ) as response: + # Validate HTTP status + if response.status >= 400: + error_body = await response.text() + raise Exception( + f"Presidio anonymizer returned HTTP {response.status}: {error_body[:200]}" + ) + + # Validate Content-Type is JSON + content_type = getattr( + response, + "content_type", + response.headers.get("Content-Type", ""), + ) + if "application/json" not in content_type: + error_body = await response.text() + raise Exception( + f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'" + ) + redacted_text = await response.json() new_text = text @@ -456,7 +501,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): except Exception as e: # Sanitize exception to avoid leaking the original text (which may # contain API keys or other secrets) in error responses. - if "Invalid anonymizer response" in str(e): + error_str = str(e) + if ( + "Invalid anonymizer response" in error_str + or "Presidio anonymizer returned" in error_str + ): raise raise Exception( f"Presidio PII anonymization failed: {type(e).__name__}" diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index de4517efc88..d228bdb2129 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -283,12 +283,17 @@ async def perform_health_check( model: Optional[str] = None, cli_model: Optional[str] = None, details: Optional[bool] = True, + model_id: Optional[str] = None, max_concurrency: Optional[int] = None, instrumentation_context: Optional[dict] = None, ): """ Perform a health check on the system. + When model_id is provided, only the deployment with that id is checked + (so models that share the same name but have different ids are checked separately). + When model (name) is provided, all deployments matching that name are checked. + Returns: (bool): True if the health check passes, False otherwise. """ @@ -314,7 +319,12 @@ async def perform_health_check( cycle_start_time = time.monotonic() requested_model_count = len(model_list) - if model is not None: + # Filter by model_id first so a single deployment is checked when id is specified + if model_id is not None: + _by_id = [x for x in model_list if (x.get("model_info") or {}).get("id") == model_id] + if _by_id: + model_list = _by_id + elif model is not None: _new_model_list = [ x for x in model_list if x["litellm_params"]["model"] == model ] diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 3570b2dd6aa..f3bed3656f6 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -243,7 +243,7 @@ async def health_services_endpoint( # noqa: PLR0915 service_in_success_callbacks = True else: for cb in litellm.success_callback: - if hasattr(cb, "callback_name") and cb.callback_name == service: + if getattr(cb, "callback_name", None) == service: service_in_success_callbacks = True break cb_id = get_callback_identifier(cb) @@ -747,6 +747,7 @@ async def _perform_health_check_and_save( model=target_model, details=details, max_concurrency=max_concurrency, + model_id=model_id, ) # Optionally save health check result to database (non-blocking) diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py new file mode 100644 index 00000000000..8d481f6b261 --- /dev/null +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -0,0 +1,208 @@ +""" +Max Iterations Limiter for LiteLLM Proxy. + +Enforces a per-session cap on the number of LLM calls an agentic loop can make. +Callers send a `session_id` with each request (via `x-litellm-session-id` header +or `metadata.session_id`), and this hook counts calls per session. When the count +exceeds `max_iterations` (configured in key/team metadata), returns 429. + +Works across multiple proxy instances via DualCache (in-memory + Redis). +Follows the same pattern as parallel_request_limiter_v3.py. +""" + +import os +from typing import TYPE_CHECKING, Any, Optional, Union + +from fastapi import HTTPException + +from litellm import DualCache +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth + +if TYPE_CHECKING: + from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache + + InternalUsageCache = _InternalUsageCache +else: + InternalUsageCache = Any + + +# Redis Lua script for atomic increment with TTL. +# Returns the new count after increment. +# Only sets EXPIRE on first increment (when count becomes 1). +MAX_ITERATIONS_INCREMENT_SCRIPT = """ +local key = KEYS[1] +local ttl = tonumber(ARGV[1]) + +local current = redis.call('INCR', key) +if current == 1 then + redis.call('EXPIRE', key, ttl) +end + +return current +""" + +# Default TTL for session iteration counters (1 hour) +DEFAULT_MAX_ITERATIONS_TTL = 3600 + + +class _PROXY_MaxIterationsHandler(CustomLogger): + """ + Pre-call hook that enforces max_iterations per session. + + Configuration: + - max_iterations: set in key metadata via /key/generate or /key/update + e.g. metadata={"max_iterations": 25} + - session_id: sent by caller via x-litellm-session-id header or + metadata.session_id in request body + + Cache key pattern: + {session_iterations:}:count + + Multi-instance support: + Uses Redis Lua script for atomic increment (same pattern as + parallel_request_limiter_v3). Falls back to in-memory cache + when Redis is unavailable. + """ + + def __init__(self, internal_usage_cache: InternalUsageCache): + self.internal_usage_cache = internal_usage_cache + self.ttl = int( + os.getenv("LITELLM_MAX_ITERATIONS_TTL", DEFAULT_MAX_ITERATIONS_TTL) + ) + + # Register Lua script with Redis if available (same pattern as v3 limiter) + if self.internal_usage_cache.dual_cache.redis_cache is not None: + self.increment_script = ( + self.internal_usage_cache.dual_cache.redis_cache.async_register_script( + MAX_ITERATIONS_INCREMENT_SCRIPT + ) + ) + else: + self.increment_script = None + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ) -> Optional[Union[Exception, str, dict]]: + """ + Check session iteration count before making the API call. + + Extracts session_id from request metadata and max_iterations from + key metadata. If the session has exceeded max_iterations, raises 429. + """ + # Extract session_id from request data + session_id = self._get_session_id(data) + if session_id is None: + return None + + # Extract max_iterations from key metadata + max_iterations = self._get_max_iterations(user_api_key_dict) + if max_iterations is None: + return None + + verbose_proxy_logger.debug( + "MaxIterationsHandler: session_id=%s, max_iterations=%s", + session_id, + max_iterations, + ) + + # Increment and check + cache_key = self._make_cache_key(session_id) + current_count = await self._increment_and_get(cache_key) + + if current_count > max_iterations: + raise HTTPException( + status_code=429, + detail=( + f"Max iterations exceeded for session {session_id}. " + f"Current count: {current_count}, max_iterations: {max_iterations}." + ), + ) + + verbose_proxy_logger.debug( + "MaxIterationsHandler: session_id=%s, count=%s/%s", + session_id, + current_count, + max_iterations, + ) + + return None + + def _get_session_id(self, data: dict) -> Optional[str]: + """Extract session_id from request metadata.""" + metadata = data.get("metadata") or {} + session_id = metadata.get("session_id") + if session_id is not None: + return str(session_id) + + # Also check litellm_metadata (used for /thread and /assistant endpoints) + litellm_metadata = data.get("litellm_metadata") or {} + session_id = litellm_metadata.get("session_id") + if session_id is not None: + return str(session_id) + + return None + + def _get_max_iterations( + self, user_api_key_dict: UserAPIKeyAuth + ) -> Optional[int]: + """Extract max_iterations from key metadata.""" + metadata = user_api_key_dict.metadata or {} + max_iterations = metadata.get("max_iterations") + if max_iterations is not None: + return int(max_iterations) + return None + + def _make_cache_key(self, session_id: str) -> str: + """ + Create cache key for session iteration counter. + + Uses Redis hash-tag pattern {session_iterations:} so all + keys for a session land on the same Redis Cluster slot. + """ + return f"{{session_iterations:{session_id}}}:count" + + async def _increment_and_get(self, cache_key: str) -> int: + """ + Atomically increment the session counter and return the new value. + + Tries Redis first (via registered Lua script for atomicity across + instances), falls back to in-memory cache. + """ + if self.increment_script is not None: + try: + result = await self.increment_script( + keys=[cache_key], + args=[self.ttl], + ) + return int(result) + except Exception as e: + verbose_proxy_logger.warning( + "MaxIterationsHandler: Redis failed, falling back to in-memory: %s", + str(e), + ) + + # Fallback: in-memory cache + return await self._in_memory_increment(cache_key) + + async def _in_memory_increment(self, cache_key: str) -> int: + """Increment counter in in-memory cache with TTL.""" + current = await self.internal_usage_cache.async_get_cache( + key=cache_key, + litellm_parent_otel_span=None, + local_only=True, + ) + new_value = (int(current) if current is not None else 0) + 1 + await self.internal_usage_cache.async_set_cache( + key=cache_key, + value=new_value, + ttl=self.ttl, + litellm_parent_otel_span=None, + local_only=True, + ) + return new_value diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 815d64f22ad..0734756d8ed 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -12,7 +12,7 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import log_db_metrics +from litellm.proxy.auth.auth_checks import get_key_object, get_team_object, log_db_metrics from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( @@ -76,6 +76,10 @@ class _ProxyDBLogger(CustomLogger): traceback_str=traceback_str, ) + _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata=_metadata, + ) + existing_metadata: dict = request_data.get("metadata", None) or {} existing_metadata.update(_metadata) @@ -255,6 +259,72 @@ class _ProxyDBLogger(CustomLogger): "Error in tracking cost callback - %s", str(e) ) + @staticmethod + async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: + """ + Enriches failure spend log metadata by looking up the key object (and team object) + from cache/DB when key fields are missing. + + This handles two scenarios: + 1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other + fields are null. We look up the full key object to fill in alias, user_id, + team_id, etc. + 2. Post-auth failures (provider errors, rate limits): key fields are populated + but team_alias is missing because LiteLLM_VerificationTokenView SQL view + doesn't include it. We look up the team object to fill in team_alias. + """ + api_key_hash = metadata.get("user_api_key") + if not api_key_hash: + return metadata + + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + # Step 1: If key fields are missing, look up the full key object + if metadata.get("user_api_key_alias") is None: + try: + key_obj = await get_key_object( + hashed_token=api_key_hash, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if metadata.get("user_api_key_alias") is None: + metadata["user_api_key_alias"] = key_obj.key_alias + if metadata.get("user_api_key_user_id") is None: + metadata["user_api_key_user_id"] = key_obj.user_id + if metadata.get("user_api_key_team_id") is None: + metadata["user_api_key_team_id"] = key_obj.team_id + if metadata.get("user_api_key_org_id") is None: + metadata["user_api_key_org_id"] = key_obj.org_id + except Exception: + verbose_proxy_logger.debug( + "Failed to enrich failure metadata with key info for api_key=%s", + api_key_hash, + ) + + # Step 2: If team_id is known but team_alias is missing, look up the team object + team_id = metadata.get("user_api_key_team_id") + if team_id and metadata.get("user_api_key_team_alias") is None: + try: + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if team_obj.team_alias is not None: + metadata["user_api_key_team_alias"] = team_obj.team_alias + except Exception: + verbose_proxy_logger.debug( + "Failed to enrich failure metadata with team_alias for team_id=%s", + team_id, + ) + return metadata + @staticmethod def _should_track_errors_in_db(): """ diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 52f0b1d46e9..d2312a00c3b 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -10,10 +10,15 @@ import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors, - LitellmDataForBackendLLMCall, - LitellmUserRoles, SpecialHeaders, - TeamCallbackMetadata, UserAPIKeyAuth) +from litellm.proxy._types import ( + AddTeamCallback, + CommonProxyErrors, + LitellmDataForBackendLLMCall, + LitellmUserRoles, + SpecialHeaders, + TeamCallbackMetadata, + UserAPIKeyAuth, +) from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers # Cache special headers as a frozenset for O(1) lookup performance @@ -23,9 +28,12 @@ _SPECIAL_HEADERS_CACHE = frozenset( from litellm.router import Router from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS from litellm.types.services import ServiceTypes -from litellm.types.utils import (LlmProviders, ProviderSpecificHeader, - StandardLoggingUserAPIKeyMetadata, - SupportedCacheControls) +from litellm.types.utils import ( + LlmProviders, + ProviderSpecificHeader, + StandardLoggingUserAPIKeyMetadata, + SupportedCacheControls, +) service_logger_obj = ServiceLogging() # used for tracking latency on OTEL @@ -228,7 +236,9 @@ def _get_dynamic_logging_metadata( def clean_headers( - headers: Headers, litellm_key_header_name: Optional[str] = None + headers: Headers, + litellm_key_header_name: Optional[str] = None, + forward_llm_provider_auth_headers: bool = False, ) -> dict: """ Removes litellm api key from headers @@ -238,15 +248,18 @@ def clean_headers( clean_headers = {} litellm_key_lower = ( litellm_key_header_name.lower() if litellm_key_header_name is not None else None - ) - + ) for header, value in headers.items(): header_lower = header.lower() - # Preserve Authorization header if it contains Anthropic OAuth token (sk-ant-oat*) - # This allows OAuth tokens to be forwarded to Anthropic-compatible providers - # via add_provider_specific_headers_to_request() + if header_lower == "authorization" and is_anthropic_oauth_key(value): clean_headers[header] = value + elif forward_llm_provider_auth_headers and header_lower in _SPECIAL_HEADERS_CACHE: + if litellm_key_lower and header_lower == litellm_key_lower: + continue + if header_lower == "authorization": + continue + clean_headers[header] = value # Check if header should be excluded: either in special headers cache or matches custom litellm key elif header_lower not in _SPECIAL_HEADERS_CACHE and ( litellm_key_lower is None or header_lower != litellm_key_lower @@ -654,7 +667,8 @@ class LiteLLMProxyRequestSetup: return data from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields, - LiteLLM_ManagementEndpoint_MetadataFields_Premium) + LiteLLM_ManagementEndpoint_MetadataFields_Premium, + ) # ignore any special fields added_metadata = {} @@ -826,6 +840,11 @@ async def add_litellm_data_to_request( # noqa: PLR0915 from litellm.types.proxy.litellm_pre_call_utils import SecretFields _raw_headers: Dict[str, str] = _safe_get_request_headers(request) + + forward_llm_auth = False + if general_settings: + forward_llm_auth = general_settings.get("forward_llm_provider_auth_headers", False) + _headers: Dict[str, str] = clean_headers( request.headers, litellm_key_header_name=( @@ -833,7 +852,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if general_settings is not None else None ), + forward_llm_provider_auth_headers=forward_llm_auth, ) + verbose_proxy_logger.debug(f"Request Headers: {_headers}") + verbose_proxy_logger.debug(f"Raw Headers: {_raw_headers}") ########################################################## # Init - Proxy Server Request @@ -1479,8 +1501,7 @@ async def move_guardrails_to_metadata( # Only check policy engine if no local config (avoid import + registry lookup) if not (has_key_config or has_team_config or has_request_config): - from litellm.proxy.policy_engine.policy_registry import \ - get_policy_registry + from litellm.proxy.policy_engine.policy_registry import get_policy_registry if not get_policy_registry().is_initialized(): # Nothing configured anywhere - clean up request body fields and return @@ -1544,16 +1565,14 @@ async def move_guardrails_to_metadata( def _is_policy_version_id(s: str) -> bool: """Return True if string is a policy version ID (starts with policy_ prefix).""" - from litellm.proxy.policy_engine.policy_registry import \ - POLICY_VERSION_ID_PREFIX + from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX) def _extract_policy_id(s: str) -> Optional[str]: """Extract raw UUID from policy_ string, or None if not a valid version ID.""" - from litellm.proxy.policy_engine.policy_registry import \ - POLICY_VERSION_ID_PREFIX + from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX if not _is_policy_version_id(s): return None @@ -1574,9 +1593,10 @@ def _match_and_track_policies( """ from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.callback_utils import ( - add_policy_sources_to_metadata, add_policy_to_applied_policies_header) - from litellm.proxy.policy_engine.attachment_registry import \ - get_attachment_registry + add_policy_sources_to_metadata, + add_policy_to_applied_policies_header, + ) + from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher # Get matching policies via attachments (with match reasons for attribution) @@ -1721,8 +1741,7 @@ async def add_guardrails_from_policy_engine( user_api_key_dict: The user's API key authentication info """ from litellm._logging import verbose_proxy_logger - from litellm.proxy.common_utils.http_parsing_utils import \ - get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import PolicyMatchContext diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index a230c2e9336..c1165ab26d0 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2535,6 +2535,7 @@ async def generate_key_helper_fn( # noqa: PLR0915 user_id: Optional[str] = None, user_alias: Optional[str] = None, team_id: Optional[str] = None, + agent_id: Optional[str] = None, user_email: Optional[str] = None, user_role: Optional[str] = None, max_parallel_requests: Optional[int] = None, @@ -2668,6 +2669,7 @@ async def generate_key_helper_fn( # noqa: PLR0915 "max_budget": key_max_budget, "user_id": user_id, "team_id": team_id, + "agent_id": agent_id, "project_id": project_id, "max_parallel_requests": max_parallel_requests, "metadata": metadata_json, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 607306f3806..82cfd455be6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1216,9 +1216,7 @@ try: # Case 2: Runtime UI exists and is ready if has_content and is_pre_restructured: - verbose_proxy_logger.info( - f"Using pre-restructured UI at {runtime_ui_path}" - ) + verbose_proxy_logger.info(f"Using pre-restructured UI at {runtime_ui_path}") ui_path = runtime_ui_path # Case 3: Runtime UI exists but needs restructuring @@ -2994,6 +2992,10 @@ class ProxyConfig: if master_key is not None and isinstance(master_key, str): litellm_master_key_hash = hash_token(master_key) + else: + verbose_proxy_logger.critical( + "LITELLM_MASTER_KEY is not set! All requests will be treated as INTERNAL_USER with no admin access. Set LITELLM_MASTER_KEY for production use." + ) ### USER API KEY CACHE IN-MEMORY TTL ### user_api_key_cache_ttl = general_settings.get( "user_api_key_cache_ttl", None @@ -3796,6 +3798,7 @@ class ProxyConfig: parsed = value elif isinstance(value, str): import json + try: parsed = yaml.safe_load(value) except (yaml.YAMLError, json.JSONDecodeError): @@ -4381,10 +4384,12 @@ class ProxyConfig: if self._should_load_db_object(object_type="model_cost_map"): await self._check_and_reload_model_cost_map(prisma_client=prisma_client) - + if self._should_load_db_object(object_type="anthropic_beta_headers"): - await self._check_and_reload_anthropic_beta_headers(prisma_client=prisma_client) - + await self._check_and_reload_anthropic_beta_headers( + prisma_client=prisma_client + ) + if self._should_load_db_object(object_type="sso_settings"): await self._init_sso_settings_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type="cache_settings"): @@ -4614,7 +4619,9 @@ class ProxyConfig: f"Error in _check_and_reload_model_cost_map: {str(e)}" ) - async def _check_and_reload_anthropic_beta_headers(self, prisma_client: PrismaClient): + async def _check_and_reload_anthropic_beta_headers( + self, prisma_client: PrismaClient + ): """ Check if anthropic beta headers config needs to be reloaded based on database configuration. This function runs every 10 seconds as part of _init_non_llm_objects_in_db. @@ -4705,7 +4712,11 @@ class ProxyConfig: ) # Count providers in config - provider_count = sum(1 for k in new_config.keys() if k != "provider_aliases" and k != "description") + provider_count = sum( + 1 + for k in new_config.keys() + if k != "provider_aliases" and k != "description" + ) verbose_proxy_logger.info( f"Anthropic beta headers config reloaded successfully. Providers: {provider_count}" ) @@ -5687,8 +5698,7 @@ class ProxyStartupEvent: ): _db_val = _db_gs_record.param_value.get("store_model_in_db") if _db_val is True or ( - isinstance(_db_val, str) - and _db_val.lower() == "true" + isinstance(_db_val, str) and _db_val.lower() == "true" ): store_model_in_db = True verbose_proxy_logger.info( @@ -6155,6 +6165,7 @@ class ProxyStartupEvent: "Pyroscope profiling will not run. Install with: pip install pyroscope-io" ) + #### API ENDPOINTS #### @router.get( "/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"] @@ -10993,18 +11004,14 @@ async def get_favicon(): from fastapi.responses import Response current_dir = os.path.dirname(os.path.abspath(__file__)) - default_favicon = os.path.join( - current_dir, "_experimental", "out", "favicon.ico" - ) + default_favicon = os.path.join(current_dir, "_experimental", "out", "favicon.ico") favicon_url = os.getenv("LITELLM_FAVICON_URL", "") if not favicon_url: if os.path.exists(default_favicon): return FileResponse(default_favicon, media_type="image/x-icon") - raise HTTPException( - status_code=404, detail="Default favicon not found" - ) + raise HTTPException(status_code=404, detail="Default favicon not found") if favicon_url.startswith(("http://", "https://")): try: @@ -11019,9 +11026,7 @@ async def get_favicon(): ) response = await async_client.get(favicon_url) if response.status_code == 200: - content_type = response.headers.get( - "content-type", "image/x-icon" - ) + content_type = response.headers.get("content-type", "image/x-icon") return Response( content=response.content, media_type=content_type, @@ -11033,12 +11038,8 @@ async def get_favicon(): response.status_code, ) if os.path.exists(default_favicon): - return FileResponse( - default_favicon, media_type="image/x-icon" - ) - raise HTTPException( - status_code=404, detail="Favicon not found" - ) + return FileResponse(default_favicon, media_type="image/x-icon") + raise HTTPException(status_code=404, detail="Favicon not found") except HTTPException: raise except Exception as e: @@ -11046,20 +11047,14 @@ async def get_favicon(): "Error downloading favicon from %s: %s", favicon_url, e ) if os.path.exists(default_favicon): - return FileResponse( - default_favicon, media_type="image/x-icon" - ) - raise HTTPException( - status_code=404, detail="Favicon not found" - ) + return FileResponse(default_favicon, media_type="image/x-icon") + raise HTTPException(status_code=404, detail="Favicon not found") else: if os.path.exists(favicon_url): return FileResponse(favicon_url, media_type="image/x-icon") if os.path.exists(default_favicon): return FileResponse(default_favicon, media_type="image/x-icon") - raise HTTPException( - status_code=404, detail="Favicon not found" - ) + raise HTTPException(status_code=404, detail="Favicon not found") #### INVITATION MANAGEMENT #### @@ -12545,7 +12540,9 @@ async def reload_anthropic_beta_headers( }, ) - provider_count = sum(1 for k in new_config.keys() if k not in ["provider_aliases", "description"]) + provider_count = sum( + 1 for k in new_config.keys() if k not in ["provider_aliases", "description"] + ) verbose_proxy_logger.info( f"Anthropic beta headers config reloaded successfully in current pod. Providers: {provider_count}" ) @@ -12557,7 +12554,9 @@ async def reload_anthropic_beta_headers( "timestamp": current_time.isoformat(), } except Exception as e: - verbose_proxy_logger.exception(f"Failed to reload anthropic beta headers: {str(e)}") + verbose_proxy_logger.exception( + f"Failed to reload anthropic beta headers: {str(e)}" + ) raise HTTPException( status_code=500, detail=f"Failed to reload anthropic beta headers: {str(e)}" ) @@ -12679,7 +12678,8 @@ async def cancel_anthropic_beta_headers_reload( f"Failed to cancel anthropic beta headers reload: {str(e)}" ) raise HTTPException( - status_code=500, detail=f"Failed to cancel anthropic beta headers reload: {str(e)}" + status_code=500, + detail=f"Failed to cancel anthropic beta headers reload: {str(e)}", ) @@ -12726,7 +12726,9 @@ async def get_anthropic_beta_headers_reload_status( ) if config_record is None or config_record.param_value is None: - verbose_proxy_logger.info("No anthropic beta headers reload configuration found") + verbose_proxy_logger.info( + "No anthropic beta headers reload configuration found" + ) return { "scheduled": False, "interval_hours": None, @@ -12752,7 +12754,9 @@ async def get_anthropic_beta_headers_reload_status( # Use pod's in-memory last reload time if last_anthropic_beta_headers_reload is not None: try: - last_reload_time = datetime.fromisoformat(last_anthropic_beta_headers_reload) + last_reload_time = datetime.fromisoformat( + last_anthropic_beta_headers_reload + ) time_since_last_reload = current_time - last_reload_time hours_since_last_reload = time_since_last_reload.total_seconds() / 3600 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 23917cf7c7f..440c9c1d829 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -64,6 +64,8 @@ model LiteLLM_AgentsTable { litellm_params Json? agent_card_params Json agent_access_groups String[] @default([]) + object_permission_id String? + object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) created_at DateTime @default(now()) @map("created_at") created_by String updated_at DateTime @default(now()) @updatedAt @map("updated_at") @@ -264,6 +266,7 @@ model LiteLLM_ObjectPermissionTable { organizations LiteLLM_OrganizationTable[] users LiteLLM_UserTable[] end_users LiteLLM_EndUserTable[] + agents_table LiteLLM_AgentsTable[] } // Holds the MCP server configuration @@ -314,6 +317,7 @@ model LiteLLM_VerificationToken { router_settings Json? @default("{}") user_id String? team_id String? + agent_id String? project_id String? permissions Json @default("{}") max_parallel_requests Int? @@ -476,6 +480,7 @@ model LiteLLM_SpendLogs { completion_tokens Int @default(0) startTime DateTime // Assuming start_time is a DateTime field endTime DateTime // Assuming end_time is a DateTime field + request_duration_ms Int? completionStartTime DateTime? // Assuming completionStartTime is a DateTime field model String @default("") model_id String? @default("") // the model id stored in proxy model db diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index b4aeda62004..5b58fbe70a0 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1726,7 +1726,7 @@ async def ui_view_spend_logs( # noqa: PLR0915 ) # Validate sort_by and sort_order - valid_sort_fields = {"spend", "total_tokens", "startTime", "endTime"} + valid_sort_fields = {"spend", "total_tokens", "startTime", "endTime", "request_duration_ms"} if sort_by not in valid_sort_fields: raise ProxyException( message=f"Invalid sort_by: {sort_by}. Must be one of: {', '.join(sorted(valid_sort_fields))}", @@ -1939,7 +1939,8 @@ async def ui_view_spend_logs( # noqa: PLR0915 custom_llm_provider, api_base, "user", metadata, cache_hit, cache_key, request_tags, team_id, organization_id, end_user, requester_ip_address, - session_id, status, mcp_namespaced_tool_name, agent_id + session_id, status, mcp_namespaced_tool_name, agent_id, + COALESCE(request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms FROM "LiteLLM_SpendLogs" WHERE {" AND ".join(sql_conditions)} ORDER BY {_sql_col} {_sql_dir} diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index d517c76a08c..05dbb8ed71e 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -447,6 +447,7 @@ def get_logging_payload( # noqa: PLR0915 kwargs=kwargs, standard_logging_payload=standard_logging_payload, ), + request_duration_ms=_get_request_duration_ms(start_time, end_time), status=_get_status_for_spend_log( metadata=metadata, ), @@ -496,6 +497,16 @@ def _get_session_id_for_spend_log( return str(uuid.uuid4()) +def _get_request_duration_ms( + start_time: datetime, end_time: datetime +) -> Optional[int]: + """Compute request duration in milliseconds from start and end times.""" + try: + return int((end_time - start_time).total_seconds() * 1000) + except Exception: + return None + + def _ensure_datetime_utc(timestamp: datetime) -> datetime: """Helper to ensure datetime is in UTC""" timestamp = timestamp.astimezone(timezone.utc) diff --git a/litellm/types/agents.py b/litellm/types/agents.py index f4e410a3e2d..3ad898b1935 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -167,16 +167,25 @@ class AugmentedAgentCard(AgentCard): is_public: bool +# Object permission shape for agent MCP tool access (mirrors LiteLLM_ObjectPermissionBase) +class AgentObjectPermission(TypedDict, total=False): + mcp_servers: Optional[List[str]] + mcp_access_groups: Optional[List[str]] + mcp_tool_permissions: Optional[Dict[str, List[str]]] + + class AgentConfig(TypedDict, total=False): agent_name: Required[str] agent_card_params: Required[AgentCard] litellm_params: Dict[str, Any] # allow for any future litellm params + object_permission: AgentObjectPermission class PatchAgentRequest(TypedDict, total=False): agent_name: str agent_card_params: AgentCard litellm_params: Dict[str, Any] + object_permission: AgentObjectPermission # Request/Response models for CRUD endpoints @@ -187,6 +196,7 @@ class AgentResponse(BaseModel): agent_name: str litellm_params: Optional[Dict[str, Any]] = None agent_card_params: Dict[str, Any] + object_permission: Optional[Dict[str, Any]] = None created_at: Optional[datetime] = None updated_at: Optional[datetime] = None created_by: Optional[str] = None diff --git a/litellm/types/caching.py b/litellm/types/caching.py index ad2ffeaf9bd..7126ba3e9b9 100644 --- a/litellm/types/caching.py +++ b/litellm/types/caching.py @@ -52,6 +52,24 @@ class RedisPipelineSetOperation(TypedDict): ttl: Optional[int] +class RedisPipelineRpushOperation(TypedDict): + """ + TypedDict for 1 Redis Pipeline RPUSH Operation + """ + + key: str + values: List[Any] + + +class RedisPipelineLpopOperation(TypedDict): + """ + TypedDict for 1 Redis Pipeline LPOP Operation + """ + + key: str + count: Optional[int] + + DynamicCacheControl = TypedDict( "DynamicCacheControl", { diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 4ab81f8fd57..15e8d1be930 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1125,6 +1125,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False): prompt: Optional[PromptObject] max_tool_calls: Optional[int] prompt_cache_key: Optional[str] + prompt_cache_retention: Optional[str] stream_options: Optional[dict] top_logprobs: Optional[int] partial_images: Optional[ diff --git a/litellm/types/responses/main.py b/litellm/types/responses/main.py index bda53bae082..449a5ac49c1 100644 --- a/litellm/types/responses/main.py +++ b/litellm/types/responses/main.py @@ -6,7 +6,7 @@ from typing_extensions import Any, List, Optional, TypedDict from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject -Phase = Optional[Literal["commentary", "final_answer"]] # TODO: Once openai sdk has updated, we can remove this and use the openai sdk type +Phase = Optional[Literal["commentary", "final_answer"]] class GenericResponseOutputItemContentAnnotation(BaseLiteLLMOpenAIResponseObject): """Annotation for content in a message""" diff --git a/litellm/types/utils.py b/litellm/types/utils.py index e8d6ac79708..dda32d98383 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1,7 +1,17 @@ import json import time from enum import Enum -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional, Union +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Mapping, + Optional, + Union, + get_args, +) from openai._models import BaseModel as OpenAIObject from openai.types.audio.transcription_create_params import ( @@ -3186,6 +3196,12 @@ OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: set[str] = { LlmProviders.HOSTED_VLLM.value, } +ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "vertex_ai"] + +LIST_BATCHES_SUPPORTED_PROVIDERS: frozenset[str] = frozenset( + get_args(ListBatchesSupportedProvider) +) + class SearchProviders(str, Enum): """ diff --git a/litellm/utils.py b/litellm/utils.py index 5046929257d..7ac828aefc1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4644,11 +4644,12 @@ def add_provider_specific_params_to_optional_params( ) is False ): - extra_body = passed_params.pop("extra_body", {}) + extra_body = passed_params.pop("extra_body", None) or {} for k in passed_params.keys(): if k not in openai_params and passed_params[k] is not None: extra_body[k] = passed_params[k] - optional_params.setdefault("extra_body", {}) + if not isinstance(optional_params.get("extra_body"), dict): + optional_params["extra_body"] = {} initial_extra_body = { **optional_params["extra_body"], **extra_body, diff --git a/poetry.lock b/poetry.lock index 0d2895248ba..34227a69ccb 100644 --- a/poetry.lock +++ b/poetry.lock @@ -3222,15 +3222,15 @@ files = [ [[package]] name = "litellm-proxy-extras" -version = "0.4.47" +version = "0.4.48" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." optional = true python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8" groups = ["main"] markers = "extra == \"proxy\"" files = [ - {file = "litellm_proxy_extras-0.4.47-py3-none-any.whl", hash = "sha256:2e900ae3edfbc20d27556f092d914974d37bac213efe88b8fd5287f77b2b7ca7"}, - {file = "litellm_proxy_extras-0.4.47.tar.gz", hash = "sha256:42d88929f9eaf0b827046d3712095354db843c1716ccabb2a40c806ea5f809b9"}, + {file = "litellm_proxy_extras-0.4.48-py3-none-any.whl", hash = "sha256:097001fccec5dbf4cffd902114898a9cfeba62673202447d55d2d0286cf93126"}, + {file = "litellm_proxy_extras-0.4.48.tar.gz", hash = "sha256:5d5d8acf31b92d0cd6738555fb4a2411819755155438de9fb23c724c356400a2"}, ] [[package]] @@ -7989,4 +7989,4 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.9,<4.0" -content-hash = "6f39f8c731e625f37460e1f5f9ba3cce63956540dc04baf6ce1b7c18b88b8322" +content-hash = "b9b1e47b3b84748c0053be6a544c2399bf2601746a4f88dcb1be7c5e4eeab359" diff --git a/pyproject.toml b/pyproject.toml index 18844da46c8..c848bc3e789 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -61,7 +61,7 @@ boto3 = { version = "1.40.76", optional = true } redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"} mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"} a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"} -litellm-proxy-extras = {version = "0.4.47", optional = true} +litellm-proxy-extras = {version = "0.4.48", optional = true} rich = {version = "13.7.1", optional = true} litellm-enterprise = {version = "0.1.32", optional = true} diskcache = {version = "^5.6.1", optional = true} diff --git a/requirements.txt b/requirements.txt index dbde6ababc9..6cdb2f63a33 100644 --- a/requirements.txt +++ b/requirements.txt @@ -10,7 +10,7 @@ cryptography==46.0.5 #GHSA-r6ph-v2qm-q3c2 anyio==4.8.0 # openai + http req. httpx==0.28.1 -openai==2.9.0 # openai req. +openai==2.24.0 # openai req. fastapi==0.120.1 # server dep starlette==0.49.1 # starlette fastapi dep backoff==2.2.1 # server dep @@ -57,7 +57,7 @@ grpcio>=1.75.0; python_version >= "3.14" sentry_sdk==2.21.0 # for sentry error handling detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests tzdata==2025.1 # IANA time zone database -litellm-proxy-extras==0.4.47 # for proxy extras - e.g. prisma migrations +litellm-proxy-extras==0.4.48 # for proxy extras - e.g. prisma migrations llm-sandbox==0.3.31 # for skill execution in sandbox ### LITELLM PACKAGE DEPENDENCIES python-dotenv==1.0.1 # for env diff --git a/schema.prisma b/schema.prisma index 4af7484148c..440c9c1d829 100644 --- a/schema.prisma +++ b/schema.prisma @@ -64,6 +64,8 @@ model LiteLLM_AgentsTable { litellm_params Json? agent_card_params Json agent_access_groups String[] @default([]) + object_permission_id String? + object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) created_at DateTime @default(now()) @map("created_at") created_by String updated_at DateTime @default(now()) @updatedAt @map("updated_at") @@ -264,6 +266,7 @@ model LiteLLM_ObjectPermissionTable { organizations LiteLLM_OrganizationTable[] users LiteLLM_UserTable[] end_users LiteLLM_EndUserTable[] + agents_table LiteLLM_AgentsTable[] } // Holds the MCP server configuration @@ -273,7 +276,6 @@ model LiteLLM_MCPServerTable { alias String? description String? url String? - spec_path String? transport String @default("sse") auth_type String? credentials Json? @default("{}") @@ -315,6 +317,7 @@ model LiteLLM_VerificationToken { router_settings Json? @default("{}") user_id String? team_id String? + agent_id String? project_id String? permissions Json @default("{}") max_parallel_requests Int? @@ -477,6 +480,7 @@ model LiteLLM_SpendLogs { completion_tokens Int @default(0) startTime DateTime // Assuming start_time is a DateTime field endTime DateTime // Assuming end_time is a DateTime field + request_duration_ms Int? completionStartTime DateTime? // Assuming completionStartTime is a DateTime field model String @default("") model_id String? @default("") // the model id stored in proxy model db @@ -1052,6 +1056,26 @@ model LiteLLM_PolicyAttachmentTable { updated_by String? } +// Global tool registry - auto-discovered from LLM responses; admins set call_policy here +model LiteLLM_ToolTable { + tool_id String @id @default(uuid()) + tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space" + origin String? // MCP server name or "user_defined" + call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked" + call_count Int @default(0) // cumulative number of times this tool was seen + assignments Json? @default("{}") + key_hash String? // hash of the virtual key that first called this tool + team_id String? // team that first called this tool + key_alias String? // human-readable alias of the virtual key + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@index([call_policy]) + @@index([team_id]) +} + //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/scripts/test_agent_mcp_endpoints.sh b/scripts/test_agent_mcp_endpoints.sh new file mode 100755 index 00000000000..93cc68db2ed --- /dev/null +++ b/scripts/test_agent_mcp_endpoints.sh @@ -0,0 +1,186 @@ +#!/usr/bin/env bash +# +# Test agent endpoint-level changes for MCP tool permissions (object_permission). +# Requires: proxy running, valid admin API key, curl, jq. +# +# Usage: +# export LITELLM_PROXY_BASE_URL="http://localhost:4000" # optional, default below +# export LITELLM_API_KEY="sk-..." # required +# ./scripts/test_agent_mcp_endpoints.sh +# +set -euo pipefail + +BASE_URL="${LITELLM_PROXY_BASE_URL:-http://localhost:4000}" +API_KEY="${LITELLM_API_KEY:-}" + +if ! command -v jq &>/dev/null; then + echo "Error: jq is required. Install with: brew install jq (macOS) or apt install jq (Linux)" + exit 1 +fi +if [[ -z "$API_KEY" ]]; then + echo "Error: LITELLM_API_KEY is not set. Export it or pass via env." + exit 1 +fi + +AUTH_HEADER="Authorization: Bearer $API_KEY" +AGENT_NAME="test-agent-mcp-$(date +%s)" + +# Colors for output +RED='\033[0;31m' +GREEN='\033[0;32m' +YELLOW='\033[1;33m' +NC='\033[0m' + +pass() { echo -e "${GREEN}PASS${NC}: $*"; } +fail() { echo -e "${RED}FAIL${NC}: $*"; exit 1; } +info() { echo -e "${YELLOW}INFO${NC}: $*"; } + +# --- 1. Create agent with object_permission --- +info "Creating agent with object_permission (mcp_servers, mcp_tool_permissions)..." +CREATE_RESP=$(curl -s -w "\n%{http_code}" -X POST "$BASE_URL/v1/agents" \ + -H "$AUTH_HEADER" \ + -H "Content-Type: application/json" \ + -d '{ + "agent_name": "'"$AGENT_NAME"'", + "agent_card_params": { + "protocolVersion": "1.0", + "name": "Test MCP Agent", + "description": "Agent for endpoint tests", + "url": "http://localhost:9999/", + "version": "1.0.0", + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "capabilities": {"streaming": true}, + "skills": [] + }, + "object_permission": { + "mcp_servers": ["server_1", "server_2"], + "mcp_access_groups": ["group_a"], + "mcp_tool_permissions": {"server_1": ["tool_a", "tool_b"], "server_2": ["tool_c"]} + } + }') +HTTP_CODE=$(echo "$CREATE_RESP" | tail -n1) +BODY=$(echo "$CREATE_RESP" | sed '$d') +if [[ "$HTTP_CODE" != "200" ]]; then + fail "POST /v1/agents returned $HTTP_CODE. Body: $BODY" +fi +AGENT_ID=$(echo "$BODY" | jq -r '.agent_id') +if [[ -z "$AGENT_ID" || "$AGENT_ID" == "null" ]]; then + fail "POST /v1/agents did not return agent_id. Body: $BODY" +fi +pass "Created agent $AGENT_ID" + +# Check create response includes object_permission +OP=$(echo "$BODY" | jq '.object_permission') +if [[ "$OP" == "null" || -z "$OP" ]]; then + fail "POST /v1/agents response missing object_permission. Body: $BODY" +fi +SERVERS=$(echo "$OP" | jq -r '.mcp_servers | join(",")') +if [[ "$SERVERS" != "server_1,server_2" ]]; then + fail "object_permission.mcp_servers unexpected: $SERVERS" +fi +pass "Create response includes object_permission with mcp_servers and mcp_tool_permissions" + +# --- 2. GET /v1/agents (list) includes object_permission for our agent --- +info "GET /v1/agents and check one agent has object_permission..." +LIST_RESP=$(curl -s -w "\n%{http_code}" -X GET "$BASE_URL/v1/agents" -H "$AUTH_HEADER") +LIST_CODE=$(echo "$LIST_RESP" | tail -n1) +LIST_BODY=$(echo "$LIST_RESP" | sed '$d') +if [[ "$LIST_CODE" != "200" ]]; then + fail "GET /v1/agents returned $LIST_CODE" +fi +AGENT_IN_LIST=$(echo "$LIST_BODY" | jq --arg id "$AGENT_ID" '.[] | select(.agent_id == $id)') +if [[ -z "$AGENT_IN_LIST" ]]; then + fail "GET /v1/agents did not return agent $AGENT_ID (list might be key-scoped)" +fi +OP_LIST=$(echo "$AGENT_IN_LIST" | jq '.object_permission') +if [[ "$OP_LIST" == "null" || -z "$OP_LIST" ]]; then + fail "GET /v1/agents list entry for agent missing object_permission" +fi +pass "GET /v1/agents list includes object_permission for agent" + +# --- 3. GET /v1/agents/{agent_id} returns object_permission --- +info "GET /v1/agents/{agent_id}..." +GET_RESP=$(curl -s -w "\n%{http_code}" -X GET "$BASE_URL/v1/agents/$AGENT_ID" -H "$AUTH_HEADER") +GET_CODE=$(echo "$GET_RESP" | tail -n1) +GET_BODY=$(echo "$GET_RESP" | sed '$d') +if [[ "$GET_CODE" != "200" ]]; then + fail "GET /v1/agents/$AGENT_ID returned $GET_CODE. Body: $GET_BODY" +fi +OP_GET=$(echo "$GET_BODY" | jq '.object_permission') +if [[ "$OP_GET" == "null" || -z "$OP_GET" ]]; then + fail "GET /v1/agents/$AGENT_ID response missing object_permission" +fi +TOOL_PERMS=$(echo "$OP_GET" | jq -r '.mcp_tool_permissions.server_1 | join(",")') +if [[ "$TOOL_PERMS" != "tool_a,tool_b" ]]; then + fail "object_permission.mcp_tool_permissions.server_1 unexpected: $TOOL_PERMS" +fi +pass "GET /v1/agents/{agent_id} returns object_permission with mcp_tool_permissions" + +# --- 4. PATCH /v1/agents/{agent_id} with new object_permission --- +info "PATCH /v1/agents/{agent_id} with updated object_permission..." +PATCH_RESP=$(curl -s -w "\n%{http_code}" -X PATCH "$BASE_URL/v1/agents/$AGENT_ID" \ + -H "$AUTH_HEADER" \ + -H "Content-Type: application/json" \ + -d '{ + "object_permission": { + "mcp_servers": ["server_3"], + "mcp_tool_permissions": {"server_3": ["tool_x"]} + } + }') +PATCH_CODE=$(echo "$PATCH_RESP" | tail -n1) +PATCH_BODY=$(echo "$PATCH_RESP" | sed '$d') +if [[ "$PATCH_CODE" != "200" ]]; then + fail "PATCH /v1/agents/$AGENT_ID returned $PATCH_CODE. Body: $PATCH_BODY" +fi +OP_PATCH=$(echo "$PATCH_BODY" | jq '.object_permission') +if [[ "$OP_PATCH" == "null" || -z "$OP_PATCH" ]]; then + fail "PATCH response missing object_permission" +fi +PATCH_SERVERS=$(echo "$OP_PATCH" | jq -r '.mcp_servers | join(",")') +if [[ "$PATCH_SERVERS" != "server_3" ]]; then + fail "PATCH object_permission.mcp_servers unexpected: $PATCH_SERVERS" +fi +pass "PATCH /v1/agents/{agent_id} updates and returns object_permission" + +# --- 5. Create agent without object_permission; GET should still work --- +info "Creating agent without object_permission..." +AGENT_NAME_2="test-agent-no-mcp-$(date +%s)" +CREATE2_RESP=$(curl -s -w "\n%{http_code}" -X POST "$BASE_URL/v1/agents" \ + -H "$AUTH_HEADER" \ + -H "Content-Type: application/json" \ + -d '{ + "agent_name": "'"$AGENT_NAME_2"'", + "agent_card_params": { + "protocolVersion": "1.0", + "name": "No MCP Agent", + "description": "No object_permission", + "url": "http://localhost:9999/", + "version": "1.0.0", + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "capabilities": {}, + "skills": [] + } + }') +CODE2=$(echo "$CREATE2_RESP" | tail -n1) +BODY2=$(echo "$CREATE2_RESP" | sed '$d') +if [[ "$CODE2" != "200" ]]; then + fail "POST /v1/agents (no object_permission) returned $CODE2. Body: $BODY2" +fi +AGENT_ID_2=$(echo "$BODY2" | jq -r '.agent_id') +# object_permission may be null or absent +pass "Created agent without object_permission: $AGENT_ID_2" + +# --- 6. Cleanup: delete both agents --- +info "Deleting test agents..." +for AID in "$AGENT_ID" "$AGENT_ID_2"; do + DEL_CODE=$(curl -s -o /dev/null -w "%{http_code}" -X DELETE "$BASE_URL/v1/agents/$AID" -H "$AUTH_HEADER") + if [[ "$DEL_CODE" != "200" ]]; then + info "DELETE /v1/agents/$AID returned $DEL_CODE (non-fatal)" + fi +done +pass "Cleanup done" + +echo "" +echo -e "${GREEN}All endpoint checks passed.${NC}" diff --git a/tests/litellm/proxy/guardrails/test_custom_code_security.py b/tests/litellm/proxy/guardrails/test_custom_code_security.py new file mode 100644 index 00000000000..d855a4dde20 --- /dev/null +++ b/tests/litellm/proxy/guardrails/test_custom_code_security.py @@ -0,0 +1,94 @@ +import pytest +from litellm.proxy.guardrails.guardrail_hooks.custom_code.code_validator import ( + validate_custom_code, + CustomCodeValidationError, +) +from litellm.proxy.guardrails.guardrail_hooks.custom_code.custom_code_guardrail import ( + CustomCodeGuardrail, +) + +# Phase 4.1: Test forbidden pattern validation + + +def test_validate_custom_code_import_os(): + code = "import os\ndef apply_guardrail(inputs, req, ty):\n return allow()" + with pytest.raises(CustomCodeValidationError, match="import statements are not"): + validate_custom_code(code) + + +def test_validate_custom_code_from_subprocess(): + code = ( + "from subprocess import call\ndef apply_guardrail(i, r, t):\n return allow()" + ) + with pytest.raises( + CustomCodeValidationError, match="import statements are not allowed" + ): + validate_custom_code(code) + + +def test_validate_custom_code_exec(): + code = "def apply_guardrail(i, r, t):\n exec('print(1)')\n return allow()" + with pytest.raises(CustomCodeValidationError, match=r"exec\(\) is not allowed"): + validate_custom_code(code) + + +def test_validate_custom_code_builtins(): + code = "def apply_guardrail(i, r, t):\n print(__builtins__)\n return allow()" + with pytest.raises( + CustomCodeValidationError, match="__builtins__ access is not allowed" + ): + validate_custom_code(code) + + +def test_validate_custom_code_subclasses(): + code = "def apply_guardrail(i, r, t):\n print(''.__class__.__mro__[1].__subclasses__())\n return allow()" + with pytest.raises( + CustomCodeValidationError, match="__subclasses__ access is not allowed" + ): + validate_custom_code(code) + + +def test_validate_custom_code_clean(): + code = ( + "def apply_guardrail(inputs, request_data, input_type):\n return allow()\n" + ) + # Should not raise any exception + validate_custom_code(code) + + +# Phase 4.2: Test __builtins__ restriction in execution + + +def test_custom_code_compile_valid(): + code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()" + guardrail = CustomCodeGuardrail(custom_code=code, guardrail_name="test") + # if it doesn't fail, we successfully compiled + assert guardrail._compiled_function is not None + + +def test_custom_code_override_builtins(): + # Verify that even if pattern validation is bypassed, __builtins__ = {} blocks dangerous builtins. + # We test this by compiling safe code and verifying builtins are not accessible in the sandbox. + code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()" + guardrail = CustomCodeGuardrail(custom_code=code, guardrail_name="test") + # The compiled function's globals should have empty __builtins__ + fn_globals = guardrail._compiled_function.__globals__ + assert fn_globals.get("__builtins__") == {} + + +@pytest.mark.asyncio +async def test_custom_code_guardrail_apply(): + code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()" + guardrail = CustomCodeGuardrail(custom_code=code, guardrail_name="test") + from litellm.types.utils import GenericGuardrailAPIInputs + + result = await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["test"]), + request_data={}, + input_type="request", + ) + assert result["texts"][0] == "test" + + +# The RBAC endpoint tests are harder to write right here, but the core security +# validations are fully covered by the simple tests above. diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py index 7ea4574bab9..b459c3cfc99 100644 --- a/tests/litellm_utils_tests/test_health_check.py +++ b/tests/litellm_utils_tests/test_health_check.py @@ -473,6 +473,50 @@ def test_update_litellm_params_for_health_check(): ) +@pytest.mark.asyncio +async def test_perform_health_check_filters_by_model_id(): + """ + When model_id is passed, only that deployment is checked (not all deployments + that share the same model name). + """ + from litellm.proxy.health_check import perform_health_check + + # Two deployments with same model_name but different ids + model_list = [ + { + "model_name": "gpt-4", + "model_info": {"id": "deployment-id-1"}, + "litellm_params": {"model": "gpt-4", "api_key": "fake-key-1"}, + }, + { + "model_name": "gpt-4", + "model_info": {"id": "deployment-id-2"}, + "litellm_params": {"model": "gpt-4", "api_key": "fake-key-2"}, + }, + ] + + captured_list = [] + + async def mock_perform_health_check(m_list, details=True): + captured_list.append(m_list) + return [{"model": "gpt-4", "api_key": m_list[0]["litellm_params"]["api_key"]}], [] + + with patch( + "litellm.proxy.health_check._perform_health_check", + side_effect=mock_perform_health_check, + ): + healthy_endpoints, unhealthy_endpoints = await perform_health_check( + model_list=model_list, model_id="deployment-id-2", details=True + ) + + # Only one deployment (deployment-id-2) should have been passed to _perform_health_check + assert len(captured_list) == 1 + assert len(captured_list[0]) == 1 + assert (captured_list[0][0].get("model_info") or {}).get("id") == "deployment-id-2" + assert len(healthy_endpoints) == 1 + assert healthy_endpoints[0]["api_key"] == "fake-key-2" + + @pytest.mark.asyncio async def test_perform_health_check_with_health_check_model(): """ diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 0cbba7b5cc3..d39edef793c 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -334,6 +334,81 @@ def test_chat_completion_forward_headers( pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}") +@pytest.mark.parametrize("forward_llm_auth_headers", [True, False]) +@mock_patch_acompletion() +def test_chat_completion_forward_llm_provider_auth_headers( + mock_acompletion, client_no_auth, forward_llm_auth_headers +): + """ + Test that LLM provider auth headers (x-api-key, x-goog-api-key) are forwarded + when forward_llm_provider_auth_headers=True. + + This allows clients to send their own LLM provider API keys through the proxy. + """ + try: + # Configure general settings + gs = getattr(litellm.proxy.proxy_server, "general_settings") + gs["forward_client_headers_to_llm_api"] = True + gs["forward_llm_provider_auth_headers"] = forward_llm_auth_headers + setattr(litellm.proxy.proxy_server, "general_settings", gs) + + # Test data + test_data = { + "model": "gpt-3.5-turbo", + "messages": [ + {"role": "user", "content": "hello"}, + ], + "max_tokens": 10, + } + + # Headers including LLM provider auth + request_headers = { + "Authorization": "Bearer sk-proxy-auth-123", # Proxy auth (should be stripped) + "x-api-key": "sk-ant-api03-test-anthropic-key", # Anthropic API key + "x-goog-api-key": "google-api-key-123", # Google API key + "X-Custom-Header": "custom-value", # Custom header (should be forwarded) + } + + # Make request + response = client_no_auth.post( + "/v1/chat/completions", json=test_data, headers=request_headers + ) + + assert response.status_code == 200 + + # Check forwarded headers + forwarded_headers = mock_acompletion.call_args.kwargs.get("headers", {}) + + if forward_llm_auth_headers: + # LLM provider auth headers should be forwarded + assert "x-api-key" in forwarded_headers + assert forwarded_headers["x-api-key"] == "sk-ant-api03-test-anthropic-key" + assert "x-goog-api-key" in forwarded_headers + assert forwarded_headers["x-goog-api-key"] == "google-api-key-123" + else: + # LLM provider auth headers should be stripped + assert "x-api-key" not in forwarded_headers + assert "x-goog-api-key" not in forwarded_headers + + # Custom headers should always be forwarded (when forward_client_headers_to_llm_api=True) + assert "x-custom-header" in forwarded_headers + assert forwarded_headers["x-custom-header"] == "custom-value" + + # Proxy Authorization should never be forwarded + assert "authorization" not in forwarded_headers + + print(f"✓ Test passed with forward_llm_provider_auth_headers={forward_llm_auth_headers}") + print(f" Forwarded headers: {list(forwarded_headers.keys())}") + + except Exception as e: + pytest.fail(f"Test failed with forward_llm_auth_headers={forward_llm_auth_headers}: {str(e)}") + finally: + # Clean up + gs = getattr(litellm.proxy.proxy_server, "general_settings") + gs.pop("forward_llm_provider_auth_headers", None) + setattr(litellm.proxy.proxy_server, "general_settings", gs) + + @mock_patch_acompletion() @pytest.mark.asyncio async def test_team_disable_guardrails(mock_acompletion, client_no_auth): diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index f3e7953b8ae..82606511826 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -122,6 +122,249 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch): assert mock_pipeline.execute.call_count == 2 +@pytest.mark.asyncio +async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping): + """Verify that multiple rpush ops are batched into a single pipeline execute""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.rpush = MagicMock() + mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1]) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineRpushOperation + + rpush_list = [ + RedisPipelineRpushOperation(key="key1", values=["a", "b"]), + RedisPipelineRpushOperation(key="key2", values=["c"]), + RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]), + ] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) + + assert result == [3, 5, 1] + assert mock_pipeline.rpush.call_count == 3 + mock_pipeline.rpush.assert_any_call("key1", "a", "b") + mock_pipeline.rpush.assert_any_call("key2", "c") + mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f") + mock_pipeline.execute.assert_called_once() + + +@pytest.mark.asyncio +async def test_async_rpush_pipeline_empty_list_returns_empty(monkeypatch, redis_no_ping): + """Empty rpush_list should return empty list without touching Redis""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + + mock_redis_instance = AsyncMock() + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + result = await redis_cache.async_rpush_pipeline(rpush_list=[]) + + assert result == [] + mock_redis_instance.pipeline.assert_not_called() + + +@pytest.mark.asyncio +async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping): + """Pipeline errors should propagate""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.rpush = MagicMock() + mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineRpushOperation + + rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + with pytest.raises(ConnectionError, match="Redis down"): + await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) + + +@pytest.mark.asyncio +async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping): + """Verify that multiple lpop ops are batched into a single pipeline execute""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + redis_cache.redis_version = "7.0.0" + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.lpop = MagicMock() + mock_pipeline.execute = AsyncMock(return_value=[ + [b"val1", b"val2"], # key1 results + None, # key2 empty + [b"val3"], # key3 results + ]) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineLpopOperation + + lpop_list = [ + RedisPipelineLpopOperation(key="key1", count=10), + RedisPipelineLpopOperation(key="key2", count=10), + RedisPipelineLpopOperation(key="key3", count=5), + ] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) + + assert len(results) == 3 + assert results[0] == ["val1", "val2"] + assert results[1] is None + assert results[2] == ["val3"] + mock_pipeline.execute.assert_called_once() + + +@pytest.mark.asyncio +async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results(monkeypatch, redis_no_ping): + """Verify Redis < 7 fallback issues individual LPOPs and regroups correctly""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + redis_cache.redis_version = "6.2.0" + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.lpop = MagicMock() + + # With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands + # Simulate: key1 has 2 values then None, key2 has 1 value then None + mock_pipeline.execute = AsyncMock(return_value=[ + b"val1", b"val2", None, # 3 LPOPs for key1 + b"val3", None, # 2 LPOPs for key2 + ]) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineLpopOperation + + lpop_list = [ + RedisPipelineLpopOperation(key="key1", count=3), + RedisPipelineLpopOperation(key="key2", count=2), + ] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) + + assert len(results) == 2 + assert results[0] == ["val1", "val2"] # 2 values, None filtered out + assert results[1] == ["val3"] # 1 value, None filtered out + # All 5 individual LPOPs should be queued, but only 1 execute() call + assert mock_pipeline.lpop.call_count == 5 + mock_pipeline.execute.assert_called_once() + + +@pytest.mark.asyncio +async def test_async_rpush_pipeline_raises_on_per_command_error(monkeypatch, redis_no_ping): + """Verify that per-command errors in pipeline results are raised, not silently dropped""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.rpush = MagicMock() + # Simulate: first RPUSH succeeds, second returns a per-command error + mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")]) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineRpushOperation + + rpush_list = [ + RedisPipelineRpushOperation(key="key1", values=["a"]), + RedisPipelineRpushOperation(key="key2", values=["b"]), + ] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + with pytest.raises(Exception, match="WRONGTYPE"): + await redis_cache.async_rpush_pipeline(rpush_list=rpush_list) + + +@pytest.mark.asyncio +async def test_async_lpop_pipeline_raises_on_per_command_error(monkeypatch, redis_no_ping): + """Verify that per-command errors in LPOP pipeline results are raised, not silently dropped""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + redis_cache.redis_version = "7.0.0" + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.lpop = MagicMock() + # Simulate: first LPOP succeeds, second returns a per-command error + mock_pipeline.execute = AsyncMock( + return_value=[[b"val1"], Exception("WRONGTYPE")] + ) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineLpopOperation + + lpop_list = [ + RedisPipelineLpopOperation(key="key1", count=10), + RedisPipelineLpopOperation(key="key2", count=10), + ] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + with pytest.raises(Exception, match="WRONGTYPE"): + await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) + + +@pytest.mark.asyncio +async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping): + """Empty lpop_list should return empty list without touching Redis""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + + mock_redis_instance = AsyncMock() + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + result = await redis_cache.async_lpop_pipeline(lpop_list=[]) + + assert result == [] + mock_redis_instance.pipeline.assert_not_called() + + +@pytest.mark.asyncio +async def test_async_lpop_pipeline_propagates_redis_exception(monkeypatch, redis_no_ping): + """Pipeline errors should propagate""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + redis_cache.redis_version = "7.0.0" + + mock_redis_instance = AsyncMock() + mock_pipeline = MagicMock() + mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline) + mock_pipeline.__aexit__ = AsyncMock(return_value=None) + mock_pipeline.lpop = MagicMock() + mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down")) + mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline) + + from litellm.types.caching import RedisPipelineLpopOperation + + lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)] + + with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance): + with pytest.raises(ConnectionError, match="Redis down"): + await redis_cache.async_lpop_pipeline(lpop_list=lpop_list) + + @pytest.mark.asyncio @pytest.mark.parametrize( "redis_version", diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 5187f733a3c..5b2edcf1ee7 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -147,7 +147,8 @@ class TestResponseCompliance: """Verify status enum values match spec.""" schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] status_prop = schema["properties"]["status"] - expected_statuses = ["UNSPECIFIED", "IN_PROGRESS", "REQUIRES_ACTION", "COMPLETED", "FAILED", "CANCELLED", "INCOMPLETE"] + # Google Interactions API uses lowercase status values (updated Feb 2026) + expected_statuses = ["in_progress", "requires_action", "completed", "failed", "cancelled", "incomplete"] assert status_prop["enum"] == expected_statuses print(f"✓ Status enum values: {expected_statuses}") diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index ebffb56446e..3729e67f0da 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -363,6 +363,87 @@ class TestProxyOAuthHeaderForwarding: assert "authorization" not in cleaned assert cleaned["content-type"] == "application/json" + def test_clean_headers_forwards_anthropic_api_key_when_enabled(self): + """clean_headers should preserve x-api-key when forward_llm_provider_auth_headers=True.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"authorization", b"Bearer sk-proxy-auth"), + (b"x-api-key", b"sk-ant-api03-test-key"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + + # x-api-key should be preserved when flag is True + assert "x-api-key" in cleaned + assert cleaned["x-api-key"] == "sk-ant-api03-test-key" + # Authorization (proxy auth) should still be stripped + assert "authorization" not in cleaned + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_strips_anthropic_api_key_when_disabled(self): + """clean_headers should strip x-api-key when forward_llm_provider_auth_headers=False (default).""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"x-api-key", b"sk-ant-api03-test-key"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=False) + + # x-api-key should be stripped by default + assert "x-api-key" not in cleaned + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_forwards_google_api_key_when_enabled(self): + """clean_headers should preserve x-goog-api-key when forward_llm_provider_auth_headers=True.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"x-goog-api-key", b"google-api-key-123"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + + assert "x-goog-api-key" in cleaned + assert cleaned["x-goog-api-key"] == "google-api-key-123" + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_preserves_oauth_regardless_of_forward_flag(self): + """clean_headers should always preserve OAuth tokens regardless of forward_llm_provider_auth_headers.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"authorization", f"Bearer {FAKE_OAUTH_TOKEN}".encode()), + (b"content-type", b"application/json"), + ] + ) + + # Should preserve OAuth even with flag=False + cleaned_without_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=False) + assert "authorization" in cleaned_without_flag + assert cleaned_without_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + + # Should also preserve OAuth with flag=True + cleaned_with_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + assert "authorization" in cleaned_with_flag + assert cleaned_with_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + def test_add_provider_specific_headers_forwards_oauth(self): """add_provider_specific_headers_to_request should forward OAuth Authorization as a ProviderSpecificHeader scoped to Anthropic-compatible providers.""" diff --git a/tests/test_litellm/llms/hosted_vllm/responses/test_hosted_vllm_responses.py b/tests/test_litellm/llms/hosted_vllm/responses/test_hosted_vllm_responses.py new file mode 100644 index 00000000000..22effbd37f1 --- /dev/null +++ b/tests/test_litellm/llms/hosted_vllm/responses/test_hosted_vllm_responses.py @@ -0,0 +1,103 @@ +""" +Tests for hosted_vllm responses API support. + +Regression test for: https://github.com/BerriAI/litellm/issues +Bug: client.responses.create() raised TypeError: 'NoneType' object is not a mapping +when extra_body=None was passed through the responses→completion pipeline for +hosted_vllm (and any OpenAI-compatible provider using add_provider_specific_params_to_optional_params). +""" + +import json +import os +import sys +from unittest.mock import MagicMock, patch + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +import litellm + + +def _make_mock_chat_completion_response(content: str = "Hello! I'm doing well.") -> dict: + return { + "id": "chatcmpl-test123", + "object": "chat.completion", + "created": 1234567890, + "model": "Qwen/Qwen3-8B", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": content}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + } + + +def _make_mock_http_client(response_body: dict) -> MagicMock: + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = response_body + mock_response.text = json.dumps(response_body) + mock_client.post.return_value = mock_response + return mock_client + + +def test_hosted_vllm_responses_create_with_string_input(): + """ + Regression test: responses.create() with string input must not raise + TypeError: 'NoneType' object is not a mapping. + + Root cause: extra_body=None was passed explicitly through the + responses→completion pipeline. In add_provider_specific_params_to_optional_params(), + passed_params.pop("extra_body", {}) returned None (key existed with value None), + and **None raised TypeError at dict unpacking. + + Fix: normalize None to {} for both extra_body and optional_params["extra_body"]. + """ + mock_client = _make_mock_http_client( + _make_mock_chat_completion_response("I'm doing well, thanks!") + ) + + with patch( + "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + return_value=mock_client, + ): + response = litellm.responses( + model="hosted_vllm/Qwen/Qwen3-8B", + input="Hello, how are you?", + api_base="https://test-vllm.example.com/v1", + api_key="test-key", + ) + + from litellm.types.llms.openai import ResponsesAPIResponse + + assert response is not None + assert isinstance(response, ResponsesAPIResponse) + assert len(response.output) > 0 + output_message = response.output[0] + assert output_message.role == "assistant" # type: ignore[union-attr] + assert len(output_message.content) > 0 # type: ignore[union-attr] + assert "well" in output_message.content[0].text # type: ignore[union-attr] + + +def test_hosted_vllm_responses_create_with_explicit_none_extra_body(): + """ + Directly verify the fix in add_provider_specific_params_to_optional_params: + extra_body=None must not crash when building optional_params. + """ + from litellm.utils import get_optional_params + + # This should not raise TypeError: 'NoneType' object is not a mapping + optional_params = get_optional_params( + model="Qwen/Qwen3-8B", + custom_llm_provider="hosted_vllm", + extra_body=None, + ) + + # extra_body=None should be normalized to an empty dict (or absent) + assert optional_params.get("extra_body") is not None or "extra_body" not in optional_params diff --git a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py index 37959f74086..39ff0a4f4d8 100644 --- a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py +++ b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py @@ -207,10 +207,10 @@ class TestOpenAIChatCompletionStreamingHandler: def test_chunk_parser_maps_reasoning_to_reasoning_content(self): """ Test that chunk_parser maps 'reasoning' field to 'reasoning_content'. - + Some OpenAI-compatible providers (e.g., GLM-5, hosted_vllm) return delta.reasoning, but LiteLLM expects delta.reasoning_content. - + Regression test for: Streaming responses with delta.reasoning field coming back empty when using openai/ or hosted_vllm/ providers. """ @@ -293,3 +293,34 @@ class TestPromptCacheKeyIntegration: prompt_cache_key="test-cache-key-123", ) assert optional_params.get("prompt_cache_key") == "test-cache-key-123" + + +class TestPromptCacheParams: + """Tests for prompt_cache_key and prompt_cache_retention support.""" + + def setup_method(self): + self.config = OpenAIGPTConfig() + + def test_prompt_cache_key_in_supported_params(self): + """Test that prompt_cache_key is in supported params for OpenAI models.""" + supported_params = self.config.get_supported_openai_params("gpt-4o") + assert "prompt_cache_key" in supported_params + + def test_prompt_cache_retention_in_supported_params(self): + """Test that prompt_cache_retention is in supported params for OpenAI models.""" + supported_params = self.config.get_supported_openai_params("gpt-4o") + assert "prompt_cache_retention" in supported_params + + def test_prompt_cache_params_passed_through(self): + """Test that prompt_cache_key and prompt_cache_retention are passed through by map_openai_params.""" + optional_params = self.config.map_openai_params( + non_default_params={ + "prompt_cache_key": "my-cache-key", + "prompt_cache_retention": "24h", + }, + optional_params={}, + model="gpt-4o", + drop_params=False, + ) + assert optional_params.get("prompt_cache_key") == "my-cache-key" + assert optional_params.get("prompt_cache_retention") == "24h" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index c2dbc94f721..b7ae33d1f80 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -1595,3 +1595,146 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission(): assert set(result) == {"direct-server", "group-server"} mock_get_perm.assert_not_called() mock_access_groups.assert_called_once_with(["grp-alpha"]) + + +@pytest.mark.asyncio +class TestAgentMCPPermissions: + """Test agent-level MCP server and tool permission intersection.""" + + async def test_get_allowed_mcp_servers_agent_intersection(self): + """Key/team allow [server_1, server_2]; agent allows [server_1]. Result = [server_1].""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + team_id="test-team", + agent_id="agent-123", + ) + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_key" + ) as mock_key: + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_team" + ) as mock_team: + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_agent" + ) as mock_agent: + mock_key.return_value = ["server_1", "server_2"] + mock_team.return_value = [] + mock_agent.return_value = ["server_1"] + result = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth + ) + assert sorted(result) == ["server_1"] + mock_agent.assert_called_once_with(user_api_key_auth) + + async def test_get_allowed_mcp_servers_agent_no_restriction(self): + """Agent with no object_permission returns []; no intersection applied (inherit key/team).""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-456", + ) + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_key" + ) as mock_key: + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_team" + ) as mock_team: + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_agent" + ) as mock_agent: + mock_key.return_value = ["server_1", "server_2"] + mock_team.return_value = [] + mock_agent.return_value = [] # no agent-level restriction + result = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth + ) + assert sorted(result) == ["server_1", "server_2"] + mock_agent.assert_called_once_with(user_api_key_auth) + + async def test_get_allowed_mcp_servers_key_team_agent_intersection(self): + """Key allows [1, 2], agent allows [2, 3]. Result = [2].""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-789", + ) + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_key" + ) as mock_key: + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_team" + ) as mock_team: + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_agent" + ) as mock_agent: + mock_key.return_value = ["server_1", "server_2"] + mock_team.return_value = [] + mock_agent.return_value = ["server_2", "server_3"] + result = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth + ) + assert sorted(result) == ["server_2"] + + async def test_get_allowed_tools_for_server_agent_intersection(self): + """Key allows [tool_a, tool_b], agent allows [tool_a]. Result = [tool_a].""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-tools", + ) + key_perm = MagicMock() + key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} + team_perm = None + with patch.object( + MCPRequestHandler, "_get_key_object_permission", return_value=key_perm + ): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", + new_callable=AsyncMock, + return_value=team_perm, + ): + with patch.object( + MCPRequestHandler, + "_get_agent_tool_permissions_for_server", + new_callable=AsyncMock, + return_value=["tool_a"], + ) as mock_agent_tools: + result = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server_1", + user_api_key_auth=user_api_key_auth, + ) + assert result == ["tool_a"] + mock_agent_tools.assert_called_once() + call_kwargs = mock_agent_tools.call_args.kwargs + assert call_kwargs["server_id"] == "server_1" + assert call_kwargs["user_api_key_auth"] == user_api_key_auth + + async def test_get_allowed_tools_for_server_agent_no_restriction(self): + """Agent has no tool permissions for server; key/team result is unchanged.""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-no-tools", + ) + key_perm = MagicMock() + key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} + with patch.object( + MCPRequestHandler, "_get_key_object_permission", return_value=key_perm + ): + with patch.object( + MCPRequestHandler, "_get_team_object_permission", + new_callable=AsyncMock, + return_value=None, + ): + with patch.object( + MCPRequestHandler, + "_get_agent_tool_permissions_for_server", + new_callable=AsyncMock, + return_value=None, + ): + result = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server_1", + user_api_key_auth=user_api_key_auth, + ) + assert sorted(result) == ["tool_a", "tool_b"] diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py new file mode 100644 index 00000000000..2a380370c30 --- /dev/null +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py @@ -0,0 +1,194 @@ +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer +from litellm.types.caching import RedisPipelineRpushOperation + + +@pytest.fixture +def mock_redis_cache(): + """Create a mock RedisCache instance""" + mock = AsyncMock() + return mock + + +@pytest.fixture +def redis_update_buffer(mock_redis_cache): + """Create a RedisUpdateBuffer with a mock RedisCache""" + return RedisUpdateBuffer(redis_cache=mock_redis_cache) + + +@pytest.mark.asyncio +async def test_store_in_memory_spend_updates_uses_pipeline(redis_update_buffer, mock_redis_cache): + """ + Verify store_in_memory_spend_updates_in_redis calls async_rpush_pipeline once + with the correct operations and skips empty queues. + """ + mock_redis_cache.async_rpush_pipeline = AsyncMock(return_value=[3, 5, 2]) + + # Create mock queues - only 3 of 7 have data + spend_update_queue = AsyncMock() + spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions = AsyncMock( + return_value={"key_list_transactions": {"key1": 1.0}} + ) + + daily_spend_queue = AsyncMock() + daily_spend_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value={"user_key1": {"spend": 1.0}} + ) + + daily_team_queue = AsyncMock() + daily_team_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value={"team_key1": {"spend": 2.0}} + ) + + # Empty queues + daily_org_queue = AsyncMock() + daily_org_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value={} + ) + + daily_end_user_queue = AsyncMock() + daily_end_user_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value=None + ) + + daily_agent_queue = AsyncMock() + daily_agent_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value={} + ) + + daily_tag_queue = AsyncMock() + daily_tag_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value={} + ) + + await redis_update_buffer.store_in_memory_spend_updates_in_redis( + spend_update_queue=spend_update_queue, + daily_spend_update_queue=daily_spend_queue, + daily_team_spend_update_queue=daily_team_queue, + daily_org_spend_update_queue=daily_org_queue, + daily_end_user_spend_update_queue=daily_end_user_queue, + daily_agent_spend_update_queue=daily_agent_queue, + daily_tag_spend_update_queue=daily_tag_queue, + ) + + # Should be called exactly once (pipeline) + mock_redis_cache.async_rpush_pipeline.assert_called_once() + + # Verify only 3 operations were included (empty ones skipped) + call_args = mock_redis_cache.async_rpush_pipeline.call_args + rpush_list = call_args.kwargs["rpush_list"] + assert len(rpush_list) == 3 + + +@pytest.mark.asyncio +async def test_store_in_memory_spend_updates_all_empty_returns_early( + redis_update_buffer, mock_redis_cache +): + """ + When all queues are empty, pipeline should never be called. + """ + mock_redis_cache.async_rpush_pipeline = AsyncMock() + + # All queues return empty + empty_queue = AsyncMock() + empty_queue.flush_and_get_aggregated_db_spend_update_transactions = AsyncMock( + return_value={} + ) + empty_daily_queue = AsyncMock() + empty_daily_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock( + return_value={} + ) + + await redis_update_buffer.store_in_memory_spend_updates_in_redis( + spend_update_queue=empty_queue, + daily_spend_update_queue=empty_daily_queue, + daily_team_spend_update_queue=empty_daily_queue, + daily_org_spend_update_queue=empty_daily_queue, + daily_end_user_spend_update_queue=empty_daily_queue, + daily_agent_spend_update_queue=empty_daily_queue, + daily_tag_spend_update_queue=empty_daily_queue, + ) + + mock_redis_cache.async_rpush_pipeline.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_all_transactions_from_redis_buffer_pipeline( + redis_update_buffer, mock_redis_cache +): + """ + Verify get_all_transactions_from_redis_buffer_pipeline correctly parses + and aggregates results from async_lpop_pipeline. + """ + # Simulate pipeline results: slot 0 = spend updates, slots 1-6 = daily categories + db_spend_json = json.dumps( + { + "key_list_transactions": {"key1": 1.0, "key2": 2.0}, + "user_list_transactions": {"user1": 0.5}, + "end_user_list_transactions": {}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + } + ) + daily_user_json = json.dumps({"user_key1": {"spend": 1.0, "api_requests": 1}}) + daily_team_json = json.dumps({"team_key1": {"spend": 2.0, "api_requests": 2}}) + + mock_redis_cache.async_lpop_pipeline = AsyncMock( + return_value=[ + [db_spend_json], # slot 0: db spend updates + [daily_user_json], # slot 1: daily user + [daily_team_json], # slot 2: daily team + None, # slot 3: daily org (empty) + None, # slot 4: daily end-user (empty) + None, # slot 5: daily agent (empty) + None, # slot 6: daily tag (empty) + ] + ) + + result = await redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline() + + assert len(result) == 7 + db_spend, daily_user, daily_team, daily_org, daily_end_user, daily_agent, daily_tag = result + + # Verify db spend was parsed correctly + assert db_spend is not None + assert db_spend["key_list_transactions"]["key1"] == 1.0 + assert db_spend["key_list_transactions"]["key2"] == 2.0 + assert db_spend["user_list_transactions"]["user1"] == 0.5 + + # Verify daily user was parsed + assert daily_user is not None + assert daily_user["user_key1"]["spend"] == 1.0 + + # Verify daily team was parsed + assert daily_team is not None + assert daily_team["team_key1"]["spend"] == 2.0 + + # Verify empty slots + assert daily_org is None + assert daily_end_user is None + assert daily_agent is None + assert daily_tag is None + + # Verify pipeline was called once with correct keys + mock_redis_cache.async_lpop_pipeline.assert_called_once() + + +@pytest.mark.asyncio +async def test_get_all_transactions_from_redis_buffer_pipeline_no_redis(): + """When redis_cache is None, should return all Nones""" + buffer = RedisUpdateBuffer(redis_cache=None) + result = await buffer.get_all_transactions_from_redis_buffer_pipeline() + assert result == (None, None, None, None, None, None, None) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index add7293285a..abd59a66a36 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1254,3 +1254,46 @@ async def test_daily_agent_receives_deepcopied_payload(): # But it should have equivalent content assert captured_agent_payloads[0]["model"] == "gpt-4" assert captured_agent_payloads[0]["spend"] == 0.1 + + +@pytest.mark.asyncio +async def test_commit_spend_updates_uses_pipeline(): + """ + Verify that _commit_spend_updates_to_db_with_redis uses + get_all_transactions_from_redis_buffer_pipeline instead of 7 individual calls. + """ + db_writer = DBSpendUpdateWriter() + + mock_redis_update_buffer = AsyncMock() + mock_redis_update_buffer.store_in_memory_spend_updates_in_redis = AsyncMock() + # Return all-None tuple (no data to commit) + mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock( + return_value=(None, None, None, None, None, None, None) + ) + db_writer.redis_update_buffer = mock_redis_update_buffer + + mock_pod_lock_manager = AsyncMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + db_writer.pod_lock_manager = mock_pod_lock_manager + + mock_prisma_client = MagicMock() + mock_proxy_logging = MagicMock() + + await db_writer._commit_spend_updates_to_db_with_redis( + prisma_client=mock_prisma_client, + n_retry_times=1, + proxy_logging_obj=mock_proxy_logging, + ) + + # Pipeline method should be called once + mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline.assert_called_once() + + # Individual methods should NOT be called + mock_redis_update_buffer.get_all_update_transactions_from_redis_buffer.assert_not_called() + mock_redis_update_buffer.get_all_daily_spend_update_transactions_from_redis_buffer.assert_not_called() + mock_redis_update_buffer.get_all_daily_team_spend_update_transactions_from_redis_buffer.assert_not_called() + mock_redis_update_buffer.get_all_daily_org_spend_update_transactions_from_redis_buffer.assert_not_called() + mock_redis_update_buffer.get_all_daily_end_user_spend_update_transactions_from_redis_buffer.assert_not_called() + mock_redis_update_buffer.get_all_daily_agent_spend_update_transactions_from_redis_buffer.assert_not_called() + mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer.assert_not_called() diff --git a/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py b/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py index 2f2eaa905be..205d724c2b0 100644 --- a/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py +++ b/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py @@ -2,10 +2,9 @@ """ Test to verify the Google GenAI proxy API endpoints """ -import asyncio import os import sys -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, patch import pytest @@ -13,7 +12,6 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -import litellm def test_google_generate_content_endpoint(): @@ -401,3 +399,123 @@ def test_google_generate_content_with_image_config(): assert "contents" in called_data assert len(called_data["contents"]) == 1 assert called_data["contents"][0]["role"] == "user" + + +def test_google_generate_content_metadata_and_trace_id_callbacks(): + """Test that google_generate_content sets litellm_call_id and logging_obj for callbacks (e.g. S3, Langfuse)""" + try: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy.google_endpoints.endpoints import router as google_router + except ImportError as e: + pytest.skip(f"Skipping test due to missing dependency: {e}") + + # Create a FastAPI app and include the router + app = FastAPI() + app.include_router(google_router) + + # Create a test client + client = TestClient(app) + + # Mock all required proxy server dependencies + with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch( + "litellm.proxy.proxy_server.general_settings", {} + ), patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config, patch( + "litellm.proxy.proxy_server.version", "1.0.0" + ), patch( + "litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request" + ) as mock_add_data: + mock_router.agenerate_content = AsyncMock(return_value={"test": "response"}) + + # Mock add_litellm_data_to_request to return data with metadata + async def mock_add_litellm_data( + data, request, user_api_key_dict, proxy_config, general_settings, version + ): + # Simulate adding user metadata + data["litellm_metadata"] = { + "user_api_key_user_id": "test-user-id", + } + return data + + mock_add_data.side_effect = mock_add_litellm_data + + # Send a request to the endpoint with x-litellm-call-id header + test_call_id = "test-custom-call-id" + response = client.post( + "/v1beta/models/test-model:generateContent", + json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]}, + headers={ + "Authorization": "Bearer sk-test-key", + "x-litellm-call-id": test_call_id, + }, + ) + + assert response.status_code == 200 + + mock_router.agenerate_content.assert_called_once() + call_args = mock_router.agenerate_content.call_args + called_data = call_args[1] + + # Verify that the litellm_logging_obj got assigned in the final called_data to router + assert "litellm_logging_obj" in called_data + assert "litellm_call_id" in called_data + assert called_data["litellm_call_id"] == test_call_id + + +def test_google_stream_generate_content_metadata_and_trace_id_callbacks(): + """Test that google_stream_generate_content sets litellm_call_id and logging_obj for callbacks""" + try: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy.google_endpoints.endpoints import router as google_router + except ImportError as e: + pytest.skip(f"Skipping test due to missing dependency: {e}") + + app = FastAPI() + app.include_router(google_router) + client = TestClient(app) + + mock_stream = AsyncMock() + mock_stream.__aiter__ = lambda self: mock_stream + mock_stream.__anext__.side_effect = StopAsyncIteration + + with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch( + "litellm.proxy.proxy_server.general_settings", {} + ), patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config, patch( + "litellm.proxy.proxy_server.version", "1.0.0" + ), patch( + "litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request" + ) as mock_add_data: + mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream) + + async def mock_add_litellm_data( + data, request, user_api_key_dict, proxy_config, general_settings, version + ): + data["litellm_metadata"] = { + "user_api_key_user_id": "test-user-id", + } + return data + + mock_add_data.side_effect = mock_add_litellm_data + + test_call_id = "test-custom-stream-call-id" + response = client.post( + "/v1beta/models/test-model:streamGenerateContent", + json={"contents": [{"role": "user", "parts": [{"text": "Hello stream"}]}]}, + headers={ + "Authorization": "Bearer sk-test-key", + "x-litellm-call-id": test_call_id, + }, + ) + + assert response.status_code == 200 + + mock_router.agenerate_content_stream.assert_called_once() + call_args = mock_router.agenerate_content_stream.call_args + called_data = call_args[1] + + assert "litellm_logging_obj" in called_data + assert "litellm_call_id" in called_data + assert called_data["litellm_call_id"] == test_call_id diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index e2c05f0ad37..76f9c39acd0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -24,12 +24,29 @@ from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType from litellm.types.utils import Choices, Message, ModelResponse -def _make_mock_session_iterator(json_response): +def _make_mock_session_iterator( + json_response, status=200, content_type="application/json", text_response="" +): """Create a mock _get_session_iterator that yields a session returning json_response.""" @asynccontextmanager async def mock_iterator(): class MockResponse: + def __init__(self): + self.status = status + self.content_type = content_type + self.headers = {"Content-Type": content_type} + + async def text(self): + if text_response: + return text_response + import json + + try: + return json.dumps(json_response) + except Exception: + return str(json_response) + async def json(self): return json_response @@ -41,6 +58,7 @@ def _make_mock_session_iterator(json_response): class MockSession: def post(self, *args, **kwargs): + self.last_kwargs = kwargs return MockResponse() async def __aenter__(self): @@ -1444,3 +1462,149 @@ def test_deny_list_and_score_threshold_combined(): # EMAIL_ADDRESS passes both filters assert len(filtered) == 1 assert filtered[0]["entity_type"] == "EMAIL_ADDRESS" + + +@pytest.mark.asyncio +async def test_analyze_text_non_json_content_type_fail_closed(): + """ + Test that analyze_text raises GuardrailRaisedException when Presidio health + endpoint returns text/html and fail-closed is enabled. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + pii_entities_config={"PERSON": PiiAction.BLOCK}, + mock_testing=False, + ) + + mock_iterator = _make_mock_session_iterator( + json_response=None, + status=200, + content_type="text/html; charset=utf-8", + text_response="Presidio Analyzer service is up.", + ) + + with patch.object(guardrail, "_get_session_iterator", mock_iterator): + with pytest.raises(GuardrailRaisedException) as exc_info: + await guardrail.analyze_text( + text="Hello world", + presidio_config=None, + request_data={}, + ) + assert "expected application/json Content-Type" in str(exc_info.value) + assert "text/html" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_analyze_text_non_json_content_type_fail_open(): + """ + Test that analyze_text returns empty list when Presidio returns text/html + and fail-closed is NOT enabled. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + mock_testing=False, + ) + + mock_iterator = _make_mock_session_iterator( + json_response=None, + status=200, + content_type="text/html; charset=utf-8", + text_response="Presidio Analyzer service is up.", + ) + + with patch.object(guardrail, "_get_session_iterator", mock_iterator): + results = await guardrail.analyze_text( + text="Hello world", + presidio_config=None, + request_data={}, + ) + assert results == [] + + +@pytest.mark.asyncio +async def test_analyze_text_http_error_status(): + """ + Test that analyze_text handles 5xx HTTP errors properly. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + pii_entities_config={"PERSON": PiiAction.BLOCK}, + mock_testing=False, + ) + + mock_iterator = _make_mock_session_iterator( + json_response=None, + status=500, + content_type="text/plain", + text_response="Internal Server Error", + ) + + with patch.object(guardrail, "_get_session_iterator", mock_iterator): + with pytest.raises(GuardrailRaisedException) as exc_info: + await guardrail.analyze_text( + text="Hello world", + presidio_config=None, + request_data={}, + ) + assert "HTTP 500" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_anonymize_text_non_json_content_type(): + """ + Test that anonymize_text raises Exception for non-JSON responses. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + mock_testing=False, + ) + + mock_iterator = _make_mock_session_iterator( + json_response=None, + status=200, + content_type="text/html", + text_response="Presidio Anonymizer service is up.", + ) + + with patch.object(guardrail, "_get_session_iterator", mock_iterator): + with pytest.raises( + Exception, match="Presidio anonymizer returned non-JSON Content-Type" + ): + await guardrail.anonymize_text( + text="Hello world", + analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}], + output_parse_pii=False, + masked_entity_count={}, + ) + + +@pytest.mark.asyncio +async def test_anonymize_text_http_error_status(): + """ + Test that anonymize_text raises Exception on HTTP error. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + mock_testing=False, + ) + + mock_iterator = _make_mock_session_iterator( + json_response=None, + status=502, + content_type="text/plain", + text_response="Bad Gateway", + ) + + with patch.object(guardrail, "_get_session_iterator", mock_iterator): + with pytest.raises(Exception, match="Presidio anonymizer returned HTTP 502"): + await guardrail.anonymize_text( + text="Hello world", + analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}], + output_parse_pii=False, + masked_entity_count={}, + ) diff --git a/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py new file mode 100644 index 00000000000..deb1c483b87 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_max_iterations_limiter.py @@ -0,0 +1,106 @@ +""" +Unit Tests for the max iterations limiter for the proxy. + +Tests that session-scoped iteration counting works correctly: +- Enforces max_iterations per session_id +- Different sessions have independent counters +""" + +import pytest +from fastapi import HTTPException + +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler +from litellm.proxy.utils import InternalUsageCache + + +@pytest.mark.asyncio +async def test_max_iterations_basic_enforcement(): + """ + Test that max_iterations is enforced per session_id. + + - 3 requests with the same session_id should succeed when max_iterations=3 + - 4th request should raise 429 + """ + local_cache = DualCache() + handler = _PROXY_MaxIterationsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-key-1234", metadata={"max_iterations": 3} + ) + + # First 3 requests should succeed + for i in range(3): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"metadata": {"session_id": "session-abc"}}, + call_type="", + ) + + # 4th request should fail with 429 + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"metadata": {"session_id": "session-abc"}}, + call_type="", + ) + assert exc_info.value.status_code == 429 + assert "max_iterations" in str(exc_info.value.detail).lower() + + +@pytest.mark.asyncio +async def test_max_iterations_different_sessions_independent(): + """ + Test that different session_ids have independent iteration counters. + + - Session A and Session B each get their own max_iterations budget + - Exhausting Session A does not affect Session B + """ + local_cache = DualCache() + handler = _PROXY_MaxIterationsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-key-5678", metadata={"max_iterations": 2} + ) + + # Session A: 2 calls succeed + for _ in range(2): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"metadata": {"session_id": "session-A"}}, + call_type="", + ) + + # Session B: 2 calls succeed (independent counter) + for _ in range(2): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"metadata": {"session_id": "session-B"}}, + call_type="", + ) + + # Session A: 3rd call fails + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"metadata": {"session_id": "session-A"}}, + call_type="", + ) + assert exc_info.value.status_code == 429 + + # Session B: 3rd call also fails + with pytest.raises(HTTPException): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"metadata": {"session_id": "session-B"}}, + call_type="", + ) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index e8765cf78ca..c46b8df5efc 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -169,6 +169,223 @@ async def test_track_cost_callback_skips_when_no_standard_logging_object(): mock_proxy_logging.failed_tracking_alert.assert_not_called() +@pytest.mark.asyncio +async def test_enrich_failure_metadata_with_team_alias(): + """ + When team_id is set but team_alias is missing (and key_alias is present), + _enrich_failure_metadata_with_key_info should look up the team from cache + and populate user_api_key_team_alias. + """ + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "my-team-alias" + + with patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ): + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": "my-key-alias", # already set + "user_api_key_team_id": "test_team_id", + "user_api_key_team_alias": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + assert result["user_api_key_team_alias"] == "my-team-alias" + + +@pytest.mark.asyncio +async def test_enrich_failure_metadata_with_full_key_lookup(): + """ + When all key fields are null (auth error 401 scenario), _enrich_failure_metadata_with_key_info + should look up the key object from cache/DB and populate alias, user_id, team_id, + then look up the team to get team_alias. + """ + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "fetched-key-alias" + mock_key_obj.user_id = "fetched-user-id" + mock_key_obj.team_id = "fetched-team-id" + mock_key_obj.org_id = "fetched-org-id" + + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "fetched-team-alias" + + with patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ), patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ): + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": None, # all null - simulates auth error path + "user_api_key_user_id": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + "user_api_key_org_id": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + assert result["user_api_key_alias"] == "fetched-key-alias" + assert result["user_api_key_user_id"] == "fetched-user-id" + assert result["user_api_key_team_id"] == "fetched-team-id" + assert result["user_api_key_org_id"] == "fetched-org-id" + assert result["user_api_key_team_alias"] == "fetched-team-alias" + + +@pytest.mark.asyncio +async def test_enrich_failure_metadata_skips_when_team_alias_present(): + """ + When team_alias is already populated, _enrich_failure_metadata_with_key_info + should not perform a team cache lookup. + """ + with patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + ) as mock_get_key, patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team: + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": "existing-alias", + "user_api_key_team_id": "test_team_id", + "user_api_key_team_alias": "already-set", + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + assert result["user_api_key_team_alias"] == "already-set" + mock_get_key.assert_not_called() + mock_get_team.assert_not_called() + + +@pytest.mark.asyncio +async def test_enrich_failure_metadata_skips_when_no_api_key(): + """ + When api_key hash is absent, _enrich_failure_metadata_with_key_info should + not perform any lookups. + """ + with patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + ) as mock_get_key: + metadata = { + "user_api_key": None, + "user_api_key_alias": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + mock_get_key.assert_not_called() + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): + """ + Simulates a 401 ProxyException (e.g. can_key_call_model). In this case + UserAPIKeyAuth is created with only api_key set. The failure hook should + look up the key and team from cache/DB to populate all missing fields. + """ + logger = _ProxyDBLogger() + + # This is what auth_exception_handler creates for 401 errors + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed_key", + # key_alias, user_id, team_id, team_alias are all None + ) + + request_data = { + "model": "claude-haiku-4-5", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + "litellm_params": {}, + } + + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "my-key-alias" + mock_key_obj.user_id = "my-user-id" + mock_key_obj.team_id = "my-team-id" + mock_key_obj.org_id = None + + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "my-team-alias" + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ), patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ): + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("401 - model not allowed"), + user_api_key_dict=user_api_key_dict, + ) + + mock_update_database.assert_called_once() + call_args = mock_update_database.call_args[1] + metadata = call_args["kwargs"]["litellm_params"]["metadata"] + assert metadata["user_api_key_alias"] == "my-key-alias" + assert metadata["user_api_key_user_id"] == "my-user-id" + assert metadata["user_api_key_team_id"] == "my-team-id" + assert metadata["user_api_key_team_alias"] == "my-team-alias" + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_enriches_missing_team_alias(): + """ + When user_api_key_dict has a team_id but no team_alias, async_post_call_failure_hook + should look up the team from cache and populate user_api_key_team_alias in the + spend log metadata written to the DB. + """ + logger = _ProxyDBLogger() + + user_api_key_dict = UserAPIKeyAuth( + api_key="test_api_key", + key_alias="test_alias", + user_id="test_user_id", + team_id="test_team_id", + team_alias=None, # Missing - simulates regular key auth where SQL view omits team_alias + ) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + "litellm_params": {}, + } + + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "enriched-team-alias" + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ): + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("Provider rate limit"), + user_api_key_dict=user_api_key_dict, + ) + + mock_update_database.assert_called_once() + call_args = mock_update_database.call_args[1] + metadata = call_args["kwargs"]["litellm_params"]["metadata"] + assert metadata["user_api_key_team_alias"] == "enriched-team-alias" + assert metadata["user_api_key_team_id"] == "test_team_id" + + @pytest.mark.asyncio @pytest.mark.parametrize("model_value", [None, ""]) async def test_track_cost_callback_skips_for_falsy_model_and_no_slo(model_value): diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index ffb4e955426..05df3c2dcbb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -5623,6 +5623,7 @@ async def test_rotate_master_key_model_data_valid_for_prisma( litellm_params/model_info are JSON strings (create_many expects dicts). """ from unittest.mock import AsyncMock, MagicMock + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.key_management_endpoints import ( _rotate_master_key, @@ -6157,3 +6158,56 @@ async def test_get_member_team_ids(): # Should return team-A and team-B (user is a member of both) # Should NOT return team-C (user is not in members list) assert sorted(result) == ["team-A", "team-B"] + + +@pytest.mark.asyncio +async def test_generate_key_with_agent_id(): + """Test that agent_id is accepted in GenerateKeyRequest and passed to generate_key_helper_fn.""" + from litellm.proxy._types import GenerateKeyRequest + + # Verify GenerateKeyRequest accepts agent_id + request = GenerateKeyRequest( + key_alias="agent-test-key", + agent_id="test-agent-123", + models=[], + ) + assert request.agent_id == "test-agent-123" + data_json = request.model_dump(exclude_unset=True, exclude_none=True) + assert data_json["agent_id"] == "test-agent-123" + + +@pytest.mark.asyncio +async def test_generate_key_helper_fn_agent_id(): + """Test that generate_key_helper_fn passes agent_id into the insert_data call.""" + from unittest.mock import AsyncMock, MagicMock, call, patch + + import litellm.proxy.management_endpoints.key_management_endpoints as km + + mock_prisma_client = AsyncMock() + mock_insert = AsyncMock( + return_value=MagicMock( + token="sk-test", + created_at=None, + updated_at=None, + litellm_budget_table=None, + ) + ) + mock_prisma_client.insert_data = mock_insert + + with patch.object(km, "prisma_client", mock_prisma_client): + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + await generate_key_helper_fn( + request_type="key", + agent_id="test-agent-456", + key_alias="test-agent-key", + models=[], + table_name="key", + ) + + assert mock_insert.called, "insert_data was never called" + # insert_data is called as insert_data(data=key_data, ...) + call_kwargs = mock_insert.call_args.kwargs + key_data = call_kwargs.get("data", {}) + assert key_data.get("agent_id") == "test-agent-456", ( + f"Expected agent_id='test-agent-456' in key_data, got: {key_data.get('agent_id')}" + ) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index c3639e88f6e..e439dfd693c 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -274,6 +274,7 @@ ignored_keys = [ "endTime", "completionStartTime", "endTime", + "request_duration_ms", "organization_id", "metadata.model_map_information", "metadata.usage_object", @@ -606,6 +607,82 @@ async def test_ui_view_spend_logs_sort_validation_errors( app.dependency_overrides.pop(ps.user_api_key_auth, None) +@pytest.mark.asyncio +async def test_ui_view_spend_logs_sort_by_request_duration_ms(client, monkeypatch): + """Test that request_duration_ms is accepted as a valid sort_by field.""" + base_logs = [ + { + "request_id": "req_fast", + "api_key": "sk-test-key", + "user": "user1", + "spend": 0.10, + "total_tokens": 100, + "request_duration_ms": 100, + "startTime": "2025-01-01T00:00:00+00:00", + "endTime": "2025-01-01T00:00:00.100000+00:00", + "model": "gpt-4", + }, + { + "request_id": "req_slow", + "api_key": "sk-test-key", + "user": "user1", + "spend": 0.05, + "total_tokens": 50, + "request_duration_ms": 5000, + "startTime": "2025-01-01T00:00:01+00:00", + "endTime": "2025-01-01T00:00:06+00:00", + "model": "gpt-4", + }, + ] + + async def mock_count(*args, **kwargs): + return len(base_logs) + + async def mock_query_raw(sql_query, *params): + reverse = "DESC" in sql_query + sorted_logs = sorted( + base_logs, key=lambda x: x.get("request_duration_ms", 0), reverse=reverse + ) + page_size = params[-2] if len(params) >= 2 else 50 + skip = params[-1] if len(params) >= 1 else 0 + return sorted_logs[skip : skip + page_size] + + class MockPrismaClient: + def __init__(self): + self.db = MagicMock() + self.db.litellm_spendlogs = MagicMock() + self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) + self.db.query_raw = AsyncMock(side_effect=mock_query_raw) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + response = client.get( + "/spend/logs/ui", + params={ + "start_date": "2024-12-25 00:00:00", + "end_date": "2025-01-02 23:59:59", + "sort_by": "request_duration_ms", + "sort_order": "asc", + }, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200, response.text + data = response.json() + actual_ids = [log["request_id"] for log in data["data"]] + assert actual_ids == ["req_fast", "req_slow"] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + @pytest.mark.asyncio async def test_ui_view_spend_logs_with_team_id(client, monkeypatch): mock_spend_logs = [ diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 47a327f01f6..06a544e7f4d 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -20,6 +20,7 @@ from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, REDACTED_BY_LITEL from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy.spend_tracking.spend_tracking_utils import ( _get_proxy_server_request_for_spend_logs_payload, + _get_request_duration_ms, _get_response_for_spend_logs_payload, _get_spend_logs_metadata, _get_vector_store_request_for_spend_logs_payload, @@ -1232,3 +1233,50 @@ def test_get_logging_payload_handles_missing_retry_info_gracefully(): metadata.get("max_retries") is None ), "max_retries should be None when not provided" + +def test_get_request_duration_ms_normal(): + """Test that request duration is correctly computed in milliseconds.""" + start = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + end = datetime.datetime(2025, 1, 1, 0, 0, 2, 500000, tzinfo=timezone.utc) # 2.5s later + result = _get_request_duration_ms(start, end) + assert result == 2500 + + +def test_get_request_duration_ms_sub_millisecond(): + """Test that sub-millisecond durations are truncated to int.""" + start = datetime.datetime(2025, 1, 1, 0, 0, 0, 0, tzinfo=timezone.utc) + end = datetime.datetime(2025, 1, 1, 0, 0, 0, 500, tzinfo=timezone.utc) # 0.5ms + result = _get_request_duration_ms(start, end) + assert result == 0 + + +def test_get_request_duration_ms_zero(): + """Test that identical start and end times produce 0.""" + t = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + result = _get_request_duration_ms(t, t) + assert result == 0 + + +def test_get_logging_payload_includes_request_duration_ms(): + """Test that get_logging_payload populates request_duration_ms.""" + start_time = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + end_time = datetime.datetime(2025, 1, 1, 0, 0, 3, tzinfo=timezone.utc) # 3s later + + kwargs = { + "model": "gpt-4", + "litellm_params": {"api_base": "https://api.openai.com"}, + "standard_logging_object": None, + } + response_obj = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}} + + with patch("litellm.proxy.proxy_server.master_key", None), \ + patch("litellm.proxy.proxy_server.general_settings", {}): + payload = get_logging_payload( + kwargs=kwargs, + response_obj=response_obj, + start_time=start_time, + end_time=end_time, + ) + + assert payload["request_duration_ms"] == 3000 + diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index ccae9fb5425..4c91d0ae91e 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -12,6 +12,7 @@ sys.path.insert(0, os.path.abspath("../../..")) from litellm.proxy.health_endpoints._health_endpoints import ( _aggregate_health_check_results, _build_model_param_to_info_mapping, + _perform_health_check_and_save, _save_background_health_checks_to_db, _save_health_check_results_if_changed, _save_health_check_to_db, @@ -466,5 +467,43 @@ async def test_get_all_latest_health_checks_without_model_id(mock_prisma): assert result[0].checked_at == mock_check2.checked_at # Latest +@pytest.mark.asyncio +async def test_perform_health_check_and_save_passes_model_id_to_perform_health_check(): + """Test that _perform_health_check_and_save passes model_id to perform_health_check so health checks run by model id.""" + model_list = [ + { + "model_name": "gpt-4", + "model_info": {"id": "deployment-abc"}, + "litellm_params": {"model": "gpt-4"}, + }, + ] + healthy = [{"model": "gpt-4"}] + unhealthy = [] + + async def mock_perform_health_check(model_list, model=None, cli_model=None, details=True, model_id=None): + return healthy, unhealthy + + with patch( + "litellm.proxy.health_endpoints._health_endpoints.perform_health_check", + side_effect=mock_perform_health_check, + ) as mock_perform: + result = await _perform_health_check_and_save( + model_list=model_list, + target_model=None, + cli_model=None, + details=True, + prisma_client=None, + start_time=0.0, + user_id="user-1", + model_id="deployment-abc", + ) + + mock_perform.assert_called_once() + call_kwargs = mock_perform.call_args[1] + assert call_kwargs["model_id"] == "deployment-abc" + assert result["healthy_count"] == 1 + assert result["unhealthy_count"] == 0 + + if __name__ == "__main__": pytest.main([__file__]) \ No newline at end of file diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 3787f451ad3..cc04e674003 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -90,6 +90,7 @@ "version": "5.2.0", "resolved": "https://registry.npmjs.org/@alloc/quick-lru/-/quick-lru-5.2.0.tgz", "integrity": "sha512-UrcABB+4bUrFABwbluTIBErXwvbsU/V7TZWfmbgJfbkwiBuziS9gxdODUyuiecfdGQ85jglMW6juS3+z5TsKLw==", + "dev": true, "license": "MIT", "engines": { "node": ">=10" @@ -1771,6 +1772,7 @@ "version": "0.3.13", "resolved": "https://registry.npmjs.org/@jridgewell/gen-mapping/-/gen-mapping-0.3.13.tgz", "integrity": "sha512-2kkt/7niJ6MgEPxF0bYdQ6etZaA+fQvDcLKckhy1yIQOzaoKjBBjSj63/aLVjYE3qhRt5dvM+uUyfCg6UKCBbA==", + "dev": true, "license": "MIT", "dependencies": { "@jridgewell/sourcemap-codec": "^1.5.0", @@ -1781,6 +1783,7 @@ "version": "3.1.2", "resolved": "https://registry.npmjs.org/@jridgewell/resolve-uri/-/resolve-uri-3.1.2.tgz", "integrity": "sha512-bRISgCIjP20/tbWSPWMEi54QVPRZExkuD9lJL+UIxUKtwVJA8wW1Trb1jMs1RFXo1CBTNZ/5hpC9QvmKWdopKw==", + "dev": true, "license": "MIT", "engines": { "node": ">=6.0.0" @@ -1790,12 +1793,14 @@ "version": "1.5.5", "resolved": "https://registry.npmjs.org/@jridgewell/sourcemap-codec/-/sourcemap-codec-1.5.5.tgz", "integrity": "sha512-cYQ9310grqxueWbl+WuIUIaiUaDcj7WOq5fVhEljNVgRfOUhY9fy2zTvfoqWsnebh8Sl70VScFbICvJnLKB0Og==", + "dev": true, "license": "MIT" }, "node_modules/@jridgewell/trace-mapping": { "version": "0.3.31", "resolved": "https://registry.npmjs.org/@jridgewell/trace-mapping/-/trace-mapping-0.3.31.tgz", "integrity": "sha512-zzNR+SdQSDJzc8joaeP8QQoCQr8NuYx2dIIytl1QeBEZHJ9uW6hebsrYgbz8hJwUQao3TWCMtmfV8Nu1twOLAw==", + "dev": true, "license": "MIT", "dependencies": { "@jridgewell/resolve-uri": "^3.1.0", @@ -1973,6 +1978,7 @@ "version": "2.1.5", "resolved": "https://registry.npmjs.org/@nodelib/fs.scandir/-/fs.scandir-2.1.5.tgz", "integrity": "sha512-vq24Bq3ym5HEQm2NKCr3yXDwjc7vTsEThRDnkp2DK9p1uqLR+DHurm/NOTo0KG7HYHU7eppKZj3MyqYuMBf62g==", + "dev": true, "license": "MIT", "dependencies": { "@nodelib/fs.stat": "2.0.5", @@ -1986,6 +1992,7 @@ "version": "2.0.5", "resolved": "https://registry.npmjs.org/@nodelib/fs.stat/-/fs.stat-2.0.5.tgz", "integrity": "sha512-RkhPPp2zrqDAQA/2jNhnztcPAlv64XdhIp7a7454A5ovI7Bukxgt7MX7udwAu3zg1DcpPU0rz3VV1SeaqvY4+A==", + "dev": true, "license": "MIT", "engines": { "node": ">= 8" @@ -1995,6 +2002,7 @@ "version": "1.2.8", "resolved": "https://registry.npmjs.org/@nodelib/fs.walk/-/fs.walk-1.2.8.tgz", "integrity": "sha512-oGB+UxlgWcgQkgwo8GcEGwemoTFt3FIO9ababBmaGwXIoBKZ+GTy0pP185beGg7Llih/NSHSV2XAs1lnznocSg==", + "dev": true, "license": "MIT", "dependencies": { "@nodelib/fs.scandir": "2.1.5", @@ -2318,7 +2326,7 @@ "version": "1.58.1", "resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.58.1.tgz", "integrity": "sha512-6LdVIUERWxQMmUSSQi0I53GgCBYgM2RpGngCPY7hSeju+VrKjq3lvs7HpJoPbDiY5QM5EYRtRX5fvrinnMAz3w==", - "devOptional": true, + "dev": true, "license": "Apache-2.0", "dependencies": { "playwright": "1.58.1" @@ -3423,12 +3431,14 @@ "version": "15.7.15", "resolved": "https://registry.npmjs.org/@types/prop-types/-/prop-types-15.7.15.tgz", "integrity": "sha512-F6bEyamV9jKGAFBEmlQnesRPGOQqS2+Uwi0Em15xenOxHaf2hv6L8YCVn3rPdPJOiJfPiCnLIRyvwVaqMY3MIw==", + "dev": true, "license": "MIT" }, "node_modules/@types/react": { "version": "18.2.48", "resolved": "https://registry.npmjs.org/@types/react/-/react-18.2.48.tgz", "integrity": "sha512-qboRCl6Ie70DQQG9hhNREz81jqC1cs9EVNcjQ1AU+jH6NFfSAhVVbrrY/+nSF+Bsk4AOwm9Qa61InvMCyV+H3w==", + "dev": true, "license": "MIT", "dependencies": { "@types/prop-types": "*", @@ -3470,6 +3480,7 @@ "version": "0.26.0", "resolved": "https://registry.npmjs.org/@types/scheduler/-/scheduler-0.26.0.tgz", "integrity": "sha512-WFHp9YUJQ6CKshqoC37iOlHnQSmxNc795UhB26CyBBttrN9svdIrUjl/NjnNmfcwtncN0h/0PPAFWv9ovP8mLA==", + "dev": true, "license": "MIT" }, "node_modules/@types/unist": { @@ -4330,12 +4341,14 @@ "version": "1.3.0", "resolved": "https://registry.npmjs.org/any-promise/-/any-promise-1.3.0.tgz", "integrity": "sha512-7UvmKalWRt1wgjL1RrGxoSJW/0QZFIegpeGvZG9kjp8vrRu55XTHbwnqq2GpXm9uLbcuhxm3IqX9OB4MZR1b2A==", + "dev": true, "license": "MIT" }, "node_modules/anymatch": { "version": "3.1.3", "resolved": "https://registry.npmjs.org/anymatch/-/anymatch-3.1.3.tgz", "integrity": "sha512-KMReFUr0B4t+D+OBkjR3KYqvocp2XaSzO55UcB6mgQMd3KbcE+mWTyvVV7D/zsdEbNnV6acZUutkiHQXvTr1Rw==", + "dev": true, "license": "ISC", "dependencies": { "normalize-path": "^3.0.0", @@ -4349,6 +4362,7 @@ "version": "2.3.1", "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.1.tgz", "integrity": "sha512-JU3teHTNjmE2VCGFzuY8EXzCDVwEqB2a8fsIvwaStHhAWJEeVd1o1QD80CU6+ZdEXXSLbSsuLwJjkCBWqRQUVA==", + "dev": true, "license": "MIT", "engines": { "node": ">=8.6" @@ -4361,6 +4375,7 @@ "version": "5.0.2", "resolved": "https://registry.npmjs.org/arg/-/arg-5.0.2.tgz", "integrity": "sha512-PYjyFOLKQ9y57JvQ6QLo8dAgNqswh8M1RMJYdQduT6xbWSgK36P/Z/v+p888pM69jMMfS8Xd8F6I1kQ/I9HUGg==", + "dev": true, "license": "MIT" }, "node_modules/argparse": { @@ -4732,6 +4747,7 @@ "version": "2.3.0", "resolved": "https://registry.npmjs.org/binary-extensions/-/binary-extensions-2.3.0.tgz", "integrity": "sha512-Ceh+7ox5qe7LJuLHoY0feh3pHuUDHAcRUeyL2VYghZwfpkNIy/+8Ocg0a3UuSoYzavmylwuLWQOf3hl0jjMMIw==", + "dev": true, "license": "MIT", "engines": { "node": ">=8" @@ -4757,6 +4773,7 @@ "version": "3.0.3", "resolved": "https://registry.npmjs.org/braces/-/braces-3.0.3.tgz", "integrity": "sha512-yQbXgO/OSZVD2IsiLlro+7Hf6Q18EJrKSEsdoMzKePKXct3gvD8oLcOQdIzGupr5Fj+EDe8gO/lxc1BzfMpxvA==", + "dev": true, "license": "MIT", "dependencies": { "fill-range": "^7.1.1" @@ -4872,6 +4889,7 @@ "version": "2.0.1", "resolved": "https://registry.npmjs.org/camelcase-css/-/camelcase-css-2.0.1.tgz", "integrity": "sha512-QOSvevhslijgYwRx6Rv7zKdMF8lbRmx+uQGx2+vDc+KI/eBnsy9kit5aj23AgGu3pa4t9AgwbnXWqS+iOY+2aA==", + "dev": true, "license": "MIT", "engines": { "node": ">= 6" @@ -4995,6 +5013,7 @@ "version": "3.6.0", "resolved": "https://registry.npmjs.org/chokidar/-/chokidar-3.6.0.tgz", "integrity": "sha512-7VT13fmjotKpGipCW9JEQAusEPE+Ei8nl6/g4FBAmIm0GOOLMua9NDDo/DWp0ZAxCr3cPq5ZpBqmPAQgDda2Pw==", + "dev": true, "license": "MIT", "dependencies": { "anymatch": "~3.1.2", @@ -5019,6 +5038,7 @@ "version": "5.1.2", "resolved": "https://registry.npmjs.org/glob-parent/-/glob-parent-5.1.2.tgz", "integrity": "sha512-AOIgSQCepiJYwP3ARnGx+5VnTu2HBYdzbGP45eLw1vr3zB3vZLeyed1sC9hnbcOc9/SrMyM5RPQrkGz4aS9Zow==", + "dev": true, "license": "ISC", "dependencies": { "is-glob": "^4.0.1" @@ -5094,6 +5114,7 @@ "version": "4.1.1", "resolved": "https://registry.npmjs.org/commander/-/commander-4.1.1.tgz", "integrity": "sha512-NOKm8xhkzAjzFx8B2v5OAHT+u5pRQc2UCa2Vq9jYL/31o2wi9mxBA7LIFs3sV5VSC49z6pEhfbMULvShKj26WA==", + "dev": true, "license": "MIT", "engines": { "node": ">= 6" @@ -5154,6 +5175,7 @@ "version": "3.0.0", "resolved": "https://registry.npmjs.org/cssesc/-/cssesc-3.0.0.tgz", "integrity": "sha512-/Tb/JcjK111nNScGob5MNtsntNM1aCNUDipB/TkwZFhyDrrE47SOx/18wF2bbjgc3ZzCSKW1T5nt5EbFoAz/Vg==", + "dev": true, "license": "MIT", "bin": { "cssesc": "bin/cssesc" @@ -5567,12 +5589,14 @@ "version": "1.2.2", "resolved": "https://registry.npmjs.org/didyoumean/-/didyoumean-1.2.2.tgz", "integrity": "sha512-gxtyfqMg7GKyhQmb056K7M3xszy/myH8w+B4RT+QXBQsvAOdc3XymqDDPHx1BgPgsdAA5SIifona89YtRATDzw==", + "dev": true, "license": "Apache-2.0" }, "node_modules/dlv": { "version": "1.1.3", "resolved": "https://registry.npmjs.org/dlv/-/dlv-1.1.3.tgz", "integrity": "sha512-+HlytyjlPKnIG8XuRG8WvmBP8xs8P71y+SKKS6ZXWoEgLuePxtDoUEiH7WkdePWrQ5JBpE6aoVqfZfJUQkjXwA==", + "dev": true, "license": "MIT" }, "node_modules/doctrine": { @@ -6486,6 +6510,7 @@ "version": "1.20.1", "resolved": "https://registry.npmjs.org/fastq/-/fastq-1.20.1.tgz", "integrity": "sha512-GGToxJ/w1x32s/D2EKND7kTil4n8OVk/9mycTc4VDza13lOvpUZTGX3mFSCtV9ksdGBVzvsyAVLM6mHFThxXxw==", + "dev": true, "license": "ISC", "dependencies": { "reusify": "^1.0.4" @@ -6518,6 +6543,7 @@ "version": "6.5.0", "resolved": "https://registry.npmjs.org/fdir/-/fdir-6.5.0.tgz", "integrity": "sha512-tIbYtZbucOs0BRGqPJkshJUYdL+SDH7dVM8gjy+ERp3WAUjLEFJE+02kanyHtwjWOnwrKYBiwAmM0p4kLJAnXg==", + "dev": true, "license": "MIT", "engines": { "node": ">=12.0.0" @@ -6555,6 +6581,7 @@ "version": "7.1.1", "resolved": "https://registry.npmjs.org/fill-range/-/fill-range-7.1.1.tgz", "integrity": "sha512-YsGpe3WHLK8ZYi4tWDg2Jy3ebRz2rXowDxnld4bkQB00cc/1Zw9AWnC0i9ztDJitivtQvaI9KaLyKrc+hBW0yg==", + "dev": true, "license": "MIT", "dependencies": { "to-regex-range": "^5.0.1" @@ -6715,6 +6742,7 @@ "version": "2.3.2", "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz", "integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==", + "dev": true, "hasInstallScript": true, "license": "MIT", "optional": true, @@ -6865,6 +6893,7 @@ "version": "6.0.2", "resolved": "https://registry.npmjs.org/glob-parent/-/glob-parent-6.0.2.tgz", "integrity": "sha512-XxwI8EOhVQgWp6iDL+3b0r86f4d6AX6zSU55HfB4ydCEuXLXc5FcYeOu+nnGftS4TEju/11rt4KJPTMgbfmv4A==", + "dev": true, "license": "ISC", "dependencies": { "is-glob": "^4.0.3" @@ -7362,6 +7391,7 @@ "version": "2.1.0", "resolved": "https://registry.npmjs.org/is-binary-path/-/is-binary-path-2.1.0.tgz", "integrity": "sha512-ZMERYes6pDydyuGidse7OsHxtbI7WVeUEozgR/g7rd0xUimYNlvZRE/K2MgZTjWy725IfelLeVcEM97mmtRGXw==", + "dev": true, "license": "MIT", "dependencies": { "binary-extensions": "^2.0.0" @@ -7414,6 +7444,7 @@ "version": "2.16.1", "resolved": "https://registry.npmjs.org/is-core-module/-/is-core-module-2.16.1.tgz", "integrity": "sha512-UfoeMA6fIJ8wTYFEUjelnaGI67v6+N7qXJEvQuIGa99l4xsCruSYOVSQ0uPANn4dAzm8lkYPaKLrrijLq7x23w==", + "dev": true, "license": "MIT", "dependencies": { "hasown": "^2.0.2" @@ -7474,6 +7505,7 @@ "version": "2.1.1", "resolved": "https://registry.npmjs.org/is-extglob/-/is-extglob-2.1.1.tgz", "integrity": "sha512-SbKbANkN603Vi4jEZv49LeVJMn4yGwsbzZworEoyEiutsN3nJYdbO36zfhGJ6QEDpOZIFkDtnq5JRxmvl3jsoQ==", + "dev": true, "license": "MIT", "engines": { "node": ">=0.10.0" @@ -7519,6 +7551,7 @@ "version": "4.0.3", "resolved": "https://registry.npmjs.org/is-glob/-/is-glob-4.0.3.tgz", "integrity": "sha512-xelSayHH36ZgE7ZWhli7pW34hNbNl8Ojv5KVmkJD4hBdD3th8Tfk9vYasLM+mXWOZhFkgZfxhLSnrwRr4elSSg==", + "dev": true, "license": "MIT", "dependencies": { "is-extglob": "^2.1.1" @@ -7567,6 +7600,7 @@ "version": "7.0.0", "resolved": "https://registry.npmjs.org/is-number/-/is-number-7.0.0.tgz", "integrity": "sha512-41Cifkg6e8TylSpdtTpeLVMqvSBEVzTttHvERD741+pnZ8ANv0004MRL43QKPDlK9cGvNp6NZWZUBlbGXYxxng==", + "dev": true, "license": "MIT", "engines": { "node": ">=0.12.0" @@ -7843,6 +7877,7 @@ "version": "1.21.7", "resolved": "https://registry.npmjs.org/jiti/-/jiti-1.21.7.tgz", "integrity": "sha512-/imKNG4EbWNrVjoNC/1H5/9GFy+tqjGBHCaSsN+P2RnPqjsLmv6UD3Ej+Kj8nBWaRAwyk7kK5ZUc+OEatnTR3A==", + "dev": true, "license": "MIT", "bin": { "jiti": "bin/jiti.js" @@ -8128,6 +8163,7 @@ "version": "3.1.3", "resolved": "https://registry.npmjs.org/lilconfig/-/lilconfig-3.1.3.tgz", "integrity": "sha512-/vlFKAoH5Cgt3Ie+JLhRbwOsCQePABiU3tJ1egGvyQ+33R/vcwM2Zl2QR/LzjsBeItPt3oSVXapn+m4nQDvpzw==", + "dev": true, "license": "MIT", "engines": { "node": ">=14" @@ -8140,6 +8176,7 @@ "version": "1.2.4", "resolved": "https://registry.npmjs.org/lines-and-columns/-/lines-and-columns-1.2.4.tgz", "integrity": "sha512-7ylylesZQ/PV29jhEDl3Ufjo6ZX7gCqJr5F7PKrqc93v7fzSymt1BpwEU8nAUXs8qzzvqhbjhK5QZg6Mt/HkBg==", + "dev": true, "license": "MIT" }, "node_modules/locate-path": { @@ -8454,6 +8491,7 @@ "version": "1.4.1", "resolved": "https://registry.npmjs.org/merge2/-/merge2-1.4.1.tgz", "integrity": "sha512-8q7VEgMJW4J8tcfVPy8g09NcQwZdbwFEqhe/WZkoIzjn/3TGDwtOCYtXGxA3O8tPzpczCCDgv+P2P5y00ZJOOg==", + "dev": true, "license": "MIT", "engines": { "node": ">= 8" @@ -8905,6 +8943,7 @@ "version": "4.0.8", "resolved": "https://registry.npmjs.org/micromatch/-/micromatch-4.0.8.tgz", "integrity": "sha512-PXwfBhYu0hBCPw8Dn0E+WDYb7af3dSLVWKi3HGv84IdF4TyFoC0ysxFd0Goxw7nSv4T/PzEJQxsYsEiFCKo2BA==", + "dev": true, "license": "MIT", "dependencies": { "braces": "^3.0.3", @@ -8918,6 +8957,7 @@ "version": "2.3.1", "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.1.tgz", "integrity": "sha512-JU3teHTNjmE2VCGFzuY8EXzCDVwEqB2a8fsIvwaStHhAWJEeVd1o1QD80CU6+ZdEXXSLbSsuLwJjkCBWqRQUVA==", + "dev": true, "license": "MIT", "engines": { "node": ">=8.6" @@ -9032,6 +9072,7 @@ "version": "2.7.0", "resolved": "https://registry.npmjs.org/mz/-/mz-2.7.0.tgz", "integrity": "sha512-z81GNO7nnYMEhrGh9LeymoE4+Yr0Wn5McHIZMK5cfQCl+NDX08sCZgUc9/6MHni9IWuFLm1Z3HTCXu2z9fN62Q==", + "dev": true, "license": "MIT", "dependencies": { "any-promise": "^1.0.0", @@ -9243,6 +9284,7 @@ "version": "3.0.0", "resolved": "https://registry.npmjs.org/normalize-path/-/normalize-path-3.0.0.tgz", "integrity": "sha512-6eZs5Ls3WtCisHWp9S2GUy8dqkpGi4BVSz3GaqiE6ezub0512ESztXUwUB6C6IKbQkY2Pnb/mD4WYojCRwcwLA==", + "dev": true, "license": "MIT", "engines": { "node": ">=0.10.0" @@ -9261,6 +9303,7 @@ "version": "3.0.0", "resolved": "https://registry.npmjs.org/object-hash/-/object-hash-3.0.0.tgz", "integrity": "sha512-RSn9F68PjH9HqtltsSnqYC1XXoWe9Bju5+213R98cNGttag9q9yAOTzdbsqvIa7aNm5WffBZFpWYr2aWrklWAw==", + "dev": true, "license": "MIT", "engines": { "node": ">= 6" @@ -9605,6 +9648,7 @@ "version": "1.0.7", "resolved": "https://registry.npmjs.org/path-parse/-/path-parse-1.0.7.tgz", "integrity": "sha512-LDJzPVEEEPR+y48z93A0Ed0yXb8pAByGWo/k5YYdYgpY2/2EsOsksJrq7lOHxryrVOn1ejG6oAp8ahvOIQD8sw==", + "dev": true, "license": "MIT" }, "node_modules/path-scurry": { @@ -9651,6 +9695,7 @@ "version": "4.0.3", "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-4.0.3.tgz", "integrity": "sha512-5gTmgEY/sqK6gFXLIsQNH19lWb4ebPDLA4SdLP7dsWkIXHWlG66oPuVvXSGFPppYZz8ZDZq0dYYrbHfBCVUb1Q==", + "dev": true, "license": "MIT", "engines": { "node": ">=12" @@ -9663,6 +9708,7 @@ "version": "2.3.0", "resolved": "https://registry.npmjs.org/pify/-/pify-2.3.0.tgz", "integrity": "sha512-udgsAY+fTnvv7kI7aaxbqwWNb0AHiB0qBO89PZKPkoTmGOgdbrHDKD+0B2X4uTfJ/FT1R09r9gTsjUjNJotuog==", + "dev": true, "license": "MIT", "engines": { "node": ">=0.10.0" @@ -9672,6 +9718,7 @@ "version": "4.0.7", "resolved": "https://registry.npmjs.org/pirates/-/pirates-4.0.7.tgz", "integrity": "sha512-TfySrs/5nm8fQJDcBDuUng3VOUKsd7S+zqvbOTiGXHfxX4wK31ard+hoNuvkicM/2YFzlpDgABOevKSsB4G/FA==", + "dev": true, "license": "MIT", "engines": { "node": ">= 6" @@ -9681,7 +9728,7 @@ "version": "1.58.1", "resolved": "https://registry.npmjs.org/playwright/-/playwright-1.58.1.tgz", "integrity": "sha512-+2uTZHxSCcxjvGc5C891LrS1/NlxglGxzrC4seZiVjcYVQfUa87wBL6rTDqzGjuoWNjnBzRqKmF6zRYGMvQUaQ==", - "devOptional": true, + "dev": true, "license": "Apache-2.0", "dependencies": { "playwright-core": "1.58.1" @@ -9700,7 +9747,7 @@ "version": "1.58.1", "resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.58.1.tgz", "integrity": "sha512-bcWzOaTxcW+VOOGBCQgnaKToLJ65d6AqfLVKEWvexyS3AS6rbXl+xdpYRMGSRBClPvyj44njOWoxjNdL/H9UNg==", - "devOptional": true, + "dev": true, "license": "Apache-2.0", "bin": { "playwright-core": "cli.js" @@ -9723,6 +9770,7 @@ "version": "8.5.6", "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.6.tgz", "integrity": "sha512-3Ybi1tAuwAP9s0r1UQ2J4n5Y0G05bJkpUIO0/bI9MhwmD70S5aTWbXGBwxHrelT+XM1k6dM0pk+SwNkpTRN7Pg==", + "dev": true, "funding": [ { "type": "opencollective", @@ -9751,6 +9799,7 @@ "version": "15.1.0", "resolved": "https://registry.npmjs.org/postcss-import/-/postcss-import-15.1.0.tgz", "integrity": "sha512-hpr+J05B2FVYUAXHeK1YyI267J/dDDhMU6B6civm8hSY1jYJnBXxzKDKDswzJmtLHryrjhnDjqqp/49t8FALew==", + "dev": true, "license": "MIT", "dependencies": { "postcss-value-parser": "^4.0.0", @@ -9768,6 +9817,7 @@ "version": "4.1.0", "resolved": "https://registry.npmjs.org/postcss-js/-/postcss-js-4.1.0.tgz", "integrity": "sha512-oIAOTqgIo7q2EOwbhb8UalYePMvYoIeRY2YKntdpFQXNosSu3vLrniGgmH9OKs/qAkfoj5oB3le/7mINW1LCfw==", + "dev": true, "funding": [ { "type": "opencollective", @@ -9793,6 +9843,7 @@ "version": "6.0.1", "resolved": "https://registry.npmjs.org/postcss-load-config/-/postcss-load-config-6.0.1.tgz", "integrity": "sha512-oPtTM4oerL+UXmx+93ytZVN82RrlY/wPUV8IeDxFrzIjXOLF1pN+EmKPLbubvKHT2HC20xXsCAH2Z+CKV6Oz/g==", + "dev": true, "funding": [ { "type": "opencollective", @@ -9835,6 +9886,7 @@ "version": "6.2.0", "resolved": "https://registry.npmjs.org/postcss-nested/-/postcss-nested-6.2.0.tgz", "integrity": "sha512-HQbt28KulC5AJzG+cZtj9kvKB93CFCdLvog1WFLf1D+xmMvPGlBstkpTEZfK5+AN9hfJocyBFCNiqyS48bpgzQ==", + "dev": true, "funding": [ { "type": "opencollective", @@ -9860,6 +9912,7 @@ "version": "6.1.2", "resolved": "https://registry.npmjs.org/postcss-selector-parser/-/postcss-selector-parser-6.1.2.tgz", "integrity": "sha512-Q8qQfPiZ+THO/3ZrOrO0cJJKfpYCagtMUkXbnEfmgUjwXg6z/WBeOyS9APBBPCTSiDV+s4SwQGu8yFsiMRIudg==", + "dev": true, "license": "MIT", "dependencies": { "cssesc": "^3.0.0", @@ -9873,6 +9926,7 @@ "version": "4.2.0", "resolved": "https://registry.npmjs.org/postcss-value-parser/-/postcss-value-parser-4.2.0.tgz", "integrity": "sha512-1NNCs6uurfkVbeXG4S8JFT9t19m45ICnif8zWLd5oPSZ50QnwMfK+H3jv408d4jw/7Bttv5axS5IiHoLaVNHeQ==", + "dev": true, "license": "MIT" }, "node_modules/prelude-ls": { @@ -9986,6 +10040,7 @@ "version": "1.2.3", "resolved": "https://registry.npmjs.org/queue-microtask/-/queue-microtask-1.2.3.tgz", "integrity": "sha512-NuaNSa6flKT5JaSYQzJok04JzTL1CA6aGhv5rfLW3PgqA+M2ChpZQnAC8h8i4ZFkBS8X5RqkDBHA7r4hej3K9A==", + "dev": true, "funding": [ { "type": "github", @@ -10774,6 +10829,7 @@ "version": "1.0.0", "resolved": "https://registry.npmjs.org/read-cache/-/read-cache-1.0.0.tgz", "integrity": "sha512-Owdv/Ft7IjOgm/i0xvNDZ1LrRANRfew4b2prF3OWMQLxLfu3bS8FVhCsrSCMK4lR56Y9ya+AThoTpDCTxCmpRA==", + "dev": true, "license": "MIT", "dependencies": { "pify": "^2.3.0" @@ -10783,6 +10839,7 @@ "version": "3.6.0", "resolved": "https://registry.npmjs.org/readdirp/-/readdirp-3.6.0.tgz", "integrity": "sha512-hOS089on8RduqdbhvQ5Z37A0ESjsqz6qnRcffsMU3495FuTdqSm+7bhJ29JvIOsBDEEnan5DPu9t3To9VRlMzA==", + "dev": true, "license": "MIT", "dependencies": { "picomatch": "^2.2.1" @@ -10795,6 +10852,7 @@ "version": "2.3.1", "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.1.tgz", "integrity": "sha512-JU3teHTNjmE2VCGFzuY8EXzCDVwEqB2a8fsIvwaStHhAWJEeVd1o1QD80CU6+ZdEXXSLbSsuLwJjkCBWqRQUVA==", + "dev": true, "license": "MIT", "engines": { "node": ">=8.6" @@ -11059,6 +11117,7 @@ "version": "1.22.11", "resolved": "https://registry.npmjs.org/resolve/-/resolve-1.22.11.tgz", "integrity": "sha512-RfqAvLnMl313r7c9oclB1HhUEAezcpLjz95wFH4LVuhk9JF/r22qmVP9AMmOU4vMX7Q8pN8jwNg/CSpdFnMjTQ==", + "dev": true, "license": "MIT", "dependencies": { "is-core-module": "^2.16.1", @@ -11099,6 +11158,7 @@ "version": "1.1.0", "resolved": "https://registry.npmjs.org/reusify/-/reusify-1.1.0.tgz", "integrity": "sha512-g6QUff04oZpHs0eG5p83rFLhHeV00ug/Yf9nZM6fLeUrPguBTkTQOdpAWWspMh55TZfVQDPaN3NQJfbVRAxdIw==", + "dev": true, "license": "MIT", "engines": { "iojs": ">=1.0.0", @@ -11154,6 +11214,7 @@ "version": "1.2.0", "resolved": "https://registry.npmjs.org/run-parallel/-/run-parallel-1.2.0.tgz", "integrity": "sha512-5l4VyZR86LZ/lDxZTR6jqL8AFE2S0IFLMP26AbjsLVADxHdhB/c0GUsH+y39UfCi3dzz8OlQuPmnaJOMoDHQBA==", + "dev": true, "funding": [ { "type": "github", @@ -11794,6 +11855,7 @@ "version": "3.35.1", "resolved": "https://registry.npmjs.org/sucrase/-/sucrase-3.35.1.tgz", "integrity": "sha512-DhuTmvZWux4H1UOnWMB3sk0sbaCVOoQZjv8u1rDoTV0HTdGem9hkAZtl4JZy8P2z4Bg0nT+YMeOFyVr4zcG5Tw==", + "dev": true, "license": "MIT", "dependencies": { "@jridgewell/gen-mapping": "^0.3.2", @@ -11829,6 +11891,7 @@ "version": "1.0.0", "resolved": "https://registry.npmjs.org/supports-preserve-symlinks-flag/-/supports-preserve-symlinks-flag-1.0.0.tgz", "integrity": "sha512-ot0WnXS9fgdkgIcePe6RHNk1WA8+muPa6cSjeR3V8K27q9BB1rTE3R1p7Hv0z1ZyAc8s6Vvv8DIyWf681MAt0w==", + "dev": true, "license": "MIT", "engines": { "node": ">= 0.4" @@ -11864,6 +11927,7 @@ "version": "3.4.19", "resolved": "https://registry.npmjs.org/tailwindcss/-/tailwindcss-3.4.19.tgz", "integrity": "sha512-3ofp+LL8E+pK/JuPLPggVAIaEuhvIz4qNcf3nA1Xn2o/7fb7s/TYpHhwGDv1ZU3PkBluUVaF8PyCHcm48cKLWQ==", + "dev": true, "license": "MIT", "dependencies": { "@alloc/quick-lru": "^5.2.0", @@ -11901,6 +11965,7 @@ "version": "3.3.3", "resolved": "https://registry.npmjs.org/fast-glob/-/fast-glob-3.3.3.tgz", "integrity": "sha512-7MptL8U0cqcFdzIzwOTHoilX9x5BrNqye7Z/LuC7kCMRio1EMSyqRK3BEAUD7sXRq4iT4AzTVuZdhgQ2TCvYLg==", + "dev": true, "license": "MIT", "dependencies": { "@nodelib/fs.stat": "^2.0.2", @@ -11917,6 +11982,7 @@ "version": "5.1.2", "resolved": "https://registry.npmjs.org/glob-parent/-/glob-parent-5.1.2.tgz", "integrity": "sha512-AOIgSQCepiJYwP3ARnGx+5VnTu2HBYdzbGP45eLw1vr3zB3vZLeyed1sC9hnbcOc9/SrMyM5RPQrkGz4aS9Zow==", + "dev": true, "license": "ISC", "dependencies": { "is-glob": "^4.0.1" @@ -11944,6 +12010,7 @@ "version": "3.3.1", "resolved": "https://registry.npmjs.org/thenify/-/thenify-3.3.1.tgz", "integrity": "sha512-RVZSIV5IG10Hk3enotrhvz0T9em6cyHBLkH/YAZuKqd8hRkKhSfCGIcP2KUY0EPxndzANBmNllzWPwak+bheSw==", + "dev": true, "license": "MIT", "dependencies": { "any-promise": "^1.0.0" @@ -11953,6 +12020,7 @@ "version": "1.6.0", "resolved": "https://registry.npmjs.org/thenify-all/-/thenify-all-1.6.0.tgz", "integrity": "sha512-RNxQH/qI8/t3thXJDwcstUO4zeqo64+Uy/+sNVRBx4Xn2OX+OZ9oP+iJnNFqplFra2ZUVeKCSa2oVWi3T4uVmA==", + "dev": true, "license": "MIT", "dependencies": { "thenify": ">= 3.1.0 < 4" @@ -11994,6 +12062,7 @@ "version": "0.2.15", "resolved": "https://registry.npmjs.org/tinyglobby/-/tinyglobby-0.2.15.tgz", "integrity": "sha512-j2Zq4NyQYG5XMST4cbs02Ak8iJUdxRM0XI5QyxXuZOzKOINmWurp3smXu3y5wDcJrptwpSjgXHzIQxR0omXljQ==", + "dev": true, "license": "MIT", "dependencies": { "fdir": "^6.5.0", @@ -12060,6 +12129,7 @@ "version": "5.0.1", "resolved": "https://registry.npmjs.org/to-regex-range/-/to-regex-range-5.0.1.tgz", "integrity": "sha512-65P7iz6X5yEr1cwcgvQxbbIw7Uk3gOy5dIdtZ4rDveLqhrdJP+Li/Hx6tyK0NEb+2GCyneCMJiGqrADCSNk8sQ==", + "dev": true, "license": "MIT", "dependencies": { "is-number": "^7.0.0" @@ -12147,6 +12217,7 @@ "version": "0.1.13", "resolved": "https://registry.npmjs.org/ts-interface-checker/-/ts-interface-checker-0.1.13.tgz", "integrity": "sha512-Y/arvbn+rrz3JCKl9C4kVNfTfSm2/mEp5FSz5EsZSANGPSlQrpRI5M4PKF+mJnE52jOO90PnPSc3Ur3bTQw0gA==", + "dev": true, "license": "Apache-2.0" }, "node_modules/tsconfig-paths": { @@ -12263,7 +12334,7 @@ "version": "5.3.3", "resolved": "https://registry.npmjs.org/typescript/-/typescript-5.3.3.tgz", "integrity": "sha512-pXWcraxM0uxAS+tN0AG/BF2TyqmHO014Z070UsJ+pFvYuRSq8KH8DmWpnbXe0pEPDHXZV3FcAbJkijJ5oNEnWw==", - "devOptional": true, + "dev": true, "license": "Apache-2.0", "bin": { "tsc": "bin/tsc", @@ -12465,6 +12536,7 @@ "version": "1.0.2", "resolved": "https://registry.npmjs.org/util-deprecate/-/util-deprecate-1.0.2.tgz", "integrity": "sha512-EPD5q1uXyFxJpCrLnCc1nHnq3gOa6DZBocAIiI2TaSCA7VCJ1UJDMagCzIkXNsUYfD1daK//LTEQ8xiIbrHtcw==", + "dev": true, "license": "MIT" }, "node_modules/uuid": { @@ -12918,7 +12990,7 @@ "version": "8.19.0", "resolved": "https://registry.npmjs.org/ws/-/ws-8.19.0.tgz", "integrity": "sha512-blAT2mjOEIi0ZzruJfIhb3nps74PRWTCz1IjglWEEpQl5XS/UNama6u2/rjFkDDouqr4L67ry+1aGIALViWjDg==", - "devOptional": true, + "dev": true, "license": "MIT", "engines": { "node": ">=10.0.0" @@ -12975,17 +13047,6 @@ "url": "https://github.com/sponsors/sindresorhus" } }, - "node_modules/zod": { - "version": "3.25.76", - "resolved": "https://registry.npmjs.org/zod/-/zod-3.25.76.tgz", - "integrity": "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ==", - "license": "MIT", - "optional": true, - "peer": true, - "funding": { - "url": "https://github.com/sponsors/colinhacks" - } - }, "node_modules/zwitch": { "version": "2.0.4", "resolved": "https://registry.npmjs.org/zwitch/-/zwitch-2.0.4.tgz", @@ -12995,6 +13056,21 @@ "type": "github", "url": "https://github.com/sponsors/wooorm" } + }, + "node_modules/@next/swc-win32-ia32-msvc": { + "version": "14.2.33", + "resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.33.tgz", + "integrity": "sha512-pc9LpGNKhJ0dXQhZ5QMmYxtARwwmWLpeocFmVG5Z0DzWq5Uf0izcI8tLc+qOpqxO1PWqZ5A7J1blrUIKrIFc7Q==", + "cpu": [ + "ia32" + ], + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } } } } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.test.tsx index 1e8eabaea2e..b0df37ad6d8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.test.tsx @@ -1,6 +1,6 @@ /* @vitest-environment jsdom */ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render } from "@testing-library/react"; +import { act, render } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import ModelsAndEndpointsView from "./ModelsAndEndpointsView"; @@ -13,6 +13,8 @@ vi.mock("@/components/networking", () => ({ getCallbacksCall: vi.fn().mockResolvedValue({ router_settings: {} }), setCallbacksCall: vi.fn().mockResolvedValue(undefined), getUiSettings: vi.fn().mockResolvedValue({ values: {} }), + latestHealthChecksCall: vi.fn().mockResolvedValue({ latest_health_checks: {} }), + getModelCostMapReloadStatus: vi.fn().mockResolvedValue({}), })); vi.mock("@/app/(dashboard)/models-and-endpoints/components/ModelAnalyticsTab/ModelAnalyticsTab", () => ({ @@ -27,6 +29,14 @@ vi.mock("@/components/add_model/AddModelForm", () => ({ default: () => null, })); +const mockHealthCheckComponent = vi.fn((_props: { all_models_on_proxy?: string[] }) => null); +vi.mock("@/components/model_dashboard/HealthCheckComponent", () => ({ + default: (props: { all_models_on_proxy?: string[] }) => { + mockHealthCheckComponent(props); + return null; + }, +})); + vi.mock("@/app/(dashboard)/hooks/useTeams", () => ({ default: () => ({ teams: [], @@ -104,4 +114,43 @@ describe("ModelsAndEndpointsView", () => { ); expect(await findByText("Model Management", {}, { timeout: 10000 })).toBeInTheDocument(); }, 15000); + + it("should pass model IDs (not model names) to HealthCheckComponent as all_models_on_proxy", async () => { + mockHealthCheckComponent.mockClear(); + const modelDataWithIds = { + data: [ + { model_name: "gpt-4", model_info: { id: "deployment-id-1" } }, + { model_name: "gpt-4", model_info: { id: "deployment-id-2" } }, + ], + }; + mockUseModelsInfo.mockReturnValue({ + data: { data: modelDataWithIds.data }, + isLoading: false, + refetch: vi.fn(), + }); + + const queryClient = createQueryClient(); + const { getByRole } = render( + + {}} + premiumUser={false} + teams={[]} + /> + , + ); + + const healthStatusTab = getByRole("tab", { name: "Health Status" }); + await act(async () => { + healthStatusTab.click(); + }); + + expect(mockHealthCheckComponent).toHaveBeenCalled(); + const healthCheckProps = mockHealthCheckComponent.mock.calls[0][0]; + expect(healthCheckProps.all_models_on_proxy).toEqual(["deployment-id-1", "deployment-id-2"]); + expect(healthCheckProps.all_models_on_proxy).not.toContain("gpt-4"); + }); }); 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 9d77774cb4c..b697a859dc5 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 @@ -98,6 +98,13 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te return modelDataResponse.data.map((model: any) => model.model_name); }, [modelDataResponse?.data]); + const allModelIdsOnProxy = useMemo(() => { + if (!modelDataResponse?.data) return []; + return modelDataResponse.data + .map((model: any) => model.model_info?.id) + .filter((id: string | undefined): id is string => Boolean(id)); + }, [modelDataResponse?.data]); + const getProviderFromModel = (model: string) => { if (modelCostMapData !== null && modelCostMapData !== undefined) { if (typeof modelCostMapData == "object" && model in modelCostMapData) { @@ -397,7 +404,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te ({ + useMultiCostEstimate: vi.fn(() => ({ + debouncedFetchForEntry: vi.fn(), + removeEntry: vi.fn(), + getMultiModelResult: vi.fn((entries: ModelEntry[]): MultiModelResult => ({ + entries: entries.map((e) => ({ entry: e, result: null, loading: false, error: null })), + totals: { + cost_per_request: 0, + daily_cost: null, + monthly_cost: null, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + })), + })), +})); + +vi.mock("./multi_export_utils", () => ({ + exportMultiToPDF: vi.fn(), + exportMultiToCSV: vi.fn(), +})); + +vi.mock("@/utils/dataUtils", () => ({ + formatNumberWithCommas: vi.fn((v: number, d: number = 0) => + Number.isFinite(v) ? v.toFixed(d) : "-" + ), +})); + +const DEFAULT_PROPS = { + accessToken: "test-token", + models: ["gpt-4", "gpt-3.5-turbo", "claude-3-sonnet"], +}; + +describe("PricingCalculator", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render the calculator with an initial model row", () => { + renderWithProviders(); + expect(screen.getByRole("table")).toBeInTheDocument(); + }); + + it("should render the time period toggle with Per Day and Per Month options", () => { + renderWithProviders(); + expect(screen.getByText("Per Day")).toBeInTheDocument(); + expect(screen.getByText("Per Month")).toBeInTheDocument(); + }); + + it("should render an Add Another Model button", () => { + renderWithProviders(); + expect(screen.getByRole("button", { name: /add another model/i })).toBeInTheDocument(); + }); + + it("should show the Requests/Month column header by default", () => { + renderWithProviders(); + expect(screen.getByText("Requests/Month")).toBeInTheDocument(); + }); + + it("should add a new row when Add Another Model is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const table = screen.getByRole("table"); + const initialRows = within(table).getAllByRole("row"); + + await user.click(screen.getByRole("button", { name: /add another model/i })); + + const updatedRows = within(table).getAllByRole("row"); + // One new data row added (header row + data rows) + expect(updatedRows.length).toBeGreaterThan(initialRows.length); + }); + + it("should have the delete button disabled when there is only one row", () => { + renderWithProviders(); + const allButtons = screen.getAllByRole("button"); + const disabledButtons = allButtons.filter((btn) => btn.hasAttribute("disabled")); + expect(disabledButtons.length).toBeGreaterThan(0); + }); + + it("should have no disabled buttons after adding a second row", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /add another model/i })); + + // With two rows, no delete buttons should be disabled + const allButtons = screen.getAllByRole("button"); + const disabledButtons = allButtons.filter((btn) => btn.hasAttribute("disabled")); + expect(disabledButtons.length).toBe(0); + }); + + describe("time period toggle", () => { + it("should switch the column header to Requests/Day when Per Day is selected", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByText("Per Day")); + + expect(screen.getByText("Requests/Day")).toBeInTheDocument(); + }); + + it("should switch the column header back to Requests/Month when Per Month is selected", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByText("Per Day")); + expect(screen.getByText("Requests/Day")).toBeInTheDocument(); + + await user.click(screen.getByText("Per Month")); + expect(screen.getByText("Requests/Month")).toBeInTheDocument(); + }); + }); + + it("should render column headers for Model, Input Tokens, and Output Tokens", () => { + renderWithProviders(); + expect(screen.getByText("Model")).toBeInTheDocument(); + expect(screen.getByText("Input Tokens")).toBeInTheDocument(); + expect(screen.getByText("Output Tokens")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_cost_results.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_cost_results.test.tsx new file mode 100644 index 00000000000..6f6522f395e --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_cost_results.test.tsx @@ -0,0 +1,305 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../../tests/test-utils"; +import MultiCostResults from "./multi_cost_results"; +import type { MultiModelResult } from "./types"; +import type { CostEstimateResponse } from "../types"; + +vi.mock("./multi_export_utils", () => ({ + exportMultiToPDF: vi.fn(), + exportMultiToCSV: vi.fn(), +})); + +vi.mock("@/utils/dataUtils", () => ({ + formatNumberWithCommas: vi.fn((v: number, d: number = 0) => + Number.isFinite(v) ? v.toFixed(d) : "-" + ), +})); + +function makeCostResponse(overrides: Partial = {}): CostEstimateResponse { + return { + model: "gpt-4", + input_tokens: 1000, + output_tokens: 500, + num_requests_per_day: 100, + num_requests_per_month: null, + cost_per_request: 0.05, + input_cost_per_request: 0.03, + output_cost_per_request: 0.02, + margin_cost_per_request: 0, + daily_cost: 5.0, + daily_input_cost: 3.0, + daily_output_cost: 2.0, + daily_margin_cost: 0, + monthly_cost: null, + monthly_input_cost: null, + monthly_output_cost: null, + monthly_margin_cost: null, + input_cost_per_token: null, + output_cost_per_token: null, + provider: "openai", + ...overrides, + }; +} + +function makeMultiResult(overrides: Partial = {}): MultiModelResult { + return { + entries: [ + { + entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, + result: makeCostResponse(), + loading: false, + error: null, + }, + ], + totals: { + cost_per_request: 0.05, + daily_cost: 5.0, + monthly_cost: null, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + ...overrides, + }; +} + +function emptyMultiResult(): MultiModelResult { + return { + entries: [ + { + entry: { id: "e1", model: "", input_tokens: 1000, output_tokens: 500 }, + result: null, + loading: false, + error: null, + }, + ], + totals: { + cost_per_request: 0, + daily_cost: null, + monthly_cost: null, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + }; +} + +describe("MultiCostResults", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + describe("when no model has been selected", () => { + it("should show a prompt to select models", () => { + renderWithProviders( + + ); + expect(screen.getByText(/select models above to see cost estimates/i)).toBeInTheDocument(); + }); + }); + + describe("when results are loading and no data has arrived yet", () => { + it("should show a calculating costs spinner", () => { + const multiResult: MultiModelResult = { + entries: [ + { + entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, + result: null, + loading: true, + error: null, + }, + ], + totals: { + cost_per_request: 0, + daily_cost: null, + monthly_cost: null, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + }; + + renderWithProviders(); + expect(screen.getByText(/calculating costs/i)).toBeInTheDocument(); + }); + }); + + describe("when there are errors but no valid results", () => { + it("should display the error message with the model name", () => { + const multiResult: MultiModelResult = { + entries: [ + { + entry: { id: "e1", model: "bad-model", input_tokens: 0, output_tokens: 0 }, + result: null, + loading: false, + error: "Pricing not found", + }, + ], + totals: { + cost_per_request: 0, + daily_cost: null, + monthly_cost: null, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + }; + + renderWithProviders(); + expect(screen.getByText(/bad-model/i)).toBeInTheDocument(); + expect(screen.getByText(/Pricing not found/i)).toBeInTheDocument(); + }); + }); + + describe("when valid results are available", () => { + it("should show the Cost Estimates heading", () => { + renderWithProviders( + + ); + expect(screen.getByText("Cost Estimates")).toBeInTheDocument(); + }); + + it("should display the Total Per Request statistic", () => { + renderWithProviders( + + ); + expect(screen.getByText("Total Per Request")).toBeInTheDocument(); + }); + + it("should display Total Daily statistic when timePeriod is day", () => { + renderWithProviders( + + ); + expect(screen.getByText("Total Daily")).toBeInTheDocument(); + }); + + it("should display Total Monthly statistic when timePeriod is month", () => { + renderWithProviders( + + ); + expect(screen.getByText("Total Monthly")).toBeInTheDocument(); + }); + + it("should show the model name in the summary table", () => { + renderWithProviders( + + ); + expect(screen.getByText("gpt-4")).toBeInTheDocument(); + }); + + it("should show the provider tag next to the model name", () => { + renderWithProviders( + + ); + expect(screen.getByText("openai")).toBeInTheDocument(); + }); + + it("should show the Export button when results are available", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /export/i })).toBeInTheDocument(); + }); + + it("should expand the model breakdown row when the expand button is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + // The expand column renders a button (RightOutlined icon) for rows without errors + const expandButtons = screen.getAllByRole("button"); + // Find the small expand button (not the Export button) + const expandButton = expandButtons.find( + (btn) => !btn.textContent?.toLowerCase().includes("export") + ); + expect(expandButton).toBeDefined(); + + await user.click(expandButton!); + + // After expanding, the SingleModelBreakdown should be visible + expect(screen.getByText("Total/Request")).toBeInTheDocument(); + }); + + it("should show the collapse icon after expanding a row", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + const getExpandButton = () => { + const allButtons = screen.getAllByRole("button"); + return allButtons.find((btn) => !btn.textContent?.toLowerCase().includes("export")); + }; + + // Before expand: button has the "down" aria-label (RightOutlined renders as down in ant icons) + // Just verify clicking works and the breakdown content appears + await user.click(getExpandButton()!); + expect(screen.getByText("Total/Request")).toBeInTheDocument(); + + // After a second click, the row collapses — content may be hidden or removed + await user.click(getExpandButton()!); + // The expanded content should no longer be visible + expect(screen.queryByText("Total/Request")).not.toBeVisible(); + }); + }); + + describe("margin section", () => { + it("should show margin fee details when margin per request is greater than zero", () => { + const multiResult = makeMultiResult({ + entries: [ + { + entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, + result: makeCostResponse({ margin_cost_per_request: 0.01, daily_margin_cost: 1.0 }), + loading: false, + error: null, + }, + ], + totals: { + cost_per_request: 0.06, + daily_cost: 6.0, + monthly_cost: null, + margin_per_request: 0.01, + daily_margin: 1.0, + monthly_margin: null, + }, + }); + + renderWithProviders(); + expect(screen.getByText("Margin Fee/Request")).toBeInTheDocument(); + }); + + it("should not show margin fee details when margin per request is zero", () => { + renderWithProviders( + + ); + expect(screen.queryByText("Margin Fee/Request")).not.toBeInTheDocument(); + }); + }); + + describe("when a model has zero cost", () => { + it("should show a warning about missing pricing data", () => { + const multiResult = makeMultiResult({ + entries: [ + { + entry: { id: "e1", model: "custom-model", input_tokens: 1000, output_tokens: 500 }, + result: makeCostResponse({ model: "custom-model", cost_per_request: 0 }), + loading: false, + error: null, + }, + ], + }); + + renderWithProviders(); + expect(screen.getByText(/no pricing data found/i)).toBeInTheDocument(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_export_dropdown.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_export_dropdown.test.tsx new file mode 100644 index 00000000000..be1a89cf77f --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_export_dropdown.test.tsx @@ -0,0 +1,146 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { screen, fireEvent } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../../tests/test-utils"; +import MultiExportDropdown from "./multi_export_dropdown"; +import type { MultiModelResult } from "./types"; + +vi.mock("./multi_export_utils", () => ({ + exportMultiToPDF: vi.fn(), + exportMultiToCSV: vi.fn(), +})); + +import { exportMultiToPDF, exportMultiToCSV } from "./multi_export_utils"; + +function makeMultiResult(hasResult: boolean): MultiModelResult { + return { + entries: [ + { + entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, + result: hasResult + ? { + model: "gpt-4", + input_tokens: 1000, + output_tokens: 500, + num_requests_per_day: null, + num_requests_per_month: null, + cost_per_request: 0.05, + input_cost_per_request: 0.03, + output_cost_per_request: 0.02, + margin_cost_per_request: 0, + daily_cost: null, + daily_input_cost: null, + daily_output_cost: null, + daily_margin_cost: null, + monthly_cost: null, + monthly_input_cost: null, + monthly_output_cost: null, + monthly_margin_cost: null, + input_cost_per_token: null, + output_cost_per_token: null, + provider: "openai", + } + : null, + loading: false, + error: null, + }, + ], + totals: { + cost_per_request: hasResult ? 0.05 : 0, + daily_cost: null, + monthly_cost: null, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + }; +} + +describe("MultiExportDropdown", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should not render anything when no entries have results", () => { + const { container } = renderWithProviders( + + ); + expect(container.firstChild).toBeNull(); + }); + + it("should render the Export button when at least one entry has a result", () => { + renderWithProviders(); + expect(screen.getByRole("button", { name: /^export$/i })).toBeInTheDocument(); + }); + + it("should show the export menu when the Export button is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + + expect(screen.getByText("Export as PDF")).toBeInTheDocument(); + expect(screen.getByText("Export as CSV")).toBeInTheDocument(); + }); + + it("should hide the export menu when the Export button is clicked again", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + expect(screen.getByText("Export as PDF")).toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + expect(screen.queryByText("Export as PDF")).not.toBeInTheDocument(); + }); + + it("should call exportMultiToPDF and close the menu when Export as PDF is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + await user.click(screen.getByText("Export as PDF")); + + expect(exportMultiToPDF).toHaveBeenCalledTimes(1); + expect(screen.queryByText("Export as PDF")).not.toBeInTheDocument(); + }); + + it("should call exportMultiToCSV and close the menu when Export as CSV is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + await user.click(screen.getByText("Export as CSV")); + + expect(exportMultiToCSV).toHaveBeenCalledTimes(1); + expect(screen.queryByText("Export as CSV")).not.toBeInTheDocument(); + }); + + it("should pass the multiResult to the export functions", async () => { + const user = userEvent.setup(); + const multiResult = makeMultiResult(true); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + await user.click(screen.getByText("Export as PDF")); + + expect(exportMultiToPDF).toHaveBeenCalledWith(multiResult); + }); + + it("should close the menu when clicking outside", async () => { + const user = userEvent.setup(); + renderWithProviders( +
+ +
Outside
+
+ ); + + await user.click(screen.getByRole("button", { name: /^export$/i })); + expect(screen.getByText("Export as PDF")).toBeInTheDocument(); + + fireEvent.mouseDown(screen.getByTestId("outside")); + expect(screen.queryByText("Export as PDF")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_export_utils.test.ts b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_export_utils.test.ts new file mode 100644 index 00000000000..40d6879d1af --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/multi_export_utils.test.ts @@ -0,0 +1,274 @@ +import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { exportMultiToPDF, exportMultiToCSV } from "./multi_export_utils"; +import type { MultiModelResult } from "./types"; +import type { CostEstimateResponse } from "../types"; + +vi.mock("@/utils/dataUtils", () => ({ + formatNumberWithCommas: vi.fn((v: number, d: number = 0) => + Number.isFinite(v) ? v.toFixed(d) : "-" + ), +})); + +function makeCostResponse(overrides: Partial = {}): CostEstimateResponse { + return { + model: "gpt-4", + input_tokens: 1000, + output_tokens: 500, + num_requests_per_day: 100, + num_requests_per_month: 3000, + cost_per_request: 0.05, + input_cost_per_request: 0.03, + output_cost_per_request: 0.02, + margin_cost_per_request: 0, + daily_cost: 5.0, + daily_input_cost: 3.0, + daily_output_cost: 2.0, + daily_margin_cost: 0, + monthly_cost: 150.0, + monthly_input_cost: 90.0, + monthly_output_cost: 60.0, + monthly_margin_cost: 0, + input_cost_per_token: 0.00003, + output_cost_per_token: 0.00004, + provider: "openai", + ...overrides, + }; +} + +function makeMultiResult(overrides: Partial = {}): MultiModelResult { + return { + entries: [ + { + entry: { id: "entry-1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, + result: makeCostResponse(), + loading: false, + error: null, + }, + ], + totals: { + cost_per_request: 0.05, + daily_cost: 5.0, + monthly_cost: 150.0, + margin_per_request: 0, + daily_margin: null, + monthly_margin: null, + }, + ...overrides, + }; +} + +describe("exportMultiToPDF", () => { + let mockPrintWindow: { + document: { write: ReturnType; close: ReturnType }; + print: ReturnType; + onload: (() => void) | null; + }; + + beforeEach(() => { + mockPrintWindow = { + document: { write: vi.fn(), close: vi.fn() }, + print: vi.fn(), + onload: null, + }; + vi.spyOn(window, "open").mockReturnValue(mockPrintWindow as unknown as Window); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("should open a new popup window", () => { + exportMultiToPDF(makeMultiResult()); + expect(window.open).toHaveBeenCalledWith("", "_blank"); + }); + + it("should write HTML containing the report title", () => { + exportMultiToPDF(makeMultiResult()); + const html = mockPrintWindow.document.write.mock.calls[0][0] as string; + expect(html).toContain("LLM Cost Estimate Report"); + }); + + it("should include model name and provider in the generated HTML", () => { + exportMultiToPDF(makeMultiResult()); + const html = mockPrintWindow.document.write.mock.calls[0][0] as string; + expect(html).toContain("gpt-4"); + expect(html).toContain("openai"); + }); + + it("should close the document after writing", () => { + exportMultiToPDF(makeMultiResult()); + expect(mockPrintWindow.document.close).toHaveBeenCalledTimes(1); + }); + + it("should call print after the window finishes loading", () => { + exportMultiToPDF(makeMultiResult()); + expect(mockPrintWindow.print).not.toHaveBeenCalled(); + mockPrintWindow.onload!(); + expect(mockPrintWindow.print).toHaveBeenCalledTimes(1); + }); + + it("should show the margin section when margin per request is greater than zero", () => { + const multiResult = makeMultiResult({ + totals: { + cost_per_request: 0.06, + daily_cost: 5.0, + monthly_cost: 150.0, + margin_per_request: 0.01, + daily_margin: 1.0, + monthly_margin: 30.0, + }, + }); + exportMultiToPDF(multiResult); + const html = mockPrintWindow.document.write.mock.calls[0][0] as string; + expect(html).toContain("Margin/Request"); + }); + + it("should not show the margin section when margin per request is zero", () => { + exportMultiToPDF(makeMultiResult()); + const html = mockPrintWindow.document.write.mock.calls[0][0] as string; + expect(html).not.toContain("Margin/Request"); + }); + + it("should alert when popup is blocked", () => { + vi.spyOn(window, "open").mockReturnValue(null); + const alertSpy = vi.spyOn(window, "alert").mockImplementation(() => {}); + exportMultiToPDF(makeMultiResult()); + expect(alertSpy).toHaveBeenCalledWith("Please allow popups to export PDF"); + }); + + it("should only include entries that have a result", () => { + const multiResult: MultiModelResult = { + entries: [ + { entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, result: null, loading: false, error: null }, + { entry: { id: "e2", model: "claude-3", input_tokens: 500, output_tokens: 250 }, result: makeCostResponse({ model: "claude-3", provider: "anthropic" }), loading: false, error: null }, + ], + totals: { cost_per_request: 0.05, daily_cost: 5.0, monthly_cost: 150.0, margin_per_request: 0, daily_margin: null, monthly_margin: null }, + }; + exportMultiToPDF(multiResult); + const html = mockPrintWindow.document.write.mock.calls[0][0] as string; + expect(html).toContain("1 model configured"); + expect(html).toContain("claude-3"); + }); + + it("should show plural 'models' when multiple results are present", () => { + const multiResult: MultiModelResult = { + entries: [ + { entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, result: makeCostResponse(), loading: false, error: null }, + { entry: { id: "e2", model: "claude-3", input_tokens: 500, output_tokens: 250 }, result: makeCostResponse({ model: "claude-3" }), loading: false, error: null }, + ], + totals: { cost_per_request: 0.10, daily_cost: 10.0, monthly_cost: 300.0, margin_per_request: 0, daily_margin: null, monthly_margin: null }, + }; + exportMultiToPDF(multiResult); + const html = mockPrintWindow.document.write.mock.calls[0][0] as string; + expect(html).toContain("2 models configured"); + }); +}); + +describe("exportMultiToCSV", () => { + beforeEach(() => { + document.body.innerHTML = ""; + window.URL.createObjectURL = vi.fn(() => "blob:mock-url"); + window.URL.revokeObjectURL = vi.fn(); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("should create an object URL and revoke it after download", () => { + exportMultiToCSV(makeMultiResult()); + expect(window.URL.createObjectURL).toHaveBeenCalledTimes(1); + expect(window.URL.revokeObjectURL).toHaveBeenCalledWith("blob:mock-url"); + }); + + it("should set the download filename to include today's date", () => { + const createdAnchors: HTMLAnchorElement[] = []; + const originalCreate = document.createElement.bind(document); + vi.spyOn(document, "createElement").mockImplementation((tag: string) => { + const el = originalCreate(tag); + if (tag === "a") createdAnchors.push(el as HTMLAnchorElement); + return el; + }); + + const today = new Date().toISOString().split("T")[0]; + exportMultiToCSV(makeMultiResult()); + + expect(createdAnchors[0].download).toBe(`cost_estimate_multi_model_${today}.csv`); + }); + + it("should generate CSV content containing a header row and model data", () => { + let csvContent = ""; + const OriginalBlob = globalThis.Blob; + globalThis.Blob = class extends OriginalBlob { + constructor(parts?: BlobPart[], options?: BlobPropertyBag) { + super(parts, options); + if (typeof parts?.[0] === "string") csvContent = parts[0]; + } + } as unknown as typeof Blob; + + exportMultiToCSV(makeMultiResult()); + globalThis.Blob = OriginalBlob; + + expect(csvContent).toContain("Model"); + expect(csvContent).toContain("Cost/Request"); + expect(csvContent).toContain("gpt-4"); + expect(csvContent).toContain("openai"); + }); + + it("should include the combined totals section in CSV", () => { + let csvContent = ""; + const OriginalBlob = globalThis.Blob; + globalThis.Blob = class extends OriginalBlob { + constructor(parts?: BlobPart[], options?: BlobPropertyBag) { + super(parts, options); + if (typeof parts?.[0] === "string") csvContent = parts[0]; + } + } as unknown as typeof Blob; + + exportMultiToCSV(makeMultiResult()); + globalThis.Blob = OriginalBlob; + + expect(csvContent).toContain("COMBINED TOTALS"); + }); + + it("should create a blob with the correct CSV mime type", () => { + let capturedType = ""; + const OriginalBlob = globalThis.Blob; + globalThis.Blob = class extends OriginalBlob { + constructor(parts?: BlobPart[], options?: BlobPropertyBag) { + super(parts, options); + if (options?.type) capturedType = options.type; + } + } as unknown as typeof Blob; + + exportMultiToCSV(makeMultiResult()); + globalThis.Blob = OriginalBlob; + + expect(capturedType).toBe("text/csv;charset=utf-8;"); + }); + + it("should skip entries with null results", () => { + const multiResult: MultiModelResult = { + entries: [ + { entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, result: null, loading: false, error: null }, + ], + totals: { cost_per_request: 0, daily_cost: null, monthly_cost: null, margin_per_request: 0, daily_margin: null, monthly_margin: null }, + }; + + let csvContent = ""; + const OriginalBlob = globalThis.Blob; + globalThis.Blob = class extends OriginalBlob { + constructor(parts?: BlobPart[], options?: BlobPropertyBag) { + super(parts, options); + if (typeof parts?.[0] === "string") csvContent = parts[0]; + } + } as unknown as typeof Blob; + + exportMultiToCSV(multiResult); + globalThis.Blob = OriginalBlob; + + // CSV should have metadata rows but no model data row for gpt-4 + const lines = csvContent.split("\n").filter((l) => l.includes('"gpt-4"')); + expect(lines).toHaveLength(0); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/use_multi_cost_estimate.test.ts b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/use_multi_cost_estimate.test.ts new file mode 100644 index 00000000000..f5715194f58 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/pricing_calculator/use_multi_cost_estimate.test.ts @@ -0,0 +1,342 @@ +import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { renderHook, act } from "@testing-library/react"; +import { useMultiCostEstimate } from "./use_multi_cost_estimate"; +import type { ModelEntry } from "./types"; +import type { CostEstimateResponse } from "../types"; + +vi.mock("@/components/networking", () => ({ + getProxyBaseUrl: vi.fn(() => ""), + getGlobalLitellmHeaderName: vi.fn(() => "Authorization"), +})); + +function makeEntry(overrides: Partial = {}): ModelEntry { + return { + id: "entry-1", + model: "gpt-4", + input_tokens: 1000, + output_tokens: 500, + ...overrides, + }; +} + +function makeApiResponse(overrides: Partial = {}): CostEstimateResponse { + return { + model: "gpt-4", + input_tokens: 1000, + output_tokens: 500, + num_requests_per_day: null, + num_requests_per_month: null, + cost_per_request: 0.05, + input_cost_per_request: 0.03, + output_cost_per_request: 0.02, + margin_cost_per_request: 0, + daily_cost: null, + daily_input_cost: null, + daily_output_cost: null, + daily_margin_cost: null, + monthly_cost: null, + monthly_input_cost: null, + monthly_output_cost: null, + monthly_margin_cost: null, + input_cost_per_token: 0.00003, + output_cost_per_token: 0.00004, + provider: "openai", + ...overrides, + }; +} + +describe("useMultiCostEstimate", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.useFakeTimers(); + }); + + afterEach(() => { + vi.useRealTimers(); + }); + + describe("debouncedFetchForEntry", () => { + it("should not fetch when access token is null", async () => { + const fetchSpy = vi.spyOn(global, "fetch"); + const { result } = renderHook(() => useMultiCostEstimate(null)); + + await act(async () => { + result.current.debouncedFetchForEntry(makeEntry()); + await vi.runAllTimersAsync(); + }); + + expect(fetchSpy).not.toHaveBeenCalled(); + }); + + it("should not fetch when the model field is empty", async () => { + const fetchSpy = vi.spyOn(global, "fetch"); + const { result } = renderHook(() => useMultiCostEstimate("token123")); + + await act(async () => { + result.current.debouncedFetchForEntry(makeEntry({ model: "" })); + await vi.runAllTimersAsync(); + }); + + expect(fetchSpy).not.toHaveBeenCalled(); + }); + + it("should not fetch immediately — only after the debounce delay", async () => { + const fetchSpy = vi.spyOn(global, "fetch").mockResolvedValue({ + ok: true, + json: async () => makeApiResponse(), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + + act(() => { + result.current.debouncedFetchForEntry(makeEntry()); + }); + + expect(fetchSpy).not.toHaveBeenCalled(); + + await act(async () => { + await vi.runAllTimersAsync(); + }); + + expect(fetchSpy).toHaveBeenCalledTimes(1); + }); + + it("should cancel an in-flight debounce when called again for the same entry", async () => { + const fetchSpy = vi.spyOn(global, "fetch").mockResolvedValue({ + ok: true, + json: async () => makeApiResponse(), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + + await act(async () => { + result.current.debouncedFetchForEntry(makeEntry()); + vi.advanceTimersByTime(200); + result.current.debouncedFetchForEntry(makeEntry()); + vi.advanceTimersByTime(200); + result.current.debouncedFetchForEntry(makeEntry()); + await vi.runAllTimersAsync(); + }); + + expect(fetchSpy).toHaveBeenCalledTimes(1); + }); + + it("should store the API result after a successful fetch", async () => { + vi.spyOn(global, "fetch").mockResolvedValue({ + ok: true, + json: async () => makeApiResponse(), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + await act(async () => { + result.current.debouncedFetchForEntry(entry); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry]); + expect(multiResult.entries[0].result).not.toBeNull(); + expect(multiResult.entries[0].result?.cost_per_request).toBe(0.05); + expect(multiResult.entries[0].loading).toBe(false); + expect(multiResult.entries[0].error).toBeNull(); + }); + + it("should set an error message when the API returns a non-ok response", async () => { + vi.spyOn(global, "fetch").mockResolvedValue({ + ok: false, + json: async () => ({ detail: { error: "Model not found" } }), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + await act(async () => { + result.current.debouncedFetchForEntry(entry); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry]); + expect(multiResult.entries[0].result).toBeNull(); + expect(multiResult.entries[0].error).toBe("Model not found"); + }); + + it("should fall back to detail string when error has no nested error field", async () => { + vi.spyOn(global, "fetch").mockResolvedValue({ + ok: false, + json: async () => ({ detail: "Bad request" }), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + await act(async () => { + result.current.debouncedFetchForEntry(entry); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry]); + expect(multiResult.entries[0].error).toBe("Bad request"); + }); + + it("should set 'Network error' when fetch throws", async () => { + vi.spyOn(global, "fetch").mockRejectedValue(new Error("connection refused")); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + await act(async () => { + result.current.debouncedFetchForEntry(entry); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry]); + expect(multiResult.entries[0].error).toBe("Network error"); + expect(multiResult.entries[0].result).toBeNull(); + }); + }); + + describe("removeEntry", () => { + it("should remove an entry's cached result", async () => { + vi.spyOn(global, "fetch").mockResolvedValue({ + ok: true, + json: async () => makeApiResponse(), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + await act(async () => { + result.current.debouncedFetchForEntry(entry); + await vi.runAllTimersAsync(); + }); + + // Confirm result was stored + expect(result.current.getMultiModelResult([entry]).entries[0].result).not.toBeNull(); + + act(() => { + result.current.removeEntry(entry.id); + }); + + // After removal, the entry should return as if it never fetched + const multiResult = result.current.getMultiModelResult([entry]); + expect(multiResult.entries[0].result).toBeNull(); + }); + + it("should cancel a pending debounce for the removed entry", async () => { + const fetchSpy = vi.spyOn(global, "fetch").mockResolvedValue({ + ok: true, + json: async () => makeApiResponse(), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + act(() => { + result.current.debouncedFetchForEntry(entry); + result.current.removeEntry(entry.id); + }); + + await act(async () => { + await vi.runAllTimersAsync(); + }); + + expect(fetchSpy).not.toHaveBeenCalled(); + }); + }); + + describe("getMultiModelResult", () => { + it("should return zero totals when no entries have results", () => { + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const multiResult = result.current.getMultiModelResult([makeEntry()]); + + expect(multiResult.totals.cost_per_request).toBe(0); + expect(multiResult.totals.margin_per_request).toBe(0); + expect(multiResult.totals.daily_cost).toBeNull(); + expect(multiResult.totals.monthly_cost).toBeNull(); + }); + + it("should return an empty entries array for an empty input list", () => { + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const multiResult = result.current.getMultiModelResult([]); + + expect(multiResult.entries).toHaveLength(0); + expect(multiResult.totals.daily_cost).toBeNull(); + expect(multiResult.totals.monthly_cost).toBeNull(); + }); + + it("should sum cost_per_request across multiple loaded entries", async () => { + const entry1 = makeEntry({ id: "e1", model: "gpt-4" }); + const entry2 = makeEntry({ id: "e2", model: "claude-3" }); + + let callIndex = 0; + const responses = [ + makeApiResponse({ cost_per_request: 0.05, margin_cost_per_request: 0 }), + makeApiResponse({ model: "claude-3", cost_per_request: 0.10, margin_cost_per_request: 0 }), + ]; + + vi.spyOn(global, "fetch").mockImplementation(async () => ({ + ok: true, + json: async () => responses[callIndex++], + } as Response)); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + + await act(async () => { + result.current.debouncedFetchForEntry(entry1); + result.current.debouncedFetchForEntry(entry2); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry1, entry2]); + expect(multiResult.totals.cost_per_request).toBeCloseTo(0.15); + }); + + it("should accumulate daily cost only when entries have a daily cost", async () => { + const entry1 = makeEntry({ id: "e1", model: "gpt-4" }); + const entry2 = makeEntry({ id: "e2", model: "claude-3" }); + + let callIndex = 0; + const responses = [ + makeApiResponse({ daily_cost: 5.0, daily_margin_cost: 0, monthly_cost: null, monthly_margin_cost: null }), + makeApiResponse({ model: "claude-3", daily_cost: 10.0, daily_margin_cost: 0, monthly_cost: null, monthly_margin_cost: null }), + ]; + + vi.spyOn(global, "fetch").mockImplementation(async () => ({ + ok: true, + json: async () => responses[callIndex++], + } as Response)); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + + await act(async () => { + result.current.debouncedFetchForEntry(entry1); + result.current.debouncedFetchForEntry(entry2); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry1, entry2]); + expect(multiResult.totals.daily_cost).toBeCloseTo(15.0); + expect(multiResult.totals.monthly_cost).toBeNull(); + }); + + it("should mark each entry's loading and error state from cached data", async () => { + vi.spyOn(global, "fetch").mockResolvedValue({ + ok: false, + json: async () => ({ detail: "Not found" }), + } as Response); + + const { result } = renderHook(() => useMultiCostEstimate("token123")); + const entry = makeEntry(); + + await act(async () => { + result.current.debouncedFetchForEntry(entry); + await vi.runAllTimersAsync(); + }); + + const multiResult = result.current.getMultiModelResult([entry]); + expect(multiResult.entries[0].error).toBe("Not found"); + expect(multiResult.entries[0].loading).toBe(false); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/EntityUsageExport/UsageExportHeader.tsx b/ui/litellm-dashboard/src/components/EntityUsageExport/UsageExportHeader.tsx index 989ff8703e8..3c9b695a7dd 100644 --- a/ui/litellm-dashboard/src/components/EntityUsageExport/UsageExportHeader.tsx +++ b/ui/litellm-dashboard/src/components/EntityUsageExport/UsageExportHeader.tsx @@ -17,6 +17,7 @@ interface UsageExportHeaderProps { selectedFilters?: string[]; onFiltersChange?: (filters: string[]) => void; filterOptions?: Array<{ label: string; value: string }>; + filterMode?: "multiple" | "single"; customTitle?: string; compactLayout?: boolean; teams?: Team[]; @@ -32,6 +33,7 @@ const UsageExportHeader: React.FC = ({ selectedFilters = [], onFiltersChange, filterOptions = [], + filterMode = "multiple", customTitle, compactLayout = false, teams = [], @@ -59,11 +61,17 @@ const UsageExportHeader: React.FC = ({
{filterLabel && {filterLabel}} setSelectedUserId(value ?? null)} + filterOption={false} + onSearch={handleUserSearchChange} + searchValue={userSearchInput} + onPopupScroll={handleUserPopupScroll} + loading={isLoadingUsers} + notFoundContent={isLoadingUsers ? : "No users found"} + options={userOptions} + popupRender={(menu) => ( + <> + {menu} + {isFetchingNextUsersPage && ( +
+ +
+ )} + + )} + /> +
+ )}
@@ -560,41 +590,6 @@ const UsagePage: React.FC = ({ teams, organizations }) => { )} - {isAdmin && ( -
- - + + + prev.allowed_mcp_servers_and_groups !== curr.allowed_mcp_servers_and_groups || + prev.mcp_tool_permissions !== curr.mcp_tool_permissions + } + > + {() => ( +
+ ) => form.setFieldsValue({ mcp_tool_permissions: toolPerms })} + /> +
+ )} +
+
+ ); + const handleAgentTypeChange = (value: string) => { setAgentType(value); form.resetFields(); }; - // Get the logo for the selected agent type for the header - const selectedLogo = selectedAgentTypeInfo?.logo_url || agentTypeMetadata.find(a => a.agent_type === "a2a")?.logo_url; + const isCustomAgent = agentType === CUSTOM_AGENT_TYPE; + const selectedLogo = isCustomAgent + ? null + : selectedAgentTypeInfo?.logo_url || + agentTypeMetadata.find((a) => a.agent_type === "a2a")?.logo_url; + + const renderConfigureStep = () => ( + <> + Agent Type} + required + tooltip="Select the type of agent you want to create" + > + + + +
+ {agentType === CUSTOM_AGENT_TYPE ? ( +
+ + + + + + +
+ ) : agentType === "a2a" ? ( + + ) : selectedAgentTypeInfo?.use_a2a_form_fields ? ( + <> + + {selectedAgentTypeInfo.credential_fields.length > 0 && ( +
+

+ {selectedAgentTypeInfo.agent_type_display_name} Settings +

+ {selectedAgentTypeInfo.credential_fields.map((field) => ( + + {field.field_type === "password" ? ( + + ) : ( + + )} + + ))} +
+ )} + + ) : selectedAgentTypeInfo ? ( + + ) : null} +
+ + ); + + const renderAssignKeyStep = () => { + const agentName = form.getFieldValue("agent_name") || "your-agent"; + return ( +
+ {/* Agent name chip */} +
+ } color="purple" className="px-3 py-1 text-sm"> + {agentName} + +
+ +
+ {/* Option: Create new key */} +
setKeyAssignOption("create_new")} + > +
+
+ setKeyAssignOption("create_new")} + /> +
+
+ + Create a new key for this agent +
+

+ A dedicated key scoped to this agent. +

+ {keyAssignOption === "create_new" && ( +
e.stopPropagation()}> +
+ + setNewKeyName(e.target.value)} + placeholder="e.g. my-agent-key" + /> +
+
+ + setSelectedExistingKey(value)} + filterOption={(input, option) => + (option?.label as string ?? "").toLowerCase().includes(input.toLowerCase()) + } + options={existingKeys.map((k) => ({ + label: k.key_alias || k.token?.slice(0, 12) + "…", + value: k.token, + }))} + /> +
+ )} +
+
+
+
+ +
+ +
+
+ ); + }; + + const renderReadyStep = () => ( +
+ +

Agent Created!

+
+ } color="purple" className="px-3 py-1 text-sm"> + {createdAgentName} + +
+ {createdKeyValue && ( +
+ +
+ )} + {assignedKeyAlias && ( +

+ Key {assignedKeyAlias} has been assigned to this agent. +

+ )} + {!createdKeyValue && !assignedKeyAlias && keyAssignOption === "skip" && ( +

+ No key assigned. You can create one from the Virtual Keys page. +

+ )} +
+ ); return ( - {selectedLogo && ( - Agent + {selectedLogo && currentStep < 1 && ( + Agent )}

Add New Agent

} open={visible} - onCancel={handleCancel} + onCancel={handleClose} footer={null} width={900} className="top-8" @@ -134,103 +640,71 @@ const AddAgentForm: React.FC = ({ }} >
+ {/* Step indicator */} + + + + + + +
- {/* Agent Type Selection */} - Agent Type} - required - tooltip="Select the type of agent you want to create" - > - - - - {/* Conditional Form Fields */} -
- {agentType === "a2a" ? ( - - ) : selectedAgentTypeInfo?.use_a2a_form_fields ? ( - // A2A-compatible agents (like Pydantic AI) use full A2A form fields - // plus any additional credential fields - <> - - {selectedAgentTypeInfo.credential_fields.length > 0 && ( -
-

- {selectedAgentTypeInfo.agent_type_display_name} Settings -

- {selectedAgentTypeInfo.credential_fields.map((field) => ( - - {field.field_type === "password" ? ( - - ) : ( - - )} - - ))} -
- )} - - ) : selectedAgentTypeInfo ? ( - - ) : null} -
- - {/* Footer Buttons */} -
- - -
+ {currentStep === 0 && renderConfigureStep()} + {currentStep === 1 && renderMCPToolsStep()} + {currentStep === 2 && renderAssignKeyStep()} + {currentStep === 3 && renderReadyStep()}
+ + {/* Footer navigation */} +
+
+ {currentStep > 0 && currentStep < 3 && ( + + )} +
+
+ {currentStep < 3 && ( + + )} + {currentStep === 0 && ( + + )} + {currentStep === 1 && ( + + )} + {currentStep === 2 && ( + + )} + {currentStep === 3 && ( + + )} +
+
); diff --git a/ui/litellm-dashboard/src/components/agents/agent_card.tsx b/ui/litellm-dashboard/src/components/agents/agent_card.tsx new file mode 100644 index 00000000000..dc684386edf --- /dev/null +++ b/ui/litellm-dashboard/src/components/agents/agent_card.tsx @@ -0,0 +1,103 @@ +import React from "react"; +import { Card, Badge, Tooltip, Button } from "antd"; +import { CopyOutlined, KeyOutlined, WarningOutlined, DeleteOutlined } from "@ant-design/icons"; +import { Agent, AgentKeyInfo } from "./types"; + +interface AgentCardProps { + agent: Agent; + keyInfo?: AgentKeyInfo; + onAgentClick: (agentId: string) => void; + onDeleteClick?: (agentId: string, agentName: string) => void; + accessToken: string | null; + isAdmin: boolean; + onAgentUpdated: () => void; +} + +const AgentCard: React.FC = ({ + agent, + keyInfo, + onAgentClick, + onDeleteClick, + isAdmin, +}) => { + const description = + agent.agent_card_params?.description || "No description"; + const url = agent.agent_card_params?.url; + const hasKey = keyInfo?.has_key ?? false; + const statusBadge = hasKey ? ( + + ) : ( + + ); + + const copyToClipboard = (e: React.MouseEvent, text: string) => { + e.stopPropagation(); + navigator.clipboard.writeText(text); + }; + + return ( + onAgentClick(agent.agent_id)} + > +
+
+
+ + {agent.agent_name} + + + copyToClipboard(e, agent.agent_id)} + className="cursor-pointer text-gray-400 hover:text-blue-500 text-xs shrink-0" + /> + +
+
{statusBadge}
+
+ {isAdmin && onDeleteClick && ( + +
+

+ {description} +

+ {url && ( +

+ {url} +

+ )} +
+ {hasKey ? ( +
+ + {keyInfo?.key_alias || keyInfo?.token_prefix || "Key assigned"} +
+ ) : ( +
+ + No key assigned +
+ )} +
+
+ ); +}; + +export default AgentCard; diff --git a/ui/litellm-dashboard/src/components/agents/agent_card_grid.tsx b/ui/litellm-dashboard/src/components/agents/agent_card_grid.tsx new file mode 100644 index 00000000000..0ba7902f8b2 --- /dev/null +++ b/ui/litellm-dashboard/src/components/agents/agent_card_grid.tsx @@ -0,0 +1,63 @@ +import React from "react"; +import { Skeleton } from "antd"; +import AgentCard from "./agent_card"; +import { Agent, AgentKeyInfo } from "./types"; + +interface AgentCardGridProps { + agentsList: Agent[]; + keyInfoMap: Record; + isLoading: boolean; + onDeleteClick: (agentId: string, agentName: string) => void; + accessToken: string | null; + onAgentUpdated: () => void; + isAdmin: boolean; + onAgentClick: (agentId: string) => void; +} + +const AgentCardGrid: React.FC = ({ + agentsList, + keyInfoMap, + isLoading, + onDeleteClick, + accessToken, + onAgentUpdated, + isAdmin, + onAgentClick, +}) => { + if (isLoading) { + return ( +
+ {[1, 2, 3].map((i) => ( + + ))} +
+ ); + } + + if (!agentsList || agentsList.length === 0) { + return ( +
+

No agents found. Create one to get started.

+
+ ); + } + + return ( +
+ {agentsList.map((agent) => ( + + ))} +
+ ); +}; + +export default AgentCardGrid; diff --git a/ui/litellm-dashboard/src/components/agents/agent_config.ts b/ui/litellm-dashboard/src/components/agents/agent_config.ts index 9dd41eed8ac..f85c4daac66 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_config.ts +++ b/ui/litellm-dashboard/src/components/agents/agent_config.ts @@ -54,9 +54,9 @@ export const AGENT_FORM_CONFIG: { name: "url", label: "URL", type: "url", - required: true, + required: false, placeholder: "http://localhost:9999/", - tooltip: "Base URL where the agent is hosted", + tooltip: "Base URL where the agent is hosted (optional)", }, { name: "version", @@ -237,9 +237,9 @@ export const buildAgentDataFromForm = (values: any, existingAgent?: any) => { agent_name: values.agent_name, agent_card_params: { protocolVersion: values.protocolVersion || "1.0", - name: values.name, - description: values.description, - url: values.url, + name: values.name || values.agent_name, + description: values.description || "", + url: values.url || "", version: values.version || "1.0.0", defaultInputModes: existingAgent?.agent_card_params?.defaultInputModes || ["text"], defaultOutputModes: existingAgent?.agent_card_params?.defaultOutputModes || ["text"], diff --git a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx index 4dc4ad6829b..d5429d2a3b5 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx @@ -10,13 +10,15 @@ const { Panel } = Collapse; interface AgentFormFieldsProps { showAgentName?: boolean; + visiblePanels?: string[]; } /** * Reusable form fields component for agent forms * Uses shared configuration from agent_config.ts */ -const AgentFormFields: React.FC = ({ showAgentName = true }) => { +const AgentFormFields: React.FC = ({ showAgentName = true, visiblePanels }) => { + const shouldShow = (key: string) => !visiblePanels || visiblePanels.includes(key); return ( <> {showAgentName && ( @@ -32,6 +34,7 @@ const AgentFormFields: React.FC = ({ showAgentName = true {/* Basic Information */} + {shouldShow(AGENT_FORM_CONFIG.basic.key) && ( {AGENT_FORM_CONFIG.basic.fields.map((field) => ( = ({ showAgentName = true ))} + )} {/* Skills */} + {shouldShow(AGENT_FORM_CONFIG.skills.key) && ( {(fields, { add, remove }) => ( @@ -127,8 +132,10 @@ const AgentFormFields: React.FC = ({ showAgentName = true )} + )} {/* Capabilities */} + {shouldShow(AGENT_FORM_CONFIG.capabilities.key) && ( {AGENT_FORM_CONFIG.capabilities.fields.map((field) => ( = ({ showAgentName = true ))} + )} {/* Optional Settings */} + {shouldShow(AGENT_FORM_CONFIG.optional.key) && ( {AGENT_FORM_CONFIG.optional.fields.map((field) => ( = ({ showAgentName = true ))} + )} {/* Cost Configuration */} + {shouldShow(AGENT_FORM_CONFIG.cost.key) && ( + )} {/* LiteLLM Parameters */} + {shouldShow(AGENT_FORM_CONFIG.litellm.key) && ( {AGENT_FORM_CONFIG.litellm.fields.map((field) => ( = ({ showAgentName = true ))} + )} ); diff --git a/ui/litellm-dashboard/src/components/agents/agent_info.tsx b/ui/litellm-dashboard/src/components/agents/agent_info.tsx index 8d0febd8417..deb4f900377 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_info.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_info.tsx @@ -205,6 +205,44 @@ const AgentInfoView: React.FC = ({ {formatDate(agent.updated_at)} + {agent.object_permission && + (agent.object_permission.mcp_servers?.length || + agent.object_permission.mcp_access_groups?.length || + (agent.object_permission.mcp_tool_permissions && + Object.keys(agent.object_permission.mcp_tool_permissions).length > 0)) && ( +
+ MCP Tool Permissions + + {agent.object_permission.mcp_servers && agent.object_permission.mcp_servers.length > 0 && ( + + {agent.object_permission.mcp_servers.join(", ")} + + )} + {agent.object_permission.mcp_access_groups && + agent.object_permission.mcp_access_groups.length > 0 && ( + + {agent.object_permission.mcp_access_groups.join(", ")} + + )} + {agent.object_permission.mcp_tool_permissions && + Object.keys(agent.object_permission.mcp_tool_permissions).length > 0 && ( + +
+ {Object.entries(agent.object_permission.mcp_tool_permissions).map( + ([serverId, tools]) => ( +
+ {serverId}:{" "} + {Array.isArray(tools) ? tools.join(", ") : String(tools)} +
+ ) + )} +
+
+ )} +
+
+ )} + {agent.agent_card_params?.skills && agent.agent_card_params.skills.length > 0 && ( diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts index 5c63d334129..2e903b14026 100644 --- a/ui/litellm-dashboard/src/components/agents/types.ts +++ b/ui/litellm-dashboard/src/components/agents/types.ts @@ -1,3 +1,15 @@ +export interface AgentKeyInfo { + key_alias?: string; + token_prefix?: string; + has_key: boolean; +} + +export interface AgentObjectPermission { + mcp_servers?: string[]; + mcp_access_groups?: string[]; + mcp_tool_permissions?: Record; +} + export interface Agent { agent_id: string; agent_name: string; @@ -7,8 +19,10 @@ export interface Agent { }; agent_card_params?: { description?: string; + url?: string; [key: string]: any; }; + object_permission?: AgentObjectPermission; created_at?: string; updated_at?: string; created_by?: string; diff --git a/ui/litellm-dashboard/src/components/common_components/LabeledField.test.tsx b/ui/litellm-dashboard/src/components/common_components/LabeledField.test.tsx new file mode 100644 index 00000000000..6558fe8bad9 --- /dev/null +++ b/ui/litellm-dashboard/src/components/common_components/LabeledField.test.tsx @@ -0,0 +1,57 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; +import LabeledField from "./LabeledField"; + +describe("LabeledField", () => { + it("should render the label and value", () => { + render(); + expect(screen.getByText("User Email")).toBeInTheDocument(); + expect(screen.getByText("test@example.com")).toBeInTheDocument(); + }); + + it("should render the icon when provided", () => { + render( + } />, + ); + expect(screen.getByTestId("test-icon")).toBeInTheDocument(); + }); + + it("should show '-' when value is empty", () => { + render(); + expect(screen.getByText("-")).toBeInTheDocument(); + }); + + it("should show 'Default Proxy Admin' tag when value is default_user_id and defaultUserIdCheck is true", () => { + render( + , + ); + expect(screen.getByText("Default Proxy Admin")).toBeInTheDocument(); + expect(screen.queryByText("default_user_id")).not.toBeInTheDocument(); + }); + + it("should show raw value when value is default_user_id but defaultUserIdCheck is false", () => { + render(); + expect(screen.getByText("default_user_id")).toBeInTheDocument(); + expect(screen.queryByText("Default Proxy Admin")).not.toBeInTheDocument(); + }); + + it("should not be copyable when value is empty", () => { + const { container } = render(); + // antd adds a .ant-typography-copy element when copyable; should not be present + expect(container.querySelector(".ant-typography-copy")).not.toBeInTheDocument(); + }); + + it("should not be copyable when value is default_user_id and defaultUserIdCheck is true", () => { + const { container } = render( + , + ); + expect(container.querySelector(".ant-typography-copy")).not.toBeInTheDocument(); + }); + + it("should be copyable when copyable is true and value is present", () => { + const { container } = render( + , + ); + expect(container.querySelector(".ant-typography-copy")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx b/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx new file mode 100644 index 00000000000..75a8236a2fa --- /dev/null +++ b/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx @@ -0,0 +1,52 @@ +import React from "react"; +import { Typography, Space } from "antd"; +import DefaultProxyAdminTag from "./DefaultProxyAdminTag"; + +const { Text } = Typography; + +interface LabeledFieldProps { + label: string; + value: string; + icon?: React.ReactNode; + truncate?: boolean; + copyable?: boolean; + defaultUserIdCheck?: boolean; +} + +export default function LabeledField({ + label, + value, + icon, + truncate = false, + copyable = false, + defaultUserIdCheck = false, +}: LabeledFieldProps) { + const isEmpty = !value; + const isDefaultUser = defaultUserIdCheck && value === "default_user_id"; + const displayValue = isEmpty ? "-" : value; + const isCopyable = copyable && !isEmpty && !isDefaultUser; + + const valueEl = isDefaultUser ? ( + + ) : ( + + {displayValue} + + ); + return ( +
+ + {icon} + + {label} + + +
{valueEl}
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_garden.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_garden.tsx index 703a6ad3e48..d95ae07e7f3 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_garden.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_garden.tsx @@ -1,7 +1,7 @@ import React, { useState } from "react"; import { Input } from "antd"; import { SearchOutlined, ArrowRightOutlined } from "@ant-design/icons"; -import { GuardrailCardInfo, LITELLM_CONTENT_FILTER_CARDS, PARTNER_GUARDRAIL_CARDS, ALL_CARDS } from "./guardrail_garden_data"; +import { GuardrailCardInfo, ALL_CARDS } from "./guardrail_garden_data"; import GuardrailCard from "./guardrail_garden_card"; import GuardrailDetailView from "./guardrail_garden_detail"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx index 84967f7a83e..404172623b9 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx @@ -194,7 +194,7 @@ const MCPToolConfiguration: React.FC = ({ {filteredTools.length === 0 ? (
- No tools found matching "{toolSearchTerm}" + No tools found matching "{toolSearchTerm}"
) : ( filteredTools.map((tool, index) => ( diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx index 3a572dec893..198bf96c3f0 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx @@ -1,12 +1,12 @@ import React, { useState } from "react"; import { useQuery, useMutation } from "@tanstack/react-query"; import { ToolTestPanel } from "./ToolTestPanel"; -import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse, AUTH_TYPE } from "./types"; +import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse } from "./types"; import { listMCPTools, callMCPTool } from "../networking"; import { Card, Title, Text } from "@tremor/react"; -import { RobotOutlined, ToolOutlined, SearchOutlined, LockOutlined, KeyOutlined } from "@ant-design/icons"; -import { Input, Alert, Button as AntdButton } from "antd"; +import { RobotOutlined, ToolOutlined, SearchOutlined, KeyOutlined } from "@ant-design/icons"; +import { Input, Button as AntdButton } from "antd"; const MCPToolsViewer = ({ serverId, @@ -134,7 +134,7 @@ const MCPToolsViewer = ({ {!showHeaderInput && Object.keys(passthroughHeaders).length === 0 && ( - This server requires additional headers. Click "Configure" to provide values. + This server requires additional headers. Click "Configure" to provide values. )} @@ -255,7 +255,7 @@ const MCPToolsViewer = ({

No tools found

-

No tools match "{toolSearchTerm}"

+

No tools match "{toolSearchTerm}"

) : (
({ + individualModelHealthCheckCall: (...args: unknown[]) => mockIndividualModelHealthCheckCall(...args), + latestHealthChecksCall: (...args: unknown[]) => mockLatestHealthChecksCall(...args), +})); + +describe("HealthCheckComponent", () => { + const getDisplayModelName = (model: { model_name?: string }) => model.model_name ?? ""; + + beforeEach(() => { + vi.clearAllMocks(); + mockLatestHealthChecksCall.mockResolvedValue({ latest_health_checks: {} }); + mockIndividualModelHealthCheckCall.mockResolvedValue({ + healthy_count: 1, + unhealthy_count: 0, + healthy_endpoints: [], + unhealthy_endpoints: [], + }); + }); + + it("should render the health check section", async () => { + const modelData = { + data: [ + { + model_name: "gpt-4", + model_info: { id: "deployment-1" }, + litellm_model_name: "gpt-4", + }, + ], + }; + + await act(async () => { + render( + , + ); + }); + await act(async () => { + await new Promise((r) => setTimeout(r, 0)); + }); + + expect(screen.getByText("Model Health Status")).toBeInTheDocument(); + expect( + screen.getByText("Run health checks on individual models to verify they are working correctly"), + ).toBeInTheDocument(); + }); + + it("should call individualModelHealthCheckCall with model id when run health check is triggered", async () => { + const modelData = { + data: [ + { + model_name: "gpt-4", + model_info: { id: "deployment-abc-123" }, + litellm_model_name: "gpt-4", + }, + ], + }; + + render( + , + ); + + const runButtons = screen.getAllByTestId("run-health-check-btn"); + expect(runButtons.length).toBeGreaterThanOrEqual(1); + const runButton = runButtons[0]; + + await act(async () => { + runButton.click(); + }); + + await act(async () => { + await new Promise((r) => setTimeout(r, 50)); + }); + + expect(mockIndividualModelHealthCheckCall).toHaveBeenCalledWith("token-123", "deployment-abc-123"); + expect(mockIndividualModelHealthCheckCall).not.toHaveBeenCalledWith("token-123", "gpt-4"); + }); + + describe("latest_health_checks keyed by model id", () => { + it("should show status from latest_health_checks when keys match model ids", async () => { + const modelData = { + data: [ + { + model_name: "gpt-4", + model_info: { id: "id-alpha" }, + litellm_model_name: "gpt-4", + }, + { + model_name: "gpt-4", + model_info: { id: "id-beta" }, + litellm_model_name: "gpt-4", + }, + ], + }; + + mockLatestHealthChecksCall.mockResolvedValue({ + latest_health_checks: { + "id-alpha": { + status: "healthy", + checked_at: "2024-01-15T10:00:00Z", + error_message: null, + }, + "id-beta": { + status: "unhealthy", + checked_at: "2024-01-15T10:05:00Z", + error_message: "Connection failed", + }, + }, + }); + + await act(async () => { + render( + , + ); + }); + await act(async () => { + await new Promise((r) => setTimeout(r, 0)); + }); + + expect(mockLatestHealthChecksCall).toHaveBeenCalledWith("token"); + const healthyBadges = screen.getAllByText("healthy"); + const unhealthyBadges = screen.getAllByText("unhealthy"); + expect(healthyBadges.length).toBeGreaterThanOrEqual(1); + expect(unhealthyBadges.length).toBeGreaterThanOrEqual(1); + }); + + it("should skip latest_health_checks entries whose key is not a known model id", async () => { + const modelData = { + data: [ + { + model_name: "gpt-4", + model_info: { id: "current-model-id" }, + litellm_model_name: "gpt-4", + }, + ], + }; + + mockLatestHealthChecksCall.mockResolvedValue({ + latest_health_checks: { + "current-model-id": { + status: "healthy", + checked_at: "2024-01-15T10:00:00Z", + error_message: null, + }, + "deleted-or-unknown-id": { + status: "unhealthy", + checked_at: "2024-01-15T10:05:00Z", + error_message: "Stale entry", + }, + }, + }); + + await act(async () => { + render( + , + ); + }); + await act(async () => { + await new Promise((r) => setTimeout(r, 0)); + }); + + expect(screen.getByText("healthy")).toBeInTheDocument(); + expect(screen.queryByText("unhealthy")).not.toBeInTheDocument(); + }); + + it("should not apply status when latest_health_checks key is model name not model id", async () => { + const modelData = { + data: [ + { + model_name: "gpt-4", + model_info: { id: "model-id-123" }, + litellm_model_name: "gpt-4", + }, + ], + }; + + mockLatestHealthChecksCall.mockResolvedValue({ + latest_health_checks: { + "gpt-4": { + status: "healthy", + checked_at: "2024-01-15T10:00:00Z", + error_message: null, + }, + }, + }); + + await act(async () => { + render( + , + ); + }); + await act(async () => { + await new Promise((r) => setTimeout(r, 0)); + }); + + expect(screen.queryByText("healthy")).not.toBeInTheDocument(); + expect(screen.getByText("none")).toBeInTheDocument(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx index 6fa83494408..db71733b5ba 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard/HealthCheckComponent.tsx @@ -53,65 +53,48 @@ const HealthCheckComponent: React.FC = ({ const healthTableRef = useRef>(null); - // Initialize health statuses on component mount + // Initialize health statuses on component mount (keyed by model id) useEffect(() => { if (!accessToken || !modelData?.data) return; const initializeHealthStatuses = async () => { const healthStatusMap: { [key: string]: HealthStatus } = {}; - // Initialize all models with default state using model names + // Initialize all models with default state using model ids modelData.data.forEach((model: any) => { - const modelName = model.model_name; - healthStatusMap[modelName] = { - status: "none", - lastCheck: "None", - lastSuccess: "None", - loading: false, - error: undefined, - fullError: undefined, - successResponse: undefined, - }; + const modelId = model.model_info?.id; + if (modelId) { + healthStatusMap[modelId] = { + status: "none", + lastCheck: "None", + lastSuccess: "None", + loading: false, + error: undefined, + fullError: undefined, + successResponse: undefined, + }; + } }); try { const latestHealthChecks = await latestHealthChecksCall(accessToken); - // Override with actual database data if it exists + // Override with actual database data if it exists (latest_health_checks is keyed by model_id) if ( latestHealthChecks && latestHealthChecks.latest_health_checks && typeof latestHealthChecks.latest_health_checks === "object" ) { - Object.entries(latestHealthChecks.latest_health_checks).forEach(([key, checkData]: [string, any]) => { + Object.entries(latestHealthChecks.latest_health_checks).forEach(([modelId, checkData]: [string, any]) => { if (!checkData) return; - let targetModelName: string | null = null; + // Key is model_id from the backend (guaranteed by DB schema) + const modelExists = modelData.data.some((m: any) => m.model_info?.id === modelId); + if (!modelExists) return; - // The key could be either model_id or model_name, try both approaches - const directModelMatch = modelData.data.find((m: any) => m.model_name === key); - if (directModelMatch) { - targetModelName = directModelMatch.model_name; - } else { - // If not a direct match, treat as model_id and find the corresponding model - const modelByIdMatch = modelData.data.find((m: any) => m.model_info && m.model_info.id === key); - if (modelByIdMatch) { - targetModelName = modelByIdMatch.model_name; - } else { - // Check if checkData contains model_name and use that - if (checkData.model_name) { - const modelByNameInData = modelData.data.find((m: any) => m.model_name === checkData.model_name); - if (modelByNameInData) { - targetModelName = modelByNameInData.model_name; - } - } - } - } + const fullError = checkData.error_message || undefined; - if (targetModelName) { - const fullError = checkData.error_message || undefined; - - healthStatusMap[targetModelName] = { + healthStatusMap[modelId] = { status: checkData.status || "unknown", lastCheck: checkData.checked_at ? new Date(checkData.checked_at).toLocaleString() : "None", lastSuccess: @@ -125,7 +108,6 @@ const HealthCheckComponent: React.FC = ({ fullError: fullError, successResponse: checkData.status === "healthy" ? checkData : undefined, }; - } }); } } catch (healthError) { @@ -246,33 +228,31 @@ const HealthCheckComponent: React.FC = ({ return cleaned.length > 100 ? cleaned.substring(0, 97) + "..." : cleaned; }; - const runIndividualHealthCheck = async (modelName: string) => { + const runIndividualHealthCheck = async (modelId: string) => { if (!accessToken) return; setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { - ...prev[modelName], + [modelId]: { + ...prev[modelId], loading: true, status: "checking", }, })); try { - // Run the health check and process the response directly - const response = await individualModelHealthCheckCall(accessToken, modelName); + const response = await individualModelHealthCheckCall(accessToken, modelId); const currentTime = new Date().toLocaleString(); - // Check if there are any unhealthy endpoints (which means this specific model failed) if (response.unhealthy_count > 0 && response.unhealthy_endpoints && response.unhealthy_endpoints.length > 0) { const rawError = response.unhealthy_endpoints[0]?.error || "Health check failed"; const errorMessage = extractMeaningfulError(rawError); setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { + [modelId]: { status: "unhealthy", lastCheck: currentTime, - lastSuccess: prev[modelName]?.lastSuccess || "None", + lastSuccess: prev[modelId]?.lastSuccess || "None", loading: false, error: errorMessage, fullError: rawError, @@ -281,7 +261,7 @@ const HealthCheckComponent: React.FC = ({ } else { setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { + [modelId]: { status: "healthy", lastCheck: currentTime, lastSuccess: currentTime, @@ -291,41 +271,33 @@ const HealthCheckComponent: React.FC = ({ })); } - // Refresh health status from database to get the saved check data including timestamp try { const latestHealthChecks = await latestHealthChecksCall(accessToken); + const checkData = latestHealthChecks.latest_health_checks?.[modelId]; - // Find the model ID for this model name to look up database data - const model = modelData.data.find((m: any) => m.model_name === modelName); - if (model) { - const modelId = model.model_info.id; - const checkData = latestHealthChecks.latest_health_checks?.[modelId]; - - if (checkData) { - const fullError = checkData.error_message || undefined; - setModelHealthStatuses((prev) => ({ - ...prev, - [modelName]: { - status: checkData.status || prev[modelName]?.status || "unknown", - lastCheck: checkData.checked_at - ? new Date(checkData.checked_at).toLocaleString() - : prev[modelName]?.lastCheck || "None", - lastSuccess: - checkData.status === "healthy" - ? checkData.checked_at - ? new Date(checkData.checked_at).toLocaleString() - : prev[modelName]?.lastSuccess || "None" - : prev[modelName]?.lastSuccess || "None", - loading: false, - error: fullError ? extractMeaningfulError(fullError) : prev[modelName]?.error, - fullError: fullError || prev[modelName]?.fullError, - successResponse: checkData.status === "healthy" ? checkData : prev[modelName]?.successResponse, - }, - })); - } + if (checkData) { + const fullError = checkData.error_message || undefined; + setModelHealthStatuses((prev) => ({ + ...prev, + [modelId]: { + status: checkData.status || prev[modelId]?.status || "unknown", + lastCheck: checkData.checked_at + ? new Date(checkData.checked_at).toLocaleString() + : prev[modelId]?.lastCheck || "None", + lastSuccess: + checkData.status === "healthy" + ? checkData.checked_at + ? new Date(checkData.checked_at).toLocaleString() + : prev[modelId]?.lastSuccess || "None" + : prev[modelId]?.lastSuccess || "None", + loading: false, + error: fullError ? extractMeaningfulError(fullError) : prev[modelId]?.error, + fullError: fullError || prev[modelId]?.fullError, + successResponse: checkData.status === "healthy" ? checkData : prev[modelId]?.successResponse, + }, + })); } } catch (dbError) { - // Ignore database errors - we already have the health check result from the API call console.debug("Could not fetch updated status from database (non-critical):", dbError); } } catch (error) { @@ -334,10 +306,10 @@ const HealthCheckComponent: React.FC = ({ const errorMessage = extractMeaningfulError(rawError); setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { + [modelId]: { status: "unhealthy", lastCheck: currentTime, - lastSuccess: prev[modelName]?.lastSuccess || "None", + lastSuccess: prev[modelId]?.lastSuccess || "None", loading: false, error: errorMessage, fullError: rawError, @@ -349,11 +321,10 @@ const HealthCheckComponent: React.FC = ({ const runAllHealthChecks = async () => { const modelsToCheck = selectedModelsForHealth.length > 0 ? selectedModelsForHealth : all_models_on_proxy; - // Set all models to loading state const loadingStatuses = modelsToCheck.reduce( - (acc, modelName) => { - acc[modelName] = { - ...modelHealthStatuses[modelName], + (acc, modelId) => { + acc[modelId] = { + ...modelHealthStatuses[modelId], loading: true, status: "checking", }; @@ -364,30 +335,25 @@ const HealthCheckComponent: React.FC = ({ setModelHealthStatuses((prev) => ({ ...prev, ...loadingStatuses })); - // Store results from individual health checks const healthCheckResults: { [key: string]: any } = {}; - // Run all health checks in parallel and collect results - const healthCheckPromises = modelsToCheck.map(async (modelName) => { + const healthCheckPromises = modelsToCheck.map(async (modelId) => { if (!accessToken) return; try { - // Run the health check and store the result - const response = await individualModelHealthCheckCall(accessToken, modelName); - healthCheckResults[modelName] = response; + const response = await individualModelHealthCheckCall(accessToken, modelId); + healthCheckResults[modelId] = response; - // Update status immediately based on response const currentTime = new Date().toLocaleString(); - // Check if there are any unhealthy endpoints (which means this specific model failed) if (response.unhealthy_count > 0 && response.unhealthy_endpoints && response.unhealthy_endpoints.length > 0) { const rawError = response.unhealthy_endpoints[0]?.error || "Health check failed"; const errorMessage = extractMeaningfulError(rawError); setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { + [modelId]: { status: "unhealthy", lastCheck: currentTime, - lastSuccess: prev[modelName]?.lastSuccess || "None", + lastSuccess: prev[modelId]?.lastSuccess || "None", loading: false, error: errorMessage, fullError: rawError, @@ -396,7 +362,7 @@ const HealthCheckComponent: React.FC = ({ } else { setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { + [modelId]: { status: "healthy", lastCheck: currentTime, lastSuccess: currentTime, @@ -406,17 +372,16 @@ const HealthCheckComponent: React.FC = ({ })); } } catch (error) { - console.error(`Health check failed for ${modelName}:`, error); - // Set error status for failed health checks + console.error(`Health check failed for model id ${modelId}:`, error); const currentTime = new Date().toLocaleString(); const rawError = error instanceof Error ? error.message : String(error); const errorMessage = extractMeaningfulError(rawError); setModelHealthStatuses((prev) => ({ ...prev, - [modelName]: { + [modelId]: { status: "unhealthy", lastCheck: currentTime, - lastSuccess: prev[modelName]?.lastSuccess || "None", + lastSuccess: prev[modelId]?.lastSuccess || "None", loading: false, error: errorMessage, fullError: rawError, @@ -425,27 +390,21 @@ const HealthCheckComponent: React.FC = ({ } }); - // Wait for all health checks to complete await Promise.allSettled(healthCheckPromises); - // Refresh health statuses from database to get the saved check data including timestamps try { if (!accessToken) return; const latestHealthChecks = await latestHealthChecksCall(accessToken); if (latestHealthChecks.latest_health_checks) { - // Update health statuses from database, which should have the most accurate saved data Object.entries(latestHealthChecks.latest_health_checks).forEach(([modelId, checkData]: [string, any]) => { - // Find the model name for this model ID - const model = modelData.data.find((m: any) => m.model_info.id === modelId); - if (model && modelsToCheck.includes(model.model_name) && checkData) { - const modelName = model.model_name; + if (modelsToCheck.includes(modelId) && checkData) { const fullError = checkData.error_message || undefined; setModelHealthStatuses((prev) => { - const currentStatus = prev[modelName]; + const currentStatus = prev[modelId]; return { ...prev, - [modelName]: { + [modelId]: { status: checkData.status || currentStatus?.status || "unknown", lastCheck: checkData.checked_at ? new Date(checkData.checked_at).toLocaleString() @@ -468,15 +427,14 @@ const HealthCheckComponent: React.FC = ({ } } catch (dbError) { console.warn("Failed to fetch updated health statuses from database (non-critical):", dbError); - // This is non-critical - we already have the health check results from the API calls } }; - const handleModelSelection = (modelName: string, checked: boolean) => { + const handleModelSelection = (modelId: string, checked: boolean) => { if (checked) { - setSelectedModelsForHealth((prev) => [...prev, modelName]); + setSelectedModelsForHealth((prev) => [...prev, modelId]); } else { - setSelectedModelsForHealth((prev) => prev.filter((name) => name !== modelName)); + setSelectedModelsForHealth((prev) => prev.filter((id) => id !== modelId)); setAllModelsSelected(false); } }; @@ -580,8 +538,9 @@ const HealthCheckComponent: React.FC = ({ teams, )} data={modelData.data.map((model: any) => { - const modelName = model.model_name; - const healthStatus = modelHealthStatuses[modelName] || { + const modelId = model.model_info?.id; + const healthStatus = modelId ? modelHealthStatuses[modelId] : null; + const status = healthStatus || { status: "none", lastCheck: "None", loading: false, @@ -591,12 +550,12 @@ const HealthCheckComponent: React.FC = ({ model_info: model.model_info, provider: model.provider, litellm_model_name: model.litellm_model_name, - health_status: healthStatus.status, - last_check: healthStatus.lastCheck, - last_success: healthStatus.lastSuccess || "None", - health_loading: healthStatus.loading, - health_error: healthStatus.error, - health_full_error: healthStatus.fullError, + health_status: status.status, + last_check: status.lastCheck, + last_success: status.lastSuccess || "None", + health_loading: status.loading, + health_error: status.error, + health_full_error: status.fullError, }; })} isLoading={false} diff --git a/ui/litellm-dashboard/src/components/model_dashboard/health_check_columns.tsx b/ui/litellm-dashboard/src/components/model_dashboard/health_check_columns.tsx index 3e8ae662ad4..396afb7ed0b 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard/health_check_columns.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard/health_check_columns.tsx @@ -35,9 +35,9 @@ export const healthCheckColumns = ( modelHealthStatuses: { [key: string]: HealthStatus }, selectedModelsForHealth: string[], allModelsSelected: boolean, - handleModelSelection: (modelName: string, checked: boolean) => void, + handleModelSelection: (modelId: string, checked: boolean) => void, handleSelectAll: (checked: boolean) => void, - runIndividualHealthCheck: (modelName: string) => void, + runIndividualHealthCheck: (modelId: string) => void, getStatusBadge: (status: string) => JSX.Element, getDisplayModelName: (model: any) => string, showErrorModal?: (modelName: string, cleanedError: string, fullError: string) => void, @@ -62,14 +62,14 @@ export const healthCheckColumns = ( sortingFn: "alphanumeric", cell: ({ row }) => { const model = row.original; - const modelName = model.model_name; - const isSelected = selectedModelsForHealth.includes(modelName); + const modelId = model.model_info?.id ?? ""; + const isSelected = selectedModelsForHealth.includes(modelId); return (
handleModelSelection(modelName, e.target.checked)} + onChange={(e) => handleModelSelection(modelId, e.target.checked)} onClick={(e) => e.stopPropagation()} /> @@ -169,8 +169,9 @@ export const healthCheckColumns = ( ); } - const modelName = model.model_name; - const hasSuccessResponse = healthStatus.status === "healthy" && modelHealthStatuses[modelName]?.successResponse; + const modelId = model.model_info?.id ?? ""; + const displayName = getDisplayModelName(model) || model.model_name; + const hasSuccessResponse = healthStatus.status === "healthy" && modelHealthStatuses[modelId]?.successResponse; return (
@@ -178,7 +179,7 @@ export const healthCheckColumns = ( {hasSuccessResponse && showSuccessModal && (
)} + {keyOwner === "agent" && ( +
+
+ + Select Agent * + +
+