diff --git a/.circleci/config.yml b/.circleci/config.yml index 02a9e3b0714..0b369477c28 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -119,7 +119,7 @@ jobs: username: ${DOCKERHUB_USERNAME} password: ${DOCKERHUB_PASSWORD} working_directory: ~/project - + parallelism: 4 steps: - checkout - setup_google_dns @@ -207,12 +207,22 @@ jobs: - run: name: Run tests command: | - pwd - ls - # Add --timeout to kill hanging tests after 300s (5 min) - # Add -v to show test names as they run for debugging - # Add --tb=short for shorter tracebacks - python -m pytest -vv tests/local_testing --cov=litellm --cov-report=xml --junitxml=test-results/junit.xml --durations=20 -k "not test_python_38.py and not test_basic_python_version.py and not router and not assistants and not langfuse and not caching and not cache" -n 4 --timeout=300 --timeout_method=thread + mkdir test-results + # Discover test files + TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py") + echo "$TEST_FILES" | circleci tests run \ + --split-by=filesize \ + --verbose \ + --command="xargs python -m pytest \ + -vv \ + --cov=litellm \ + --cov-report=xml \ + --junitxml=test-results/junit.xml \ + --durations=20 \ + -k \"not test_python_38.py and not test_basic_python_version.py and not router and not assistants and not langfuse and not caching and not cache\" \ + -n 4 \ + --timeout=300 \ + --timeout_method=thread" no_output_timeout: 120m - run: name: Rename the coverage files @@ -499,7 +509,7 @@ jobs: username: ${DOCKERHUB_USERNAME} password: ${DOCKERHUB_PASSWORD} working_directory: ~/project - + parallelism: 4 steps: - checkout - setup_google_dns @@ -513,6 +523,7 @@ jobs: pip install "pytest-cov==5.0.0" pip install "pytest-retry==1.6.3" pip install "pytest-asyncio==0.21.1" + pip install "pytest-xdist==3.6.1" pip install semantic_router --no-deps pip install aurelio_sdk --no-deps # Run pytest and generate JUnit XML report @@ -520,9 +531,21 @@ jobs: - run: name: Run tests command: | - pwd - ls - python -m pytest tests/local_testing --cov=litellm --cov-report=xml -vv -k "router" -v --junitxml=test-results/junit.xml --durations=5 + mkdir test-results + # Find test files only in local_testing + TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py") + echo "$TEST_FILES" | circleci tests run \ + --split-by=filesize \ + --verbose \ + --command="xargs python -m pytest -o junit_family=legacy \ + -k 'router' \ + --cov=litellm \ + --cov-report=xml \ + -n 4 \ + --dist=loadscope \ + --junitxml=test-results/junit.xml \ + --durations=5 \ + -vv" no_output_timeout: 120m - run: name: Rename the coverage files @@ -1743,13 +1766,14 @@ jobs: pip install "pytest-cov==5.0.0" pip install "pytest-asyncio==0.21.1" pip install "respx==0.22.0" + pip install "pytest-xdist==3.6.1" # Run pytest and generate JUnit XML report - run: name: Run tests command: | pwd ls - python -m pytest -vv tests/image_gen_tests --cov=litellm --cov-report=xml -x -v --junitxml=test-results/junit.xml --durations=5 + python -m pytest -vv tests/image_gen_tests -n 4 --cov=litellm --cov-report=xml -x -v --junitxml=test-results/junit.xml --durations=5 no_output_timeout: 120m - run: name: Rename the coverage files @@ -2192,6 +2216,8 @@ jobs: pip install "asyncio==3.4.3" pip install "PyGithub==1.59.1" pip install "openai==1.100.1" + pip install "litellm[proxy]" + pip install "pytest-xdist==3.6.1" - run: name: Install dockerize command: | @@ -2268,7 +2294,7 @@ jobs: command: | pwd ls - python -m pytest -s -vv tests/*.py -x --junitxml=test-results/junit.xml --durations=5 --ignore=tests/otel_tests --ignore=tests/spend_tracking_tests --ignore=tests/pass_through_tests --ignore=tests/proxy_admin_ui_tests --ignore=tests/load_tests --ignore=tests/llm_translation --ignore=tests/llm_responses_api_testing --ignore=tests/mcp_tests --ignore=tests/guardrails_tests --ignore=tests/image_gen_tests --ignore=tests/pass_through_unit_tests + python -m pytest -s -vv tests/*.py -x --junitxml=test-results/junit.xml -n 4 --durations=5 --ignore=tests/otel_tests --ignore=tests/spend_tracking_tests --ignore=tests/pass_through_tests --ignore=tests/proxy_admin_ui_tests --ignore=tests/load_tests --ignore=tests/llm_translation --ignore=tests/llm_responses_api_testing --ignore=tests/mcp_tests --ignore=tests/guardrails_tests --ignore=tests/image_gen_tests --ignore=tests/pass_through_unit_tests no_output_timeout: 120m # Store test results diff --git a/.gitignore b/.gitignore index 9d9e28dc466..0248d68c1e1 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ .python-version .venv +.venv_policy_test .env .newenv newenv/* diff --git a/ci_cd/security_scans.sh b/ci_cd/security_scans.sh index 04f3e27a944..cf026eb5263 100755 --- a/ci_cd/security_scans.sh +++ b/ci_cd/security_scans.sh @@ -138,6 +138,21 @@ run_grype_scans() { "CVE-2026-22184" # zlib untgz buffer overflow - untgz unused + no fixed Wolfi build yet "GHSA-58pv-8j8x-9vj2" # jaraco.context path traversal - setuptools vendored only (v5.3.0), not used in application code (using v6.1.0+) "GHSA-r6q2-hw4h-h46w" # node-tar not used by application runtime, Linux-only container, not affect by macOS APFS-specific exploit + "GHSA-8rrh-rw8j-w5fx" # wheel is from chainguard and will be handled by then TODO: Remove this after Chainguard updates the wheel + "CVE-2025-59465" # We do not use Node in application runtime, only used for building Admin UI + "CVE-2025-55131" # We do not use Node in application runtime, only used for building Admin UI + "CVE-2025-59466" # We do not use Node in application runtime, only used for building Admin UI + "CVE-2025-55130" # We do not use Node in application runtime, only used for building Admin UI + "CVE-2025-59467" # We do not use Node in application runtime, only used for building Admin UI + "CVE-2026-21637" # We do not use Node in application runtime, only used for building Admin UI + "CVE-2025-15281" # No fix available yet + "CVE-2026-0865" # No fix available yet + "CVE-2025-15282" # No fix available yet + "CVE-2026-0672" # No fix available yet + "CVE-2025-15366" # No fix available yet + "CVE-2025-15367" # No fix available yet + "CVE-2025-12781" # No fix available yet + "CVE-2025-11468" # No fix available yet ) # Build JSON array of allowlisted CVE IDs for jq diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 89e1e2910e4..6d3fa206113 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -611,6 +611,8 @@ router_settings: | GALILEO_USERNAME | Username for Galileo authentication | GOOGLE_SECRET_MANAGER_PROJECT_ID | Project ID for Google Secret Manager | GCS_BUCKET_NAME | Name of the Google Cloud Storage bucket +| GCS_MOCK | Enable mock mode for GCS integration testing. When set to true, intercepts GCS API calls and returns mock responses without making actual network calls. Default is false +| GCS_MOCK_LATENCY_MS | Mock latency in milliseconds for GCS API calls when mock mode is enabled. Simulates network round-trip time. Default is 150ms | GCS_PATH_SERVICE_ACCOUNT | Path to the Google Cloud service account JSON file | GCS_FLUSH_INTERVAL | Flush interval for GCS logging (in seconds). Specify how often you want a log to be sent to GCS. **Default is 20 seconds** | GCS_BATCH_SIZE | Batch size for GCS logging. Specify after how many logs you want to flush to GCS. If `BATCH_SIZE` is set to 10, logs are flushed every 10 logs. **Default is 2048** @@ -697,6 +699,8 @@ router_settings: | LANGFUSE_FLUSH_INTERVAL | Interval for flushing Langfuse logs | LANGFUSE_TRACING_ENVIRONMENT | Environment for Langfuse tracing | LANGFUSE_HOST | Host URL for Langfuse service +| LANGFUSE_MOCK | Enable mock mode for Langfuse integration testing. When set to true, intercepts Langfuse API calls and returns mock responses without making actual network calls. Default is false +| LANGFUSE_MOCK_LATENCY_MS | Mock latency in milliseconds for Langfuse API calls when mock mode is enabled. Simulates network round-trip time. Default is 100ms | LANGFUSE_PUBLIC_KEY | Public key for Langfuse authentication | LANGFUSE_RELEASE | Release version of Langfuse integration | LANGFUSE_SECRET_KEY | Secret key for Langfuse authentication diff --git a/docs/my-website/docs/troubleshoot/spend_queue_warnings.md b/docs/my-website/docs/troubleshoot/spend_queue_warnings.md new file mode 100644 index 00000000000..4be8b18f5cd --- /dev/null +++ b/docs/my-website/docs/troubleshoot/spend_queue_warnings.md @@ -0,0 +1,46 @@ +# Spend Update Queue Full Warnings + +## Overview + +The "Spend update queue is full" warning occurs in high-volume LiteLLM proxy deployments when the internal spend tracking queue reaches capacity. This is a protective mechanism to prevent memory issues during traffic spikes. + +## Warning Message + +``` +WARNING:litellm.proxy.db.db_transaction_queue.spend_update_queue:Spend update queue is full. Aggregating entries to prevent memory issues. +``` + +## Root Cause + +The spend update queue has a default maximum size of 10,000 entries (`MAX_SIZE_IN_MEMORY_QUEUE=10000`). When this limit is reached: + +1. New spend tracking entries are aggregated instead of queued individually +2. This prevents memory exhaustion but may slightly delay spend updates +3. The warning indicates your deployment is processing requests faster than the database can handle spend updates + +## Solutions + +### 1. Increase Queue Size + +Set the `MAX_SIZE_IN_MEMORY_QUEUE` environment variable to a higher value: + +```bash +MAX_SIZE_IN_MEMORY_QUEUE=50000 +``` + +**Tradeoffs:** +Higher queue sizes store more items in memory - provision at least 8GB RAM for large queues +- Recommended for deployments with consistent high traffic + +### 2. Horizontal Scaling + +Deploy multiple proxy instances with load balancing. This distributes the spend tracking load across multiple queues, reducing the pressure on any single instance's spend update queue. + + + +## Related Configuration + +```yaml +# Environment variables +MAX_SIZE_IN_MEMORY_QUEUE: 10000 # Default queue size +``` diff --git a/docs/my-website/release_notes/v1.81.0/index.md b/docs/my-website/release_notes/v1.81.0/index.md index 88ac240c614..d69a62b6941 100644 --- a/docs/my-website/release_notes/v1.81.0/index.md +++ b/docs/my-website/release_notes/v1.81.0/index.md @@ -62,7 +62,7 @@ This means you can now use Claude Code's web search tool with any provider, not Proxy Admins can configure web search interception in their LiteLLM proxy config to enable this capability for their teams using Claude Code with Bedrock, Azure, or any other supported provider. -[**Learn more →**](../../docs/tutorials/claude_code_websearch.md) +[**Learn more →**](https://docs.litellm.ai/docs/tutorials/claude_code_websearch) --- diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 9bf4e167e20..a72f295c597 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -1017,6 +1017,7 @@ const sidebars = { items: [ "troubleshoot/cpu_issues", "troubleshoot/memory_issues", + "troubleshoot/spend_queue_warnings", ], }, ], diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.26-py3-none-any.whl b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.26-py3-none-any.whl new file mode 100644 index 00000000000..64cf55598b3 Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.26-py3-none-any.whl differ diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.26.tar.gz b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.26.tar.gz new file mode 100644 index 00000000000..8b0e817d978 Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.26.tar.gz differ diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260123131407_add_policy_tables_and_policies_field/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260123131407_add_policy_tables_and_policies_field/migration.sql new file mode 100644 index 00000000000..595d8f4a0c5 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260123131407_add_policy_tables_and_policies_field/migration.sql @@ -0,0 +1,51 @@ +-- AlterTable +ALTER TABLE "LiteLLM_DeletedTeamTable" ADD COLUMN "policies" TEXT[] DEFAULT ARRAY[]::TEXT[]; + +-- AlterTable +ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN "policies" TEXT[] DEFAULT ARRAY[]::TEXT[]; + +-- AlterTable +ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN "policies" TEXT[] DEFAULT ARRAY[]::TEXT[]; + +-- AlterTable +ALTER TABLE "LiteLLM_UserTable" ADD COLUMN "policies" TEXT[] DEFAULT ARRAY[]::TEXT[]; + +-- AlterTable +ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "policies" TEXT[] DEFAULT ARRAY[]::TEXT[]; + +-- CreateTable +CREATE TABLE "LiteLLM_PolicyTable" ( + "policy_id" TEXT NOT NULL, + "policy_name" TEXT NOT NULL, + "inherit" TEXT, + "description" TEXT, + "guardrails_add" TEXT[] DEFAULT ARRAY[]::TEXT[], + "guardrails_remove" TEXT[] DEFAULT ARRAY[]::TEXT[], + "condition" JSONB DEFAULT '{}', + "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_PolicyTable_pkey" PRIMARY KEY ("policy_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_PolicyAttachmentTable" ( + "attachment_id" TEXT NOT NULL, + "policy_name" TEXT NOT NULL, + "scope" TEXT, + "teams" TEXT[] DEFAULT ARRAY[]::TEXT[], + "keys" TEXT[] DEFAULT ARRAY[]::TEXT[], + "models" TEXT[] DEFAULT ARRAY[]::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_PolicyAttachmentTable_pkey" PRIMARY KEY ("attachment_id") +); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_PolicyTable_policy_name_key" ON "LiteLLM_PolicyTable"("policy_name"); + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 71b398c59a4..d7aa6e9f0d0 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -124,8 +124,9 @@ model LiteLLM_TeamTable { updated_at DateTime @default(now()) @updatedAt @map("updated_at") model_spend Json @default("{}") model_max_budget Json @default("{}") - router_settings Json? @default("{}") + router_settings Json? @default("{}") team_member_permissions String[] @default([]) + policies String[] @default([]) model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) @@ -156,6 +157,7 @@ model LiteLLM_DeletedTeamTable { model_max_budget Json @default("{}") router_settings Json? @default("{}") team_member_permissions String[] @default([]) + policies String[] @default([]) model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases // Original timestamps from team creation/updates @@ -197,6 +199,7 @@ model LiteLLM_UserTable { budget_duration String? budget_reset_at DateTime? allowed_cache_controls String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") created_at DateTime? @default(now()) @map("created_at") @@ -283,6 +286,7 @@ model LiteLLM_VerificationToken { budget_reset_at DateTime? allowed_cache_controls String[] @default([]) allowed_routes String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") budget_id String? @@ -327,6 +331,7 @@ model LiteLLM_DeletedVerificationToken { budget_reset_at DateTime? allowed_cache_controls String[] @default([]) allowed_routes String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") router_settings Json? @default("{}") @@ -863,3 +868,32 @@ model LiteLLM_SkillsTable { updated_at DateTime @default(now()) @updatedAt updated_by String? } + +// Policy table for storing guardrail policies +model LiteLLM_PolicyTable { + policy_id String @id @default(uuid()) + policy_name String @unique + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? +} + +// Policy attachment table for defining where policies apply +model LiteLLM_PolicyAttachmentTable { + attachment_id String @id @default(uuid()) + policy_name String // Name of the policy to attach + scope String? // Use '*' for global scope + teams String[] @default([]) // Team aliases or patterns + keys String[] @default([]) // Key aliases or patterns + models String[] @default([]) // Model names or patterns + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? +} diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index eddb563d55a..862ebaffe2a 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.25" +version = "0.4.26" 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.25" +version = "0.4.26" version_files = [ "pyproject.toml:version", "../requirements.txt:litellm-proxy-extras==", diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index c9a1531b5d4..b75e296be47 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -13,18 +13,20 @@ from litellm.types.utils import StandardLoggingPayload if TYPE_CHECKING: from opentelemetry.trace import Span +from litellm.integrations._types.open_inference import ( + MessageAttributes, + ImageAttributes, + SpanAttributes, + AudioAttributes, + EmbeddingAttributes, + OpenInferenceSpanKindValues +) class ArizeOTELAttributes(BaseLLMObsOTELAttributes): - @staticmethod @override def set_messages(span: "Span", kwargs: Dict[str, Any]): - from litellm.integrations._types.open_inference import ( - MessageAttributes, - SpanAttributes, - ) - messages = kwargs.get("messages") # for /chat/completions @@ -56,7 +58,6 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes): def set_response_output_messages(span: "Span", response_obj): """ Sets output message attributes on the span from the LLM response. - Args: span: The OpenTelemetry span to set attributes on response_obj: The response object containing choices with messages @@ -88,112 +89,243 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes): ) -def _set_tool_attributes(span: "Span", optional_params: dict): - """Helper to set tool and function call attributes on span.""" - from litellm.integrations._types.open_inference import ( - MessageAttributes, - SpanAttributes, - ToolCallAttributes, - ) - - tools = optional_params.get("tools") - if tools: - for idx, tool in enumerate(tools): - function = tool.get("function") - if not function: - continue - prefix = f"{SpanAttributes.LLM_TOOLS}.{idx}" - safe_set_attribute( - span, f"{prefix}.{SpanAttributes.TOOL_NAME}", function.get("name") - ) - safe_set_attribute( - span, - f"{prefix}.{SpanAttributes.TOOL_DESCRIPTION}", - function.get("description"), - ) - safe_set_attribute( - span, - f"{prefix}.{SpanAttributes.TOOL_PARAMETERS}", - json.dumps(function.get("parameters")), - ) - - functions = optional_params.get("functions") - if functions: - for idx, function in enumerate(functions): - prefix = f"{MessageAttributes.MESSAGE_TOOL_CALLS}.{idx}" - safe_set_attribute( - span, - f"{prefix}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}", - function.get("name"), - ) - - def _set_response_attributes(span: "Span", response_obj): """Helper to set response output and token usage attributes on span.""" - from litellm.integrations._types.open_inference import ( - MessageAttributes, - SpanAttributes, - ) if not hasattr(response_obj, "get"): return + _set_choice_outputs(span, response_obj, MessageAttributes, SpanAttributes) + _set_image_outputs(span, response_obj, ImageAttributes, SpanAttributes) + _set_audio_outputs(span, response_obj, AudioAttributes, SpanAttributes) + _set_embedding_outputs(span, response_obj, EmbeddingAttributes, SpanAttributes) + _set_structured_outputs(span, response_obj, MessageAttributes, SpanAttributes) + _set_usage_outputs(span, response_obj, SpanAttributes) + + +def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs): for idx, choice in enumerate(response_obj.get("choices", [])): response_message = choice.get("message", {}) safe_set_attribute( span, - SpanAttributes.OUTPUT_VALUE, + span_attrs.OUTPUT_VALUE, response_message.get("content", ""), ) - prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{idx}" + prefix = f"{span_attrs.LLM_OUTPUT_MESSAGES}.{idx}" safe_set_attribute( span, - f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", + f"{prefix}.{msg_attrs.MESSAGE_ROLE}", response_message.get("role"), ) safe_set_attribute( span, - f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}", + f"{prefix}.{msg_attrs.MESSAGE_CONTENT}", response_message.get("content", ""), ) - output_items = response_obj.get("output", []) - if output_items: - for i, item in enumerate(output_items): - prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{i}" - if hasattr(item, "type"): - item_type = item.type - if item_type == "reasoning" and hasattr(item, "summary"): - for summary in item.summary: - if hasattr(summary, "text"): - safe_set_attribute( - span, - f"{prefix}.{MessageAttributes.MESSAGE_REASONING_SUMMARY}", - summary.text, - ) - elif item_type == "message" and hasattr(item, "content"): - message_content = "" - content_list = item.content - if content_list and len(content_list) > 0: - first_content = content_list[0] - message_content = getattr(first_content, "text", "") - message_role = getattr(item, "role", "assistant") - safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, message_content) - safe_set_attribute(span, f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}", message_content) - safe_set_attribute(span, f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", message_role) +def _set_image_outputs(span: "Span", response_obj, image_attrs, span_attrs): + images = response_obj.get("data", []) + for i, image in enumerate(images): + img_url = image.get("url") + if img_url is None and image.get("b64_json"): + img_url = f"data:image/png;base64,{image.get('b64_json')}" + + if not img_url: + continue + + if i == 0: + safe_set_attribute(span, span_attrs.OUTPUT_VALUE, img_url) + + safe_set_attribute(span, f"{image_attrs.IMAGE_URL}.{i}", img_url) + + +def _set_audio_outputs(span: "Span", response_obj, audio_attrs, span_attrs): + audio = response_obj.get("audio", []) + for i, audio_item in enumerate(audio): + audio_url = audio_item.get("url") + if audio_url is None and audio_item.get("b64_json"): + audio_url = f"data:audio/wav;base64,{audio_item.get('b64_json')}" + + if audio_url: + if i == 0: + safe_set_attribute(span, span_attrs.OUTPUT_VALUE, audio_url) + safe_set_attribute(span, f"{audio_attrs.AUDIO_URL}.{i}", audio_url) + + audio_mime = audio_item.get("mime_type") + if audio_mime: + safe_set_attribute(span, f"{audio_attrs.AUDIO_MIME_TYPE}.{i}", audio_mime) + + audio_transcript = audio_item.get("transcript") + if audio_transcript: + safe_set_attribute(span, f"{audio_attrs.AUDIO_TRANSCRIPT}.{i}", audio_transcript) + + +def _set_embedding_outputs(span: "Span", response_obj, embedding_attrs, span_attrs): + embeddings = response_obj.get("data", []) + for i, embedding_item in enumerate(embeddings): + embedding_vector = embedding_item.get("embedding") + if embedding_vector: + if i == 0: + safe_set_attribute( + span, + span_attrs.OUTPUT_VALUE, + str(embedding_vector), + ) + + safe_set_attribute( + span, + f"{embedding_attrs.EMBEDDING_VECTOR}.{i}", + str(embedding_vector), + ) + + embedding_text = embedding_item.get("text") + if embedding_text: + safe_set_attribute( + span, + f"{embedding_attrs.EMBEDDING_TEXT}.{i}", + str(embedding_text), + ) + + +def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs): + output_items = response_obj.get("output", []) + for i, item in enumerate(output_items): + prefix = f"{span_attrs.LLM_OUTPUT_MESSAGES}.{i}" + if not hasattr(item, "type"): + continue + + item_type = item.type + if item_type == "reasoning" and hasattr(item, "summary"): + for summary in item.summary: + if hasattr(summary, "text"): + safe_set_attribute( + span, + f"{prefix}.{msg_attrs.MESSAGE_REASONING_SUMMARY}", + summary.text, + ) + elif item_type == "message" and hasattr(item, "content"): + message_content = "" + content_list = item.content + if content_list and len(content_list) > 0: + first_content = content_list[0] + message_content = getattr(first_content, "text", "") + message_role = getattr(item, "role", "assistant") + safe_set_attribute(span, span_attrs.OUTPUT_VALUE, message_content) + safe_set_attribute(span, f"{prefix}.{msg_attrs.MESSAGE_CONTENT}", message_content) + safe_set_attribute(span, f"{prefix}.{msg_attrs.MESSAGE_ROLE}", message_role) + + +def _set_usage_outputs(span: "Span", response_obj, span_attrs): usage = response_obj and response_obj.get("usage") - if usage: - safe_set_attribute(span, SpanAttributes.LLM_TOKEN_COUNT_TOTAL, usage.get("total_tokens")) - completion_tokens = usage.get("completion_tokens") or usage.get("output_tokens") - if completion_tokens: - safe_set_attribute(span, SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, completion_tokens) - prompt_tokens = usage.get("prompt_tokens") or usage.get("input_tokens") - if prompt_tokens: - safe_set_attribute(span, SpanAttributes.LLM_TOKEN_COUNT_PROMPT, prompt_tokens) - reasoning_tokens = usage.get("output_tokens_details", {}).get("reasoning_tokens") - if reasoning_tokens: - safe_set_attribute(span, SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, reasoning_tokens) + if not usage: + return + + safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_TOTAL, usage.get("total_tokens")) + completion_tokens = usage.get("completion_tokens") or usage.get("output_tokens") + if completion_tokens: + safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_COMPLETION, completion_tokens) + prompt_tokens = usage.get("prompt_tokens") or usage.get("input_tokens") + if prompt_tokens: + safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_PROMPT, prompt_tokens) + reasoning_tokens = usage.get("output_tokens_details", {}).get("reasoning_tokens") + if reasoning_tokens: + safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, reasoning_tokens) + + +def _infer_open_inference_span_kind(call_type: Optional[str]) -> str: + """ + Map LiteLLM call types to OpenInference span kinds. + """ + + if not call_type: + return OpenInferenceSpanKindValues.UNKNOWN.value + + lowered = str(call_type).lower() + + if "embed" in lowered: + return OpenInferenceSpanKindValues.EMBEDDING.value + + if "rerank" in lowered: + return OpenInferenceSpanKindValues.RERANKER.value + + if "search" in lowered: + return OpenInferenceSpanKindValues.RETRIEVER.value + + if "moderation" in lowered or "guardrail" in lowered: + return OpenInferenceSpanKindValues.GUARDRAIL.value + + if lowered == "call_mcp_tool" or lowered == "mcp" or lowered.endswith("tool"): + return OpenInferenceSpanKindValues.TOOL.value + + if "asend_message" in lowered or "a2a" in lowered or "assistant" in lowered: + return OpenInferenceSpanKindValues.AGENT.value + + if any( + keyword in lowered + for keyword in ( + "completion", + "chat", + "image", + "audio", + "speech", + "transcription", + "generate_content", + "response", + "videos", + "realtime", + "pass_through", + "anthropic_messages", + "ocr", + ) + ): + return OpenInferenceSpanKindValues.LLM.value + + if any(keyword in lowered for keyword in ("file", "batch", "container", "fine_tuning_job")): + return OpenInferenceSpanKindValues.CHAIN.value + + return OpenInferenceSpanKindValues.UNKNOWN.value + +def _set_tool_attributes( + span: "Span", optional_tools: Optional[list], metadata_tools: Optional[list] +): + """set tool attributes on span from optional_params or tool call metadata""" + if optional_tools: + for idx, tool in enumerate(optional_tools): + if not isinstance(tool, dict): + continue + function = tool.get("function") if isinstance(tool.get("function"), dict) else None + if not function: + continue + tool_name = function.get("name") + if tool_name: + safe_set_attribute(span, f"{SpanAttributes.LLM_TOOLS}.{idx}.name", tool_name) + tool_description = function.get("description") + if tool_description: + safe_set_attribute(span, f"{SpanAttributes.LLM_TOOLS}.{idx}.description", tool_description) + params = function.get("parameters") + if params is not None: + safe_set_attribute(span, f"{SpanAttributes.LLM_TOOLS}.{idx}.parameters", json.dumps(params)) + + if metadata_tools and isinstance(metadata_tools, list): + for idx, tool in enumerate(metadata_tools): + if not isinstance(tool, dict): + continue + tool_name = tool.get("name") + if tool_name: + safe_set_attribute( + span, + f"{SpanAttributes.LLM_INVOCATION_PARAMETERS}.tools.{idx}.name", + tool_name, + ) + + tool_description = tool.get("description") + if tool_description: + safe_set_attribute( + span, + f"{SpanAttributes.LLM_INVOCATION_PARAMETERS}.tools.{idx}.description", + tool_description, + ) def set_attributes( @@ -202,70 +334,42 @@ def set_attributes( """ Populates span with OpenInference-compliant LLM attributes for Arize and Phoenix tracing. """ - from litellm.integrations._types.open_inference import ( - OpenInferenceSpanKindValues, - SpanAttributes, - ) - try: - # Remove secret_fields to prevent leaking sensitive data (e.g., authorization headers) - optional_params = kwargs.get("optional_params", {}) - if isinstance(optional_params, dict): - optional_params.pop("secret_fields", None) - litellm_params = kwargs.get("litellm_params", {}) + optional_params = _sanitize_optional_params(kwargs.get("optional_params")) + litellm_params = kwargs.get("litellm_params", {}) or {} standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get( "standard_logging_object" ) if standard_logging_payload is None: raise ValueError("standard_logging_object not found in kwargs") - metadata = ( - standard_logging_payload.get("metadata") - if standard_logging_payload - else None + metadata = standard_logging_payload.get("metadata") if standard_logging_payload else None + _set_metadata_attributes(span, metadata, SpanAttributes) + + metadata_tools = _extract_metadata_tools(metadata) + optional_tools = _extract_optional_tools(optional_params) + + call_type = standard_logging_payload.get("call_type") + _set_request_attributes( + span=span, + kwargs=kwargs, + standard_logging_payload=standard_logging_payload, + optional_params=optional_params, + litellm_params=litellm_params, + response_obj=response_obj, + span_attrs=SpanAttributes, ) - if metadata is not None: - safe_set_attribute(span, SpanAttributes.METADATA, safe_dumps(metadata)) - if kwargs.get("model"): - safe_set_attribute(span, SpanAttributes.LLM_MODEL_NAME, kwargs.get("model")) + span_kind = _infer_open_inference_span_kind(call_type=call_type) + _set_tool_attributes(span, optional_tools, metadata_tools) + if (optional_tools or metadata_tools) and span_kind != OpenInferenceSpanKindValues.TOOL.value: + span_kind = OpenInferenceSpanKindValues.TOOL.value - safe_set_attribute(span, "llm.request.type", standard_logging_payload["call_type"]) - safe_set_attribute(span, SpanAttributes.LLM_PROVIDER, litellm_params.get("custom_llm_provider", "Unknown")) - - if optional_params.get("max_tokens"): - safe_set_attribute(span, "llm.request.max_tokens", optional_params.get("max_tokens")) - if optional_params.get("temperature"): - safe_set_attribute(span, "llm.request.temperature", optional_params.get("temperature")) - if optional_params.get("top_p"): - safe_set_attribute(span, "llm.request.top_p", optional_params.get("top_p")) - - safe_set_attribute(span, "llm.is_streaming", str(optional_params.get("stream", False))) - - if optional_params.get("user"): - safe_set_attribute(span, "llm.user", optional_params.get("user")) - - if response_obj and response_obj.get("id"): - safe_set_attribute(span, "llm.response.id", response_obj.get("id")) - if response_obj and response_obj.get("model"): - safe_set_attribute(span, "llm.response.model", response_obj.get("model")) - - safe_set_attribute(span, SpanAttributes.OPENINFERENCE_SPAN_KIND, OpenInferenceSpanKindValues.LLM.value) + safe_set_attribute(span, SpanAttributes.OPENINFERENCE_SPAN_KIND, span_kind) attributes.set_messages(span, kwargs) - _set_tool_attributes(span=span, optional_params=optional_params) - - model_params = ( - standard_logging_payload.get("model_parameters") - if standard_logging_payload - else None - ) - if model_params: - safe_set_attribute(span, SpanAttributes.LLM_INVOCATION_PARAMETERS, safe_dumps(model_params)) - if model_params.get("user"): - user_id = model_params.get("user") - if user_id is not None: - safe_set_attribute(span, SpanAttributes.USER_ID, user_id) + model_params = standard_logging_payload.get("model_parameters") if standard_logging_payload else None + _set_model_params(span, model_params, SpanAttributes) _set_response_attributes(span=span, response_obj=response_obj) @@ -275,3 +379,72 @@ def set_attributes( ) if hasattr(span, "record_exception"): span.record_exception(e) + + +def _sanitize_optional_params(optional_params: Optional[dict]) -> dict: + if not isinstance(optional_params, dict): + return {} + optional_params.pop("secret_fields", None) + return optional_params + + +def _set_metadata_attributes(span: "Span", metadata: Optional[Any], span_attrs) -> None: + if metadata is not None: + safe_set_attribute(span, span_attrs.METADATA, safe_dumps(metadata)) + + +def _extract_metadata_tools(metadata: Optional[Any]) -> Optional[list]: + if not isinstance(metadata, dict): + return None + llm_obj = metadata.get("llm") + if isinstance(llm_obj, dict): + return llm_obj.get("tools") + return None + + +def _extract_optional_tools(optional_params: dict) -> Optional[list]: + return optional_params.get("tools") if isinstance(optional_params, dict) else None + + +def _set_request_attributes( + span: "Span", + kwargs, + standard_logging_payload: StandardLoggingPayload, + optional_params: dict, + litellm_params: dict, + response_obj, + span_attrs, +): + if kwargs.get("model"): + safe_set_attribute(span, span_attrs.LLM_MODEL_NAME, kwargs.get("model")) + + safe_set_attribute(span, "llm.request.type", standard_logging_payload.get("call_type")) + safe_set_attribute(span, span_attrs.LLM_PROVIDER, litellm_params.get("custom_llm_provider", "Unknown")) + + if optional_params.get("max_tokens"): + safe_set_attribute(span, "llm.request.max_tokens", optional_params.get("max_tokens")) + if optional_params.get("temperature"): + safe_set_attribute(span, "llm.request.temperature", optional_params.get("temperature")) + if optional_params.get("top_p"): + safe_set_attribute(span, "llm.request.top_p", optional_params.get("top_p")) + + safe_set_attribute(span, "llm.is_streaming", str(optional_params.get("stream", False))) + + if optional_params.get("user"): + safe_set_attribute(span, "llm.user", optional_params.get("user")) + + if response_obj and response_obj.get("id"): + safe_set_attribute(span, "llm.response.id", response_obj.get("id")) + if response_obj and response_obj.get("model"): + safe_set_attribute(span, "llm.response.model", response_obj.get("model")) + + +def _set_model_params(span: "Span", model_params: Optional[dict], span_attrs) -> None: + if not model_params: + return + + safe_set_attribute(span, span_attrs.LLM_INVOCATION_PARAMETERS, safe_dumps(model_params)) + if model_params.get("user"): + user_id = model_params.get("user") + if user_id is not None: + safe_set_attribute(span, span_attrs.USER_ID, user_id) diff --git a/litellm/integrations/gcs_bucket/gcs_bucket_base.py b/litellm/integrations/gcs_bucket/gcs_bucket_base.py index 2612face050..b1db9ec9588 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket_base.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket_base.py @@ -2,6 +2,13 @@ import json import os from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union +from litellm.integrations.gcs_bucket.gcs_bucket_mock_client import ( + should_use_gcs_mock, + create_mock_gcs_client, + mock_vertex_auth_methods, +) + + from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.llms.custom_httpx.http_handler import ( @@ -20,6 +27,12 @@ IAM_AUTH_KEY = "IAM_AUTH" class GCSBucketBase(CustomBatchLogger): def __init__(self, bucket_name: Optional[str] = None, **kwargs) -> None: + self.is_mock_mode = should_use_gcs_mock() + + if self.is_mock_mode: + mock_vertex_auth_methods() + create_mock_gcs_client() + self.async_httpx_client = get_async_httpx_client( llm_provider=httpxSpecialProvider.LoggingCallback ) diff --git a/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py b/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py new file mode 100644 index 00000000000..6201dc343dc --- /dev/null +++ b/litellm/integrations/gcs_bucket/gcs_bucket_mock_client.py @@ -0,0 +1,236 @@ +""" +Mock client for GCS Bucket integration testing. + +This module intercepts GCS API calls and Vertex AI auth calls, returning successful +mock responses, allowing full code execution without making actual network calls. + +Usage: + Set GCS_MOCK=true in environment variables or config to enable mock mode. +""" + +import httpx +import json +import asyncio +from datetime import timedelta +from typing import Dict, Optional + +from litellm._logging import verbose_logger + +# Store original methods for restoration +_original_async_handler_post = None +_original_async_handler_get = None +_original_async_handler_delete = None + +# Track if mocks have been initialized to avoid duplicate initialization +_mocks_initialized = False + +# Default mock latency in seconds (simulates network round-trip) +# Typical GCS API calls take 100-300ms for uploads, 50-150ms for GET/DELETE +_MOCK_LATENCY_SECONDS = float(__import__("os").getenv("GCS_MOCK_LATENCY_MS", "150")) / 1000.0 + + +class MockGCSResponse: + """Mock httpx.Response that satisfies GCS API requirements.""" + + def __init__(self, status_code: int = 200, json_data: Optional[Dict] = None, url: Optional[str] = None, elapsed_seconds: float = 0.0): + self.status_code = status_code + self._json_data = json_data or {"kind": "storage#object", "name": "mock-object"} + self.headers = httpx.Headers({}) + self.is_success = status_code < 400 + self.is_error = status_code >= 400 + self.is_redirect = 300 <= status_code < 400 + self.url = httpx.URL(url) if url else httpx.URL("") + # Set realistic elapsed time based on mock latency + elapsed_time = elapsed_seconds if elapsed_seconds > 0 else _MOCK_LATENCY_SECONDS + self.elapsed = timedelta(seconds=elapsed_time) + self._text = json.dumps(self._json_data) + self._content = self._text.encode("utf-8") + + @property + def text(self) -> str: + """Return response text.""" + return self._text + + @property + def content(self) -> bytes: + """Return response content.""" + return self._content + + def json(self) -> Dict: + """Return JSON response data.""" + return self._json_data + + def read(self) -> bytes: + """Read response content.""" + return self._content + + def raise_for_status(self): + """Raise exception for error status codes.""" + if self.status_code >= 400: + raise Exception(f"HTTP {self.status_code}") + + +async def _mock_async_handler_post(self, url, data=None, json=None, params=None, headers=None, timeout=None, stream=False, logging_obj=None, files=None, content=None): + """Monkey-patched AsyncHTTPHandler.post that intercepts GCS calls.""" + # Only mock GCS API calls + if isinstance(url, str) and "storage.googleapis.com" in url: + verbose_logger.info(f"[GCS MOCK] POST to {url}") + # Simulate network latency + await asyncio.sleep(_MOCK_LATENCY_SECONDS) + return MockGCSResponse( + status_code=200, + json_data={"kind": "storage#object", "name": "mock-object"}, + url=url, + elapsed_seconds=_MOCK_LATENCY_SECONDS + ) + # For non-GCS calls, use original method + if _original_async_handler_post is not None: + return await _original_async_handler_post(self, url=url, data=data, json=json, params=params, headers=headers, timeout=timeout, stream=stream, logging_obj=logging_obj, files=files, content=content) + # Fallback: if original not set, raise error + raise RuntimeError("Original AsyncHTTPHandler.post not available") + + +async def _mock_async_handler_get(self, url, params=None, headers=None, follow_redirects=None): + """Monkey-patched AsyncHTTPHandler.get that intercepts GCS calls.""" + # Only mock GCS API calls + if isinstance(url, str) and "storage.googleapis.com" in url: + verbose_logger.info(f"[GCS MOCK] GET to {url}") + # Simulate network latency + await asyncio.sleep(_MOCK_LATENCY_SECONDS) + return MockGCSResponse( + status_code=200, + json_data={"data": "mock-log-data"}, + url=url, + elapsed_seconds=_MOCK_LATENCY_SECONDS + ) + # For non-GCS calls, use original method + if _original_async_handler_get is not None: + return await _original_async_handler_get(self, url=url, params=params, headers=headers, follow_redirects=follow_redirects) + # Fallback: if original not set, raise error + raise RuntimeError("Original AsyncHTTPHandler.get not available") + + +async def _mock_async_handler_delete(self, url, data=None, json=None, params=None, headers=None, timeout=None, stream=False, content=None): + """Monkey-patched AsyncHTTPHandler.delete that intercepts GCS calls.""" + # Only mock GCS API calls + if isinstance(url, str) and "storage.googleapis.com" in url: + verbose_logger.info(f"[GCS MOCK] DELETE to {url}") + # Simulate network latency + await asyncio.sleep(_MOCK_LATENCY_SECONDS) + return MockGCSResponse( + status_code=204, + json_data={}, + url=url, + elapsed_seconds=_MOCK_LATENCY_SECONDS + ) + # For non-GCS calls, use original method + if _original_async_handler_delete is not None: + return await _original_async_handler_delete(self, url=url, data=data, json=json, params=params, headers=headers, timeout=timeout, stream=stream, content=content) + # Fallback: if original not set, raise error + raise RuntimeError("Original AsyncHTTPHandler.delete not available") + + +def create_mock_gcs_client(): + """ + Monkey-patch AsyncHTTPHandler methods to intercept GCS calls. + + AsyncHTTPHandler is used by LiteLLM's get_async_httpx_client() which is what + GCSBucketBase uses for making API calls. + + This function is idempotent - it only initializes mocks once, even if called multiple times. + """ + global _original_async_handler_post, _original_async_handler_get, _original_async_handler_delete + global _mocks_initialized + + # If already initialized, skip + if _mocks_initialized: + return + + verbose_logger.debug("[GCS MOCK] Initializing GCS mock client...") + + # Patch AsyncHTTPHandler methods (used by LiteLLM's custom httpx handler) + if _original_async_handler_post is None: + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + _original_async_handler_post = AsyncHTTPHandler.post + AsyncHTTPHandler.post = _mock_async_handler_post # type: ignore + verbose_logger.debug("[GCS MOCK] Patched AsyncHTTPHandler.post") + + if _original_async_handler_get is None: + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + _original_async_handler_get = AsyncHTTPHandler.get + AsyncHTTPHandler.get = _mock_async_handler_get # type: ignore + verbose_logger.debug("[GCS MOCK] Patched AsyncHTTPHandler.get") + + if _original_async_handler_delete is None: + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + _original_async_handler_delete = AsyncHTTPHandler.delete + AsyncHTTPHandler.delete = _mock_async_handler_delete # type: ignore + verbose_logger.debug("[GCS MOCK] Patched AsyncHTTPHandler.delete") + + verbose_logger.debug(f"[GCS MOCK] Mock latency set to {_MOCK_LATENCY_SECONDS*1000:.0f}ms") + verbose_logger.debug("[GCS MOCK] GCS mock client initialization complete") + + _mocks_initialized = True + + +def mock_vertex_auth_methods(): + """ + Monkey-patch Vertex AI auth methods to return fake tokens. + This prevents auth failures when GCS_MOCK is enabled. + + This function is idempotent - it only patches once, even if called multiple times. + """ + from litellm.llms.vertex_ai.vertex_llm_base import VertexBase + + # Store original methods if not already stored + if not hasattr(VertexBase, '_original_ensure_access_token_async'): + setattr(VertexBase, '_original_ensure_access_token_async', VertexBase._ensure_access_token_async) + setattr(VertexBase, '_original_ensure_access_token', VertexBase._ensure_access_token) + setattr(VertexBase, '_original_get_token_and_url', VertexBase._get_token_and_url) + + async def _mock_ensure_access_token_async(self, credentials, project_id, custom_llm_provider): + """Mock async auth method - returns fake token.""" + verbose_logger.debug("[GCS MOCK] Vertex AI auth: _ensure_access_token_async called") + return ("mock-gcs-token", "mock-project-id") + + def _mock_ensure_access_token(self, credentials, project_id, custom_llm_provider): + """Mock sync auth method - returns fake token.""" + verbose_logger.debug("[GCS MOCK] Vertex AI auth: _ensure_access_token called") + return ("mock-gcs-token", "mock-project-id") + + def _mock_get_token_and_url(self, model, auth_header, vertex_credentials, vertex_project, + vertex_location, gemini_api_key, stream, custom_llm_provider, api_base): + """Mock get_token_and_url - returns fake token.""" + verbose_logger.debug("[GCS MOCK] Vertex AI auth: _get_token_and_url called") + return ("mock-gcs-token", "https://storage.googleapis.com") + + # Patch the methods + VertexBase._ensure_access_token_async = _mock_ensure_access_token_async # type: ignore + VertexBase._ensure_access_token = _mock_ensure_access_token # type: ignore + VertexBase._get_token_and_url = _mock_get_token_and_url # type: ignore + + verbose_logger.debug("[GCS MOCK] Patched Vertex AI auth methods") + + +def should_use_gcs_mock() -> bool: + """ + Determine if GCS should run in mock mode. + + Checks the GCS_MOCK environment variable. + + Returns: + bool: True if mock mode should be enabled + """ + import os + from litellm.secret_managers.main import str_to_bool + + mock_mode = os.getenv("GCS_MOCK", "false") + result = str_to_bool(mock_mode) + + # Ensure we return a bool, not None + result = bool(result) if result is not None else False + + if result: + verbose_logger.info("GCS Mock Mode: ENABLED - API calls will be mocked") + + return result diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 8087c17cafe..46ada3c3930 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -25,6 +25,10 @@ from litellm.litellm_core_utils.core_helpers import ( reconstruct_model_name, ) from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info +from litellm.integrations.langfuse.langfuse_mock_client import ( + create_mock_langfuse_client, + should_use_langfuse_mock, +) from litellm.llms.custom_httpx.http_handler import _get_httpx_client from litellm.secret_managers.main import str_to_bool from litellm.types.integrations.langfuse import * @@ -119,8 +123,14 @@ class LangFuseLogger: self.langfuse_flush_interval = LangFuseLogger._get_langfuse_flush_interval( flush_interval ) - http_client = _get_httpx_client() - self.langfuse_client = http_client.client + + if should_use_langfuse_mock(): + self.langfuse_client = create_mock_langfuse_client() + self.is_mock_mode = True + else: + http_client = _get_httpx_client() + self.langfuse_client = http_client.client + self.is_mock_mode = False parameters = { "public_key": self.public_key, @@ -139,11 +149,15 @@ class LangFuseLogger: # set the current langfuse project id in the environ # this is used by Alerting to link to the correct project - try: - project_id = self.Langfuse.client.projects.get().data[0].id - os.environ["LANGFUSE_PROJECT_ID"] = project_id - except Exception: - project_id = None + if self.is_mock_mode: + os.environ["LANGFUSE_PROJECT_ID"] = "mock-project-id" + verbose_logger.debug("Langfuse Mock: Using mock project ID") + else: + try: + project_id = self.Langfuse.client.projects.get().data[0].id + os.environ["LANGFUSE_PROJECT_ID"] = project_id + except Exception: + project_id = None if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None: upstream_langfuse_debug = ( diff --git a/litellm/integrations/langfuse/langfuse_mock_client.py b/litellm/integrations/langfuse/langfuse_mock_client.py new file mode 100644 index 00000000000..1dc739ea328 --- /dev/null +++ b/litellm/integrations/langfuse/langfuse_mock_client.py @@ -0,0 +1,121 @@ +""" +Mock httpx client for Langfuse integration testing. + +This module intercepts Langfuse API calls and returns successful mock responses, +allowing full code execution without making actual network calls. + +Usage: + Set LANGFUSE_MOCK=true in environment variables or config to enable mock mode. +""" + +import httpx +import json +from datetime import timedelta +from typing import Dict, Optional + +from litellm._logging import verbose_logger + +_original_httpx_post = None + +# Default mock latency in seconds (simulates network round-trip) +# Typical Langfuse API calls take 50-150ms +_MOCK_LATENCY_SECONDS = float(__import__("os").getenv("LANGFUSE_MOCK_LATENCY_MS", "100")) / 1000.0 + + +class MockLangfuseResponse: + """Mock httpx.Response that satisfies Langfuse SDK requirements.""" + + def __init__(self, status_code: int = 200, json_data: Optional[Dict] = None, url: Optional[str] = None, elapsed_seconds: float = 0.0): + self.status_code = status_code + self._json_data = json_data or {"status": "success"} + self.headers = httpx.Headers({}) + self.is_success = status_code < 400 + self.is_error = status_code >= 400 + self.is_redirect = 300 <= status_code < 400 + self.url = httpx.URL(url) if url else httpx.URL("") + # Set realistic elapsed time based on mock latency + elapsed_time = elapsed_seconds if elapsed_seconds > 0 else _MOCK_LATENCY_SECONDS + self.elapsed = timedelta(seconds=elapsed_time) + self._text = json.dumps(self._json_data) + self._content = self._text.encode("utf-8") + + @property + def text(self) -> str: + return self._text + + @property + def content(self) -> bytes: + return self._content + + def json(self) -> Dict: + return self._json_data + + def read(self) -> bytes: + return self._content + + def raise_for_status(self): + if self.status_code >= 400: + raise Exception(f"HTTP {self.status_code}") + + +def _is_langfuse_url(url) -> bool: + """Check if URL is a Langfuse domain.""" + try: + parsed_url = httpx.URL(url) if isinstance(url, str) else url + hostname = parsed_url.host or "" + + return ( + hostname.endswith(".langfuse.com") or + hostname == "langfuse.com" or + (hostname in ("localhost", "127.0.0.1") and "langfuse" in str(parsed_url).lower()) + ) + except Exception: + return False + + +def _mock_httpx_post(self, url, **kwargs): + """Monkey-patched httpx.Client.post that intercepts Langfuse calls.""" + if _is_langfuse_url(url): + verbose_logger.info(f"[LANGFUSE MOCK] POST to {url}") + return MockLangfuseResponse(status_code=200, json_data={"status": "success"}, url=url, elapsed_seconds=_MOCK_LATENCY_SECONDS) + + if _original_httpx_post is not None: + return _original_httpx_post(self, url, **kwargs) + + +def create_mock_langfuse_client(): + """ + Monkey-patch httpx.Client.post to intercept Langfuse calls. + + Returns a real httpx.Client instance - the monkey-patch intercepts all calls. + """ + global _original_httpx_post + + if _original_httpx_post is None: + _original_httpx_post = httpx.Client.post + httpx.Client.post = _mock_httpx_post # type: ignore + verbose_logger.debug("[LANGFUSE MOCK] Patched httpx.Client.post") + + return httpx.Client() + + +def should_use_langfuse_mock() -> bool: + """ + Determine if Langfuse should run in mock mode. + + Checks the LANGFUSE_MOCK environment variable. + + Returns: + bool: True if mock mode should be enabled + """ + import os + from litellm.secret_managers.main import str_to_bool + + mock_mode = os.getenv("LANGFUSE_MOCK", "false") + result = str_to_bool(mock_mode) + result = bool(result) if result is not None else False + + if result: + verbose_logger.info("Langfuse Mock Mode: ENABLED - API calls will be mocked") + + return result diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 93d631eb0f2..a8a7fa77b3d 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -17,6 +17,10 @@ from litellm.types.utils import ( StandardCallbackDynamicParams, StandardLoggingPayload, ) +from litellm.integrations._types.open_inference import ( + OpenInferenceSpanKindValues, + SpanAttributes, +) # OpenTelemetry imports moved to individual functions to avoid import errors when not installed @@ -660,6 +664,9 @@ class OpenTelemetry(CustomLogger): self._maybe_log_raw_request( kwargs, response_obj, start_time, end_time, span ) + # Ensure proxy-request parent span is annotated with the actual operation kind + if parent_span is not None and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME: + self.set_attributes(parent_span, kwargs, response_obj) else: # Do not create primary span (keep hierarchy shallow when parent exists) from opentelemetry.trace import Status, StatusCode @@ -1106,6 +1113,12 @@ class OpenTelemetry(CustomLogger): context=context, ) + self.safe_set_attribute( + span=guardrail_span, + key=SpanAttributes.OPENINFERENCE_SPAN_KIND, + value=OpenInferenceSpanKindValues.GUARDRAIL.value, + ) + self.safe_set_attribute( span=guardrail_span, key="guardrail_name", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2fd03ac2128..0516a4aaa66 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -73,6 +73,7 @@ class SupportedDBObjectType(str, enum.Enum): MODELS = "models" MCP = "mcp" GUARDRAILS = "guardrails" + POLICIES = "policies" VECTOR_STORES = "vector_stores" PASS_THROUGH_ENDPOINTS = "pass_through_endpoints" PROMPTS = "prompts" @@ -844,6 +845,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): model_rpm_limit: Optional[dict] = None model_tpm_limit: Optional[dict] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None prompts: Optional[List[str]] = None blocked: Optional[bool] = None aliases: Optional[dict] = {} @@ -1477,6 +1479,7 @@ class NewTeamRequest(TeamBase): model_aliases: Optional[dict] = None tags: Optional[list] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None prompts: Optional[List[str]] = None object_permission: Optional[LiteLLM_ObjectPermissionBase] = None allowed_passthrough_routes: Optional[list] = None @@ -1526,6 +1529,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): blocked: Optional[bool] = None budget_duration: Optional[str] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None """ team_id: str # required @@ -1541,6 +1545,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): tags: Optional[list] = None model_aliases: Optional[dict] = None guardrails: Optional[List[str]] = None + policies: Optional[List[str]] = None object_permission: Optional[LiteLLM_ObjectPermissionBase] = None team_member_budget: Optional[float] = None team_member_budget_duration: Optional[str] = None @@ -3499,6 +3504,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [ LiteLLM_ManagementEndpoint_MetadataFields_Premium = [ "guardrails", + "policies", "tags", "team_member_key_duration", "prompts", diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 32cddc0ef58..1d3ef2e10c2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1311,6 +1311,118 @@ def _add_guardrails_from_key_or_team_metadata( data[metadata_variable_name]["guardrails"] = list(combined_guardrails) +def _add_guardrails_from_policies_in_metadata( + key_metadata: Optional[dict], + team_metadata: Optional[dict], + data: dict, + metadata_variable_name: str, +) -> None: + """ + Helper to resolve guardrails from policies attached to key/team metadata. + + This function: + 1. Gets policy names from key and team metadata + 2. Resolves guardrails from those policies (including inheritance) + 3. Adds resolved guardrails to request metadata + + Args: + key_metadata: The key metadata dictionary to check for policies + team_metadata: The team metadata dictionary to check for policies + data: The request data to update + metadata_variable_name: The name of the metadata field in data + """ + from litellm._logging import verbose_proxy_logger + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_resolver import PolicyResolver + from litellm.proxy.utils import _premium_user_check + from litellm.types.proxy.policy_engine import PolicyMatchContext + + # Collect policy names from key and team metadata + policy_names: set = set() + + # Add key-level policies first + if key_metadata and "policies" in key_metadata: + if ( + isinstance(key_metadata["policies"], list) + and len(key_metadata["policies"]) > 0 + ): + _premium_user_check() + policy_names.update(key_metadata["policies"]) + + # Add team-level policies + if team_metadata and "policies" in team_metadata: + if ( + isinstance(team_metadata["policies"], list) + and len(team_metadata["policies"]) > 0 + ): + _premium_user_check() + policy_names.update(team_metadata["policies"]) + + if not policy_names: + return + + verbose_proxy_logger.debug( + f"Policy engine: resolving guardrails from key/team policies: {policy_names}" + ) + + # Check if policy registry is initialized + registry = get_policy_registry() + if not registry.is_initialized(): + verbose_proxy_logger.debug( + "Policy engine not initialized, skipping policy resolution from metadata" + ) + return + + # Build context for policy resolution (model from request data) + context = PolicyMatchContext(model=data.get("model")) + + # Get all policies from registry + all_policies = registry.get_all_policies() + + # Resolve guardrails from the specified policies + resolved_guardrails: set = set() + for policy_name in policy_names: + if registry.has_policy(policy_name): + resolved_policy = PolicyResolver.resolve_policy_guardrails( + policy_name=policy_name, + policies=all_policies, + context=context, + ) + resolved_guardrails.update(resolved_policy.guardrails) + verbose_proxy_logger.debug( + f"Policy engine: resolved guardrails from policy '{policy_name}': {resolved_policy.guardrails}" + ) + else: + verbose_proxy_logger.warning( + f"Policy engine: policy '{policy_name}' not found in registry" + ) + + if not resolved_guardrails: + return + + # Add resolved guardrails to request metadata + if metadata_variable_name not in data: + data[metadata_variable_name] = {} + + existing_guardrails = data[metadata_variable_name].get("guardrails", []) + if not isinstance(existing_guardrails, list): + existing_guardrails = [] + + # Combine existing guardrails with policy-resolved guardrails (no duplicates) + combined = set(existing_guardrails) + combined.update(resolved_guardrails) + data[metadata_variable_name]["guardrails"] = list(combined) + + # Store applied policies in metadata for tracking + if "applied_policies" not in data[metadata_variable_name]: + data[metadata_variable_name]["applied_policies"] = [] + data[metadata_variable_name]["applied_policies"].extend(list(policy_names)) + + verbose_proxy_logger.debug( + f"Policy engine: added guardrails from key/team policies to request metadata: {list(resolved_guardrails)}" + ) + + def move_guardrails_to_metadata( data: dict, _metadata_variable_name: str, @@ -1321,6 +1433,7 @@ def move_guardrails_to_metadata( - If guardrails set on API Key metadata then sets guardrails on request metadata - If guardrails not set on API key, then checks request metadata + - Adds guardrails from policies attached to key/team metadata - Adds guardrails from policy engine based on team/key/model context """ # Check key-level guardrails @@ -1331,6 +1444,16 @@ def move_guardrails_to_metadata( metadata_variable_name=_metadata_variable_name, ) + ######################################################################################### + # Add guardrails from policies attached to key/team metadata + ######################################################################################### + _add_guardrails_from_policies_in_metadata( + key_metadata=user_api_key_dict.metadata, + team_metadata=user_api_key_dict.team_metadata, + data=data, + metadata_variable_name=_metadata_variable_name, + ) + ######################################################################################### # Add guardrails from policy engine based on team/key/model context ######################################################################################### diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 2672c41893d..9a986326c03 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -14,7 +14,6 @@ These are members of a Team on LiteLLM import asyncio import json import traceback -from litellm._uuid import uuid from datetime import datetime, timezone from typing import Any, Dict, List, Optional, Union, cast @@ -23,6 +22,7 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Request, status import litellm from litellm._logging import verbose_proxy_logger +from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks @@ -355,6 +355,7 @@ async def new_user( - allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request- - blocked: Optional[bool] - [Not Implemented Yet] Whether the user is blocked. - guardrails: Optional[List[str]] - [Not Implemented Yet] List of active guardrails for the user + - policies: Optional[List[str]] - List of policy names to apply to the user. Policies define guardrails, conditions, and inheritance rules. - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. @@ -1060,6 +1061,7 @@ async def user_update( - allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request- - blocked: Optional[bool] - [Not Implemented Yet] Whether the user is blocked. - guardrails: Optional[List[str]] - [Not Implemented Yet] List of active guardrails for the user + - policies: Optional[List[str]] - List of policy names to apply to the user. Policies define guardrails, conditions, and inheritance rules. - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e40a44edf5c..ab87e862ea6 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -14,11 +14,11 @@ import copy import json import secrets import traceback -import yaml from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Literal, Optional, Tuple, cast -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + import fastapi +import yaml from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status import litellm @@ -31,6 +31,7 @@ from litellm.constants import ( UI_SESSION_TOKEN_TEAM_ID, ) from litellm.litellm_core_utils.duration_parser import duration_in_seconds +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._experimental.mcp_server.db import ( rotate_mcp_server_credentials_master_key, ) @@ -1010,6 +1011,7 @@ async def generate_key_fn( - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - metadata: Optional[dict] - Metadata for key, store information for key. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } - guardrails: Optional[List[str]] - List of active guardrails for the key + - policies: Optional[List[str]] - List of policy names to apply to the key. Policies define guardrails, conditions, and inheritance rules. - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. - permissions: Optional[dict] - key-specific permissions. Currently just used for turning off pii masking (if connected). Example - {"pii": false} - model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget. @@ -1480,6 +1482,7 @@ async def update_key_fn( - permissions: Optional[dict] - Key-specific permissions - send_invite_email: Optional[bool] - Send invite email to user_id - guardrails: Optional[List[str]] - List of active guardrails for the key + - policies: Optional[List[str]] - List of policy names to apply to the key. Policies define guardrails, conditions, and inheritance rules. - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. - prompts: Optional[List[str]] - List of prompts that the key is allowed to use. - blocked: Optional[bool] - Whether the key is blocked @@ -2077,6 +2080,7 @@ async def generate_key_helper_fn( # noqa: PLR0915 model_rpm_limit: Optional[dict] = None, model_tpm_limit: Optional[dict] = None, guardrails: Optional[list] = None, + policies: Optional[list] = None, prompts: Optional[list] = None, teams: Optional[list] = None, organization_id: Optional[str] = None, @@ -2139,6 +2143,9 @@ async def generate_key_helper_fn( # noqa: PLR0915 if guardrails is not None: metadata = metadata or {} metadata["guardrails"] = guardrails + if policies is not None: + metadata = metadata or {} + metadata["policies"] = policies if prompts is not None: metadata = metadata or {} metadata["prompts"] = prompts diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 4d313fb1235..c77b60649ab 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -22,11 +22,13 @@ from pydantic import BaseModel import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ( BlockTeamRequest, CommonProxyErrors, DeleteTeamRequest, LiteLLM_AuditLogs, + LiteLLM_DeletedTeamTable, LiteLLM_ManagementEndpoint_MetadataFields, LiteLLM_ManagementEndpoint_MetadataFields_Premium, LiteLLM_ModelTable, @@ -34,7 +36,6 @@ from litellm.proxy._types import ( LiteLLM_OrganizationTableWithMembers, LiteLLM_TeamMembership, LiteLLM_TeamTable, - LiteLLM_DeletedTeamTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LiteLLM_VerificationToken, @@ -102,7 +103,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( TeamMemberAddResult, UpdateTeamMemberPermissionsRequest, ) -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + router = APIRouter() @@ -689,6 +690,7 @@ async def new_team( # noqa: PLR0915 - organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`. - model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias) - guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails) + - policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies) - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. - object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission. - team_member_budget: Optional[float] - The maximum budget allocated to an individual team member. @@ -1228,6 +1230,7 @@ async def update_team( # noqa: PLR0915 - organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`. - model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias) - guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails) + - policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies) - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. - object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission. - team_member_budget: Optional[float] - The maximum budget allocated to an individual team member. diff --git a/litellm/proxy/pass_through_endpoints/architecture.md b/litellm/proxy/pass_through_endpoints/architecture.md index 064a443a2e7..f7dd8077ab5 100644 --- a/litellm/proxy/pass_through_endpoints/architecture.md +++ b/litellm/proxy/pass_through_endpoints/architecture.md @@ -16,23 +16,33 @@ sequenceDiagram rect rgb(240, 240, 240) Note over Proxy: 1. URL Construction - Note over Proxy: Build: https://us-central1-aiplatform.googleapis.com/... + Note over Proxy: Build regional/provider-specific URL end rect rgb(240, 240, 240) Note over Proxy: 2. Auth Header Replacement - Note over Proxy: Replace litellm key → provider credentials + Note over Proxy: LiteLLM key → provider credentials + end + + rect rgb(240, 240, 240) + Note over Proxy: 3. Extra Operations + Note over Proxy: • x-pass-* headers (strip prefix, forward) + Note over Proxy: • x-litellm-tags → metadata + Note over Proxy: • Guardrails (opt-in) + Note over Proxy: • Multipart form reconstruction end Proxy->>Provider: POST https://us-central1-aiplatform.googleapis.com/... Note over Proxy,Provider: Headers: Authorization: Bearer ya29.google-oauth... Note over Proxy,Provider: Body: { "contents": [...] } ← UNCHANGED - Provider-->>Proxy: Response - + Provider-->>Proxy: Response (streaming or non-streaming) + rect rgb(240, 240, 240) - Note over Proxy: 3. Logging (async, optional) - Note over Proxy: Parse response → calculate cost → log + Note over Proxy: 4. Response Handling (async) + Note over Proxy: • Collect streaming chunks for logging + Note over Proxy: • Cost injection (if enabled) + Note over Proxy: • Parse response → calculate cost → log end Proxy-->>Client: Response (unchanged) @@ -42,7 +52,17 @@ sequenceDiagram - **URL Construction** - Build correct provider URL (e.g., regional endpoints for Vertex AI, Bedrock) - **Auth Header Replacement** - Swap LiteLLM virtual key for actual provider credentials -- **Logging** (optional) - Parse response to extract usage and calculate cost + +## Extra Operations + +| Operation | Description | +|-----------|-------------| +| `x-pass-*` headers | Strip prefix and forward (e.g., `x-pass-anthropic-beta` → `anthropic-beta`) | +| `x-litellm-tags` header | Extract tags and add to request metadata for logging | +| Streaming chunk collection | Collect chunks async for logging after stream completes | +| Multipart form handling | Reconstruct multipart/form-data requests for file uploads | +| Guardrails (opt-in) | Run content filtering when explicitly configured | +| Cost injection | Inject cost into streaming chunks when `include_cost_in_streaming_usage` enabled | ## What Does NOT Change diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index b5d6f2fb745..4a335b54747 100644 --- a/litellm/proxy/policy_engine/attachment_registry.py +++ b/litellm/proxy/policy_engine/attachment_registry.py @@ -5,14 +5,20 @@ Attachments define WHERE policies apply, separate from the policy definitions. This allows the same policy to be attached to multiple scopes. """ -from typing import Any, Dict, List, Optional +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger from litellm.types.proxy.policy_engine import ( PolicyAttachment, + PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, PolicyMatchContext, ) +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + class AttachmentRegistry: """ @@ -188,6 +194,238 @@ class AttachmentRegistry: ) return removed_count + def remove_attachment_by_id(self, attachment_id: str) -> bool: + """ + Remove an attachment by its ID (for DB-synced attachments). + + Args: + attachment_id: The ID of the attachment to remove + + Returns: + True if removed, False if not found + """ + # Note: In-memory attachments don't have IDs, so this is primarily + # for consistency after DB operations + return False + + # ───────────────────────────────────────────────────────────────────────── + # Database CRUD Methods + # ───────────────────────────────────────────────────────────────────────── + + async def add_attachment_to_db( + self, + attachment_request: PolicyAttachmentCreateRequest, + prisma_client: "PrismaClient", + created_by: Optional[str] = None, + ) -> PolicyAttachmentDBResponse: + """ + Add a policy attachment to the database. + + Args: + attachment_request: The attachment creation request + prisma_client: The Prisma client instance + created_by: User who created the attachment + + Returns: + PolicyAttachmentDBResponse with the created attachment + """ + try: + created_attachment = ( + await prisma_client.db.litellm_policyattachmenttable.create( + data={ + "policy_name": attachment_request.policy_name, + "scope": attachment_request.scope, + "teams": attachment_request.teams or [], + "keys": attachment_request.keys or [], + "models": attachment_request.models or [], + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": created_by, + "updated_by": created_by, + } + ) + ) + + # Also add to in-memory registry + attachment = PolicyAttachment( + policy=attachment_request.policy_name, + scope=attachment_request.scope, + teams=attachment_request.teams, + keys=attachment_request.keys, + models=attachment_request.models, + ) + self.add_attachment(attachment) + + return PolicyAttachmentDBResponse( + attachment_id=created_attachment.attachment_id, + policy_name=created_attachment.policy_name, + scope=created_attachment.scope, + teams=created_attachment.teams or [], + keys=created_attachment.keys or [], + models=created_attachment.models or [], + created_at=created_attachment.created_at, + updated_at=created_attachment.updated_at, + created_by=created_attachment.created_by, + updated_by=created_attachment.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error adding attachment to DB: {e}") + raise Exception(f"Error adding attachment to DB: {str(e)}") + + async def delete_attachment_from_db( + self, + attachment_id: str, + prisma_client: "PrismaClient", + ) -> Dict[str, str]: + """ + Delete a policy attachment from the database. + + Args: + attachment_id: The ID of the attachment to delete + prisma_client: The Prisma client instance + + Returns: + Dict with success message + """ + try: + # Get attachment before deleting + attachment = ( + await prisma_client.db.litellm_policyattachmenttable.find_unique( + where={"attachment_id": attachment_id} + ) + ) + + if attachment is None: + raise Exception(f"Attachment with ID {attachment_id} not found") + + # Delete from DB + await prisma_client.db.litellm_policyattachmenttable.delete( + where={"attachment_id": attachment_id} + ) + + # Note: In-memory attachments don't have IDs, so we need to sync from DB + # to properly update in-memory state + await self.sync_attachments_from_db(prisma_client) + + return {"message": f"Attachment {attachment_id} deleted successfully"} + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting attachment from DB: {e}") + raise Exception(f"Error deleting attachment from DB: {str(e)}") + + async def get_attachment_by_id_from_db( + self, + attachment_id: str, + prisma_client: "PrismaClient", + ) -> Optional[PolicyAttachmentDBResponse]: + """ + Get a policy attachment by ID from the database. + + Args: + attachment_id: The ID of the attachment to retrieve + prisma_client: The Prisma client instance + + Returns: + PolicyAttachmentDBResponse if found, None otherwise + """ + try: + attachment = ( + await prisma_client.db.litellm_policyattachmenttable.find_unique( + where={"attachment_id": attachment_id} + ) + ) + + if attachment is None: + return None + + return PolicyAttachmentDBResponse( + attachment_id=attachment.attachment_id, + policy_name=attachment.policy_name, + scope=attachment.scope, + teams=attachment.teams or [], + keys=attachment.keys or [], + models=attachment.models or [], + created_at=attachment.created_at, + updated_at=attachment.updated_at, + created_by=attachment.created_by, + updated_by=attachment.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error getting attachment from DB: {e}") + raise Exception(f"Error getting attachment from DB: {str(e)}") + + async def get_all_attachments_from_db( + self, + prisma_client: "PrismaClient", + ) -> List[PolicyAttachmentDBResponse]: + """ + Get all policy attachments from the database. + + Args: + prisma_client: The Prisma client instance + + Returns: + List of PolicyAttachmentDBResponse objects + """ + try: + attachments = ( + await prisma_client.db.litellm_policyattachmenttable.find_many( + order={"created_at": "desc"}, + ) + ) + + return [ + PolicyAttachmentDBResponse( + attachment_id=a.attachment_id, + policy_name=a.policy_name, + scope=a.scope, + teams=a.teams or [], + keys=a.keys or [], + models=a.models or [], + created_at=a.created_at, + updated_at=a.updated_at, + created_by=a.created_by, + updated_by=a.updated_by, + ) + for a in attachments + ] + except Exception as e: + verbose_proxy_logger.exception(f"Error getting attachments from DB: {e}") + raise Exception(f"Error getting attachments from DB: {str(e)}") + + async def sync_attachments_from_db( + self, + prisma_client: "PrismaClient", + ) -> None: + """ + Sync policy attachments from the database to in-memory registry. + + Args: + prisma_client: The Prisma client instance + """ + try: + attachments = await self.get_all_attachments_from_db(prisma_client) + + # Clear existing attachments and reload from DB + self._attachments = [] + + for attachment_response in attachments: + attachment = PolicyAttachment( + policy=attachment_response.policy_name, + scope=attachment_response.scope, + teams=attachment_response.teams if attachment_response.teams else None, + keys=attachment_response.keys if attachment_response.keys else None, + models=attachment_response.models if attachment_response.models else None, + ) + self._attachments.append(attachment) + + self._initialized = True + verbose_proxy_logger.info( + f"Synced {len(attachments)} attachments from DB to in-memory registry" + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error syncing attachments from DB: {e}") + raise Exception(f"Error syncing attachments from DB: {str(e)}") + # Global singleton instance _attachment_registry: Optional[AttachmentRegistry] = None diff --git a/litellm/proxy/policy_engine/policy_endpoints.py b/litellm/proxy/policy_engine/policy_endpoints.py new file mode 100644 index 00000000000..615e153862a --- /dev/null +++ b/litellm/proxy/policy_engine/policy_endpoints.py @@ -0,0 +1,578 @@ +""" +CRUD ENDPOINTS FOR POLICIES + +Provides REST API endpoints for managing policies and policy attachments. +""" + +from fastapi import APIRouter, Depends, HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry +from litellm.proxy.policy_engine.policy_registry import get_policy_registry +from litellm.types.proxy.policy_engine import ( + PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, + PolicyAttachmentListResponse, + PolicyCreateRequest, + PolicyDBResponse, + PolicyListDBResponse, + PolicyUpdateRequest, +) + +router = APIRouter() + +# Get singleton instances +POLICY_REGISTRY = get_policy_registry() +ATTACHMENT_REGISTRY = get_attachment_registry() + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy CRUD Endpoints +# ───────────────────────────────────────────────────────────────────────────── + + +@router.get( + "/policies/list", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyListDBResponse, +) +async def list_policies(): + """ + List all policies from the database. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/list" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "policies": [ + { + "policy_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "inherit": null, + "description": "Base guardrails for all requests", + "guardrails_add": ["pii_masking"], + "guardrails_remove": [], + "condition": null, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ], + "total_count": 1 + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + policies = await POLICY_REGISTRY.get_all_policies_from_db(prisma_client) + return PolicyListDBResponse(policies=policies, total_count=len(policies)) + except Exception as e: + verbose_proxy_logger.exception(f"Error listing policies: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/policies", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def create_policy( + request: PolicyCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new policy. + + Example Request: + ```bash + curl -X POST "http://localhost:4000/policies" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "policy_name": "global-baseline", + "description": "Base guardrails for all requests", + "guardrails_add": ["pii_masking", "prompt_injection"], + "guardrails_remove": [] + }' + ``` + + Example Response: + ```json + { + "policy_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "inherit": null, + "description": "Base guardrails for all requests", + "guardrails_add": ["pii_masking", "prompt_injection"], + "guardrails_remove": [], + "condition": null, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + created_by = user_api_key_dict.user_id + result = await POLICY_REGISTRY.add_policy_to_db( + policy_request=request, + prisma_client=prisma_client, + created_by=created_by, + ) + return result + except Exception as e: + verbose_proxy_logger.exception(f"Error creating policy: {e}") + if "unique constraint" in str(e).lower(): + raise HTTPException( + status_code=400, + detail=f"Policy with name '{request.policy_name}' already exists", + ) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/{policy_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def get_policy(policy_id: str): + """ + Get a policy by ID. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + result = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if result is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policy: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.put( + "/policies/{policy_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyDBResponse, +) +async def update_policy( + policy_id: str, + request: PolicyUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update an existing policy. + + Example Request: + ```bash + curl -X PUT "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "description": "Updated description", + "guardrails_add": ["pii_masking", "toxicity_filter"] + }' + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if policy exists + existing = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if existing is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + + updated_by = user_api_key_dict.user_id + result = await POLICY_REGISTRY.update_policy_in_db( + policy_id=policy_id, + policy_request=request, + prisma_client=prisma_client, + updated_by=updated_by, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error updating policy: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/policies/{policy_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_policy(policy_id: str): + """ + Delete a policy. + + Example Request: + ```bash + curl -X DELETE "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "message": "Policy 123e4567-e89b-12d3-a456-426614174000 deleted successfully" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if policy exists + existing = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if existing is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + + result = await POLICY_REGISTRY.delete_policy_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting policy: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/{policy_id}/resolved-guardrails", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def get_resolved_guardrails(policy_id: str): + """ + Get the resolved guardrails for a policy (including inherited guardrails). + + This endpoint resolves the full inheritance chain and returns the final + set of guardrails that would be applied for this policy. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000/resolved-guardrails" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "policy_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "healthcare-compliance", + "resolved_guardrails": ["pii_masking", "prompt_injection", "toxicity_filter"] + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Get the policy + policy = await POLICY_REGISTRY.get_policy_by_id_from_db( + policy_id=policy_id, + prisma_client=prisma_client, + ) + if policy is None: + raise HTTPException( + status_code=404, detail=f"Policy with ID {policy_id} not found" + ) + + # Resolve guardrails + resolved = await POLICY_REGISTRY.resolve_guardrails_from_db( + policy_name=policy.policy_name, + prisma_client=prisma_client, + ) + + return { + "policy_id": policy.policy_id, + "policy_name": policy.policy_name, + "resolved_guardrails": resolved, + } + except HTTPException: + raise + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + verbose_proxy_logger.exception(f"Error resolving guardrails: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy Attachment CRUD Endpoints +# ───────────────────────────────────────────────────────────────────────────── + + +@router.get( + "/policies/attachments/list", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyAttachmentListResponse, +) +async def list_policy_attachments(): + """ + List all policy attachments from the database. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/attachments/list" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "attachments": [ + { + "attachment_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "scope": "*", + "teams": [], + "keys": [], + "models": [], + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ], + "total_count": 1 + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + attachments = await ATTACHMENT_REGISTRY.get_all_attachments_from_db( + prisma_client + ) + return PolicyAttachmentListResponse( + attachments=attachments, total_count=len(attachments) + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error listing policy attachments: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/policies/attachments", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyAttachmentDBResponse, +) +async def create_policy_attachment( + request: PolicyAttachmentCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a new policy attachment. + + Example Request: + ```bash + curl -X POST "http://localhost:4000/policies/attachments" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "policy_name": "global-baseline", + "scope": "*" + }' + ``` + + Example with team-specific attachment: + ```bash + curl -X POST "http://localhost:4000/policies/attachments" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "policy_name": "healthcare-compliance", + "teams": ["healthcare-team", "medical-research"] + }' + ``` + + Example Response: + ```json + { + "attachment_id": "123e4567-e89b-12d3-a456-426614174000", + "policy_name": "global-baseline", + "scope": "*", + "teams": [], + "keys": [], + "models": [], + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:00Z" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Verify the policy exists + policy = await POLICY_REGISTRY.get_all_policies_from_db(prisma_client) + policy_names = [p.policy_name for p in policy] + if request.policy_name not in policy_names: + raise HTTPException( + status_code=404, + detail=f"Policy '{request.policy_name}' not found. Create the policy first.", + ) + + created_by = user_api_key_dict.user_id + result = await ATTACHMENT_REGISTRY.add_attachment_to_db( + attachment_request=request, + prisma_client=prisma_client, + created_by=created_by, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error creating policy attachment: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/policies/attachments/{attachment_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], + response_model=PolicyAttachmentDBResponse, +) +async def get_policy_attachment(attachment_id: str): + """ + Get a policy attachment by ID. + + Example Request: + ```bash + curl -X GET "http://localhost:4000/policies/attachments/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + result = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db( + attachment_id=attachment_id, + prisma_client=prisma_client, + ) + if result is None: + raise HTTPException( + status_code=404, + detail=f"Attachment with ID {attachment_id} not found", + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policy attachment: {e}") + raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/policies/attachments/{attachment_id}", + tags=["Policies"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_policy_attachment(attachment_id: str): + """ + Delete a policy attachment. + + Example Request: + ```bash + curl -X DELETE "http://localhost:4000/policies/attachments/123e4567-e89b-12d3-a456-426614174000" \\ + -H "Authorization: Bearer " + ``` + + Example Response: + ```json + { + "message": "Attachment 123e4567-e89b-12d3-a456-426614174000 deleted successfully" + } + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Database not connected") + + try: + # Check if attachment exists + existing = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db( + attachment_id=attachment_id, + prisma_client=prisma_client, + ) + if existing is None: + raise HTTPException( + status_code=404, + detail=f"Attachment with ID {attachment_id} not found", + ) + + result = await ATTACHMENT_REGISTRY.delete_attachment_from_db( + attachment_id=attachment_id, + prisma_client=prisma_client, + ) + return result + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting policy attachment: {e}") + raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 68485f92489..5fb5084f648 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -7,15 +7,22 @@ Policies define WHAT guardrails to apply. WHERE they apply is defined by policy_attachments (see AttachmentRegistry). """ -from typing import Any, Dict, List, Optional +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_proxy_logger from litellm.types.proxy.policy_engine import ( Policy, PolicyCondition, + PolicyCreateRequest, + PolicyDBResponse, PolicyGuardrails, + PolicyUpdateRequest, ) +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + class PolicyRegistry: """ @@ -178,6 +185,364 @@ class PolicyRegistry: return True return False + # ───────────────────────────────────────────────────────────────────────── + # Database CRUD Methods + # ───────────────────────────────────────────────────────────────────────── + + async def add_policy_to_db( + self, + policy_request: PolicyCreateRequest, + prisma_client: "PrismaClient", + created_by: Optional[str] = None, + ) -> PolicyDBResponse: + """ + Add a policy to the database. + + Args: + policy_request: The policy creation request + prisma_client: The Prisma client instance + created_by: User who created the policy + + Returns: + PolicyDBResponse with the created policy + """ + try: + # Build data dict, only include condition if it's set + data: Dict[str, Any] = { + "policy_name": policy_request.policy_name, + "guardrails_add": policy_request.guardrails_add or [], + "guardrails_remove": policy_request.guardrails_remove or [], + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + } + + # Only add optional fields if they have values + if policy_request.inherit is not None: + data["inherit"] = policy_request.inherit + if policy_request.description is not None: + data["description"] = policy_request.description + if created_by is not None: + data["created_by"] = created_by + data["updated_by"] = created_by + if policy_request.condition is not None: + data["condition"] = policy_request.condition.model_dump() + + created_policy = await prisma_client.db.litellm_policytable.create( + data=data + ) + + # Also add to in-memory registry + policy = self._parse_policy( + policy_request.policy_name, + { + "inherit": policy_request.inherit, + "description": policy_request.description, + "guardrails": { + "add": policy_request.guardrails_add, + "remove": policy_request.guardrails_remove, + }, + "condition": policy_request.condition.model_dump() + if policy_request.condition + else None, + }, + ) + self.add_policy(policy_request.policy_name, policy) + + return PolicyDBResponse( + policy_id=created_policy.policy_id, + policy_name=created_policy.policy_name, + inherit=created_policy.inherit, + description=created_policy.description, + guardrails_add=created_policy.guardrails_add or [], + guardrails_remove=created_policy.guardrails_remove or [], + condition=created_policy.condition, + created_at=created_policy.created_at, + updated_at=created_policy.updated_at, + created_by=created_policy.created_by, + updated_by=created_policy.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error adding policy to DB: {e}") + raise Exception(f"Error adding policy to DB: {str(e)}") + + async def update_policy_in_db( + self, + policy_id: str, + policy_request: PolicyUpdateRequest, + prisma_client: "PrismaClient", + updated_by: Optional[str] = None, + ) -> PolicyDBResponse: + """ + Update a policy in the database. + + Args: + policy_id: The ID of the policy to update + policy_request: The policy update request + prisma_client: The Prisma client instance + updated_by: User who updated the policy + + Returns: + PolicyDBResponse with the updated policy + """ + try: + # Build update data - only include fields that are set + update_data: Dict[str, Any] = { + "updated_at": datetime.now(timezone.utc), + "updated_by": updated_by, + } + + if policy_request.policy_name is not None: + update_data["policy_name"] = policy_request.policy_name + if policy_request.inherit is not None: + update_data["inherit"] = policy_request.inherit + if policy_request.description is not None: + update_data["description"] = policy_request.description + if policy_request.guardrails_add is not None: + update_data["guardrails_add"] = policy_request.guardrails_add + if policy_request.guardrails_remove is not None: + update_data["guardrails_remove"] = policy_request.guardrails_remove + if policy_request.condition is not None: + update_data["condition"] = policy_request.condition.model_dump() + + updated_policy = await prisma_client.db.litellm_policytable.update( + where={"policy_id": policy_id}, + data=update_data, + ) + + # Update in-memory registry + policy = self._parse_policy( + updated_policy.policy_name, + { + "inherit": updated_policy.inherit, + "description": updated_policy.description, + "guardrails": { + "add": updated_policy.guardrails_add, + "remove": updated_policy.guardrails_remove, + }, + "condition": updated_policy.condition, + }, + ) + self.add_policy(updated_policy.policy_name, policy) + + return PolicyDBResponse( + policy_id=updated_policy.policy_id, + policy_name=updated_policy.policy_name, + inherit=updated_policy.inherit, + description=updated_policy.description, + guardrails_add=updated_policy.guardrails_add or [], + guardrails_remove=updated_policy.guardrails_remove or [], + condition=updated_policy.condition, + created_at=updated_policy.created_at, + updated_at=updated_policy.updated_at, + created_by=updated_policy.created_by, + updated_by=updated_policy.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error updating policy in DB: {e}") + raise Exception(f"Error updating policy in DB: {str(e)}") + + async def delete_policy_from_db( + self, + policy_id: str, + prisma_client: "PrismaClient", + ) -> Dict[str, str]: + """ + Delete a policy from the database. + + Args: + policy_id: The ID of the policy to delete + prisma_client: The Prisma client instance + + Returns: + Dict with success message + """ + try: + # Get policy name before deleting + policy = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id} + ) + + if policy is None: + raise Exception(f"Policy with ID {policy_id} not found") + + # Delete from DB + await prisma_client.db.litellm_policytable.delete( + where={"policy_id": policy_id} + ) + + # Remove from in-memory registry + self.remove_policy(policy.policy_name) + + return {"message": f"Policy {policy_id} deleted successfully"} + except Exception as e: + verbose_proxy_logger.exception(f"Error deleting policy from DB: {e}") + raise Exception(f"Error deleting policy from DB: {str(e)}") + + async def get_policy_by_id_from_db( + self, + policy_id: str, + prisma_client: "PrismaClient", + ) -> Optional[PolicyDBResponse]: + """ + Get a policy by ID from the database. + + Args: + policy_id: The ID of the policy to retrieve + prisma_client: The Prisma client instance + + Returns: + PolicyDBResponse if found, None otherwise + """ + try: + policy = await prisma_client.db.litellm_policytable.find_unique( + where={"policy_id": policy_id} + ) + + if policy is None: + return None + + return PolicyDBResponse( + policy_id=policy.policy_id, + policy_name=policy.policy_name, + inherit=policy.inherit, + description=policy.description, + guardrails_add=policy.guardrails_add or [], + guardrails_remove=policy.guardrails_remove or [], + condition=policy.condition, + created_at=policy.created_at, + updated_at=policy.updated_at, + created_by=policy.created_by, + updated_by=policy.updated_by, + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policy from DB: {e}") + raise Exception(f"Error getting policy from DB: {str(e)}") + + async def get_all_policies_from_db( + self, + prisma_client: "PrismaClient", + ) -> List[PolicyDBResponse]: + """ + Get all policies from the database. + + Args: + prisma_client: The Prisma client instance + + Returns: + List of PolicyDBResponse objects + """ + try: + policies = await prisma_client.db.litellm_policytable.find_many( + order={"created_at": "desc"}, + ) + + return [ + PolicyDBResponse( + policy_id=p.policy_id, + policy_name=p.policy_name, + inherit=p.inherit, + description=p.description, + guardrails_add=p.guardrails_add or [], + guardrails_remove=p.guardrails_remove or [], + condition=p.condition, + created_at=p.created_at, + updated_at=p.updated_at, + created_by=p.created_by, + updated_by=p.updated_by, + ) + for p in policies + ] + except Exception as e: + verbose_proxy_logger.exception(f"Error getting policies from DB: {e}") + raise Exception(f"Error getting policies from DB: {str(e)}") + + async def sync_policies_from_db( + self, + prisma_client: "PrismaClient", + ) -> None: + """ + Sync policies from the database to in-memory registry. + + Args: + prisma_client: The Prisma client instance + """ + try: + policies = await self.get_all_policies_from_db(prisma_client) + + for policy_response in policies: + policy = self._parse_policy( + policy_response.policy_name, + { + "inherit": policy_response.inherit, + "description": policy_response.description, + "guardrails": { + "add": policy_response.guardrails_add, + "remove": policy_response.guardrails_remove, + }, + "condition": policy_response.condition, + }, + ) + self.add_policy(policy_response.policy_name, policy) + + verbose_proxy_logger.info( + f"Synced {len(policies)} policies from DB to in-memory registry" + ) + except Exception as e: + verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}") + raise Exception(f"Error syncing policies from DB: {str(e)}") + + async def resolve_guardrails_from_db( + self, + policy_name: str, + prisma_client: "PrismaClient", + ) -> List[str]: + """ + Resolve all guardrails for a policy from the database. + + Uses the existing PolicyResolver to handle inheritance chain resolution. + + Args: + policy_name: Name of the policy to resolve + prisma_client: The Prisma client instance + + Returns: + List of resolved guardrail names + """ + from litellm.proxy.policy_engine.policy_resolver import PolicyResolver + + try: + # Load all policies from DB to ensure we have the full inheritance chain + policies = await self.get_all_policies_from_db(prisma_client) + + # Build a temporary in-memory map for resolution + temp_policies = {} + for policy_response in policies: + policy = self._parse_policy( + policy_response.policy_name, + { + "inherit": policy_response.inherit, + "description": policy_response.description, + "guardrails": { + "add": policy_response.guardrails_add, + "remove": policy_response.guardrails_remove, + }, + "condition": policy_response.condition, + }, + ) + temp_policies[policy_response.policy_name] = policy + + # Use the existing PolicyResolver to resolve guardrails + resolved_policy = PolicyResolver.resolve_policy_guardrails( + policy_name=policy_name, + policies=temp_policies, + context=None, # No context needed for simple resolution + ) + + return sorted(resolved_policy.guardrails) + except Exception as e: + verbose_proxy_logger.exception(f"Error resolving guardrails from DB: {e}") + raise Exception(f"Error resolving guardrails from DB: {str(e)}") + # Global singleton instance _policy_registry: Optional[PolicyRegistry] = None diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 994f6ee862c..64dfc6a5d8f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -381,6 +381,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( router as pass_through_router, ) +from litellm.proxy.policy_engine.policy_endpoints import router as policy_crud_router from litellm.proxy.prompts.prompt_endpoints import router as prompts_router from litellm.proxy.public_endpoints import router as public_endpoints_router from litellm.proxy.rag_endpoints.endpoints import router as rag_router @@ -3801,6 +3802,9 @@ class ProxyConfig: if self._should_load_db_object(object_type="guardrails"): await self._init_guardrails_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="policies"): + await self._init_policies_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="vector_stores"): await self._init_vector_stores_in_db(prisma_client=prisma_client) @@ -4024,6 +4028,36 @@ class ProxyConfig: ) ) + async def _init_policies_in_db(self, prisma_client: PrismaClient): + """ + Initialize policies and policy attachments from database into the in-memory registries. + """ + from litellm.proxy.policy_engine.attachment_registry import ( + get_attachment_registry, + ) + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + + try: + # Get the global singleton instances + policy_registry = get_policy_registry() + attachment_registry = get_attachment_registry() + + # Sync policies from DB to in-memory registry + await policy_registry.sync_policies_from_db(prisma_client=prisma_client) + + # Sync attachments from DB to in-memory registry + await attachment_registry.sync_attachments_from_db(prisma_client=prisma_client) + + verbose_proxy_logger.debug( + "Successfully synced policies and attachments from DB" + ) + except Exception as e: + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_policies_in_db - {}".format( + str(e) + ) + ) + async def _init_vector_stores_in_db(self, prisma_client: PrismaClient): from litellm.vector_stores.vector_store_registry import VectorStoreRegistry @@ -10702,6 +10736,7 @@ app.include_router(caching_router) app.include_router(analytics_router) app.include_router(guardrails_router) app.include_router(policy_router) +app.include_router(policy_crud_router) app.include_router(search_tool_management_router) app.include_router(prompts_router) app.include_router(callback_management_endpoints_router) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 22888f6d3af..d7aa6e9f0d0 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -124,8 +124,9 @@ model LiteLLM_TeamTable { updated_at DateTime @default(now()) @updatedAt @map("updated_at") model_spend Json @default("{}") model_max_budget Json @default("{}") - router_settings Json? @default("{}") + router_settings Json? @default("{}") team_member_permissions String[] @default([]) + policies String[] @default([]) model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) @@ -156,6 +157,7 @@ model LiteLLM_DeletedTeamTable { model_max_budget Json @default("{}") router_settings Json? @default("{}") team_member_permissions String[] @default([]) + policies String[] @default([]) model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases // Original timestamps from team creation/updates @@ -197,6 +199,7 @@ model LiteLLM_UserTable { budget_duration String? budget_reset_at DateTime? allowed_cache_controls String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") created_at DateTime? @default(now()) @map("created_at") @@ -283,6 +286,7 @@ model LiteLLM_VerificationToken { budget_reset_at DateTime? allowed_cache_controls String[] @default([]) allowed_routes String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") budget_id String? @@ -327,6 +331,7 @@ model LiteLLM_DeletedVerificationToken { budget_reset_at DateTime? allowed_cache_controls String[] @default([]) allowed_routes String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") router_settings Json? @default("{}") @@ -864,19 +869,31 @@ model LiteLLM_SkillsTable { updated_by String? } -// Claude Code Marketplace - stores plugins for Claude Code integration -model LiteLLM_ClaudeCodePluginTable { - id String @id @default(uuid()) - name String @unique // Plugin name (kebab-case) - version String? // Semantic version - description String? // Plugin description - manifest_json String // Full plugin.json as JSON string - files_json String // All files as JSON: {"path": "content"} - enabled Boolean @default(true) - created_at DateTime @default(now()) - updated_at DateTime @default(now()) @updatedAt - created_by String? - - @@index([name]) - @@map("litellm_claudecodeplugin") +// Policy table for storing guardrail policies +model LiteLLM_PolicyTable { + policy_id String @id @default(uuid()) + policy_name String @unique + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? +} + +// Policy attachment table for defining where policies apply +model LiteLLM_PolicyAttachmentTable { + attachment_id String @id @default(uuid()) + policy_name String // Name of the policy to attach + scope String? // Use '*' for global scope + teams String[] @default([]) // Team aliases or patterns + keys String[] @default([]) // Key aliases or patterns + models String[] @default([]) // Model names or patterns + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? } diff --git a/litellm/secret_managers/hashicorp_secret_manager.py b/litellm/secret_managers/hashicorp_secret_manager.py index dac0397dd95..c59f2ef638a 100644 --- a/litellm/secret_managers/hashicorp_secret_manager.py +++ b/litellm/secret_managers/hashicorp_secret_manager.py @@ -563,7 +563,7 @@ class HashicorpSecretManager(BaseSecretManager): return create_response - except httpx.TimeoutException as e: + except httpx.TimeoutException: verbose_logger.exception("Timeout error occurred during secret rotation") return {"status": "error", "message": "Timeout error occurred"} except Exception as e: diff --git a/litellm/types/proxy/policy_engine/__init__.py b/litellm/types/proxy/policy_engine/__init__.py index 50ed4581013..bc54c3eb36b 100644 --- a/litellm/types/proxy/policy_engine/__init__.py +++ b/litellm/types/proxy/policy_engine/__init__.py @@ -19,13 +19,21 @@ from litellm.types.proxy.policy_engine.policy_types import ( PolicyScope, ) from litellm.types.proxy.policy_engine.resolver_types import ( + PolicyAttachmentCreateRequest, + PolicyAttachmentDBResponse, + PolicyAttachmentListResponse, + PolicyConditionRequest, + PolicyCreateRequest, + PolicyDBResponse, PolicyGuardrailsResponse, PolicyInfoResponse, + PolicyListDBResponse, PolicyListResponse, PolicyMatchContext, PolicyScopeResponse, PolicySummaryItem, PolicyTestResponse, + PolicyUpdateRequest, ResolvedPolicy, ) from litellm.types.proxy.policy_engine.validation_types import ( @@ -58,4 +66,13 @@ __all__ = [ "PolicyScopeResponse", "PolicySummaryItem", "PolicyTestResponse", + # CRUD Request/Response types + "PolicyConditionRequest", + "PolicyCreateRequest", + "PolicyUpdateRequest", + "PolicyDBResponse", + "PolicyListDBResponse", + "PolicyAttachmentCreateRequest", + "PolicyAttachmentDBResponse", + "PolicyAttachmentListResponse", ] diff --git a/litellm/types/proxy/policy_engine/resolver_types.py b/litellm/types/proxy/policy_engine/resolver_types.py index 81ae248d436..9488b8b0841 100644 --- a/litellm/types/proxy/policy_engine/resolver_types.py +++ b/litellm/types/proxy/policy_engine/resolver_types.py @@ -5,7 +5,8 @@ These types are used for matching requests to policies and resolving the final guardrails list. """ -from typing import Dict, List, Optional +from datetime import datetime +from typing import Any, Dict, List, Optional from pydantic import BaseModel, ConfigDict, Field @@ -108,3 +109,168 @@ class PolicyTestResponse(BaseModel): matching_policies: List[str] resolved_guardrails: List[str] message: Optional[str] = None + + +# ───────────────────────────────────────────────────────────────────────────── +# CRUD Request/Response Types for Policy Endpoints +# ───────────────────────────────────────────────────────────────────────────── + + +class PolicyConditionRequest(BaseModel): + """Condition for when a policy applies.""" + + model: Optional[str] = Field( + default=None, + description="Model name pattern (exact match or regex) for when policy applies.", + ) + + +class PolicyCreateRequest(BaseModel): + """Request body for creating a new policy.""" + + policy_name: str = Field(description="Unique name for the policy.") + inherit: Optional[str] = Field( + default=None, + description="Name of parent policy to inherit from.", + ) + description: Optional[str] = Field( + default=None, + description="Human-readable description of the policy.", + ) + guardrails_add: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to add.", + ) + guardrails_remove: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to remove (from inherited).", + ) + condition: Optional[PolicyConditionRequest] = Field( + default=None, + description="Condition for when this policy applies.", + ) + + +class PolicyUpdateRequest(BaseModel): + """Request body for updating a policy.""" + + policy_name: Optional[str] = Field( + default=None, + description="New name for the policy.", + ) + inherit: Optional[str] = Field( + default=None, + description="Name of parent policy to inherit from.", + ) + description: Optional[str] = Field( + default=None, + description="Human-readable description of the policy.", + ) + guardrails_add: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to add.", + ) + guardrails_remove: Optional[List[str]] = Field( + default=None, + description="List of guardrail names to remove (from inherited).", + ) + condition: Optional[PolicyConditionRequest] = Field( + default=None, + description="Condition for when this policy applies.", + ) + + +class PolicyDBResponse(BaseModel): + """Response for a policy from the database.""" + + policy_id: str = Field(description="Unique ID of the policy.") + policy_name: str = Field(description="Name of the policy.") + inherit: Optional[str] = Field(default=None, description="Parent policy name.") + description: Optional[str] = Field(default=None, description="Policy description.") + guardrails_add: List[str] = Field( + default_factory=list, description="Guardrails to add." + ) + guardrails_remove: List[str] = Field( + default_factory=list, description="Guardrails to remove." + ) + condition: Optional[Dict[str, Any]] = Field( + default=None, description="Policy condition." + ) + created_at: Optional[datetime] = Field( + default=None, description="When the policy was created." + ) + updated_at: Optional[datetime] = Field( + default=None, description="When the policy was last updated." + ) + created_by: Optional[str] = Field(default=None, description="Who created the policy.") + updated_by: Optional[str] = Field( + default=None, description="Who last updated the policy." + ) + + +class PolicyListDBResponse(BaseModel): + """Response for listing policies from the database.""" + + policies: List[PolicyDBResponse] = Field( + default_factory=list, description="List of policies." + ) + total_count: int = Field(default=0, description="Total number of policies.") + + +# ───────────────────────────────────────────────────────────────────────────── +# Policy Attachment CRUD Types +# ───────────────────────────────────────────────────────────────────────────── + + +class PolicyAttachmentCreateRequest(BaseModel): + """Request body for creating a policy attachment.""" + + policy_name: str = Field(description="Name of the policy to attach.") + scope: Optional[str] = Field( + default=None, + description="Use '*' for global scope (applies to all requests).", + ) + teams: Optional[List[str]] = Field( + default=None, + description="Team aliases or patterns this attachment applies to.", + ) + keys: Optional[List[str]] = Field( + default=None, + description="Key aliases or patterns this attachment applies to.", + ) + models: Optional[List[str]] = Field( + default=None, + description="Model names or patterns this attachment applies to.", + ) + + +class PolicyAttachmentDBResponse(BaseModel): + """Response for a policy attachment from the database.""" + + attachment_id: str = Field(description="Unique ID of the attachment.") + policy_name: str = Field(description="Name of the attached policy.") + scope: Optional[str] = Field(default=None, description="Scope of the attachment.") + teams: List[str] = Field(default_factory=list, description="Team patterns.") + keys: List[str] = Field(default_factory=list, description="Key patterns.") + models: List[str] = Field(default_factory=list, description="Model patterns.") + created_at: Optional[datetime] = Field( + default=None, description="When the attachment was created." + ) + updated_at: Optional[datetime] = Field( + default=None, description="When the attachment was last updated." + ) + created_by: Optional[str] = Field( + default=None, description="Who created the attachment." + ) + updated_by: Optional[str] = Field( + default=None, description="Who last updated the attachment." + ) + + +class PolicyAttachmentListResponse(BaseModel): + """Response for listing policy attachments.""" + + attachments: List[PolicyAttachmentDBResponse] = Field( + default_factory=list, description="List of policy attachments." + ) + total_count: int = Field(default=0, description="Total number of attachments.") diff --git a/poetry.lock b/poetry.lock index c5f5a87894e..e5c304ea776 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand. [[package]] name = "a2a-sdk" @@ -275,7 +275,7 @@ files = [ {file = "annotated_doc-0.0.4-py3-none-any.whl", hash = "sha256:571ac1dc6991c450b25a9c2d84a3705e2ae7a53467b5d111c24fa8baabbed320"}, {file = "annotated_doc-0.0.4.tar.gz", hash = "sha256:fbcda96e87e9c92ad167c2e53839e57503ecfda18804ea28102353485033faa4"}, ] -markers = {main = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\""} +markers = {main = "(extra == \"mlflow\" or extra == \"proxy\") and python_version >= \"3.10\" or extra == \"proxy\""} [[package]] name = "annotated-types" @@ -343,7 +343,7 @@ zookeeper = ["kazoo"] name = "async-timeout" version = "5.0.1" description = "Timeout context manager for asyncio programs" -optional = true +optional = false python-versions = ">=3.8" groups = ["main"] markers = "python_full_version < \"3.11.3\" and (extra == \"extra-proxy\" or extra == \"proxy\" or python_version < \"3.11\")" @@ -557,48 +557,48 @@ files = [ [[package]] name = "boto3" -version = "1.36.0" +version = "1.40.76" description = "The AWS SDK for Python" optional = true -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main"] markers = "extra == \"proxy\"" files = [ - {file = "boto3-1.36.0-py3-none-any.whl", hash = "sha256:d0ca7a58ce25701a52232cc8df9d87854824f1f2964b929305722ebc7959d5a9"}, - {file = "boto3-1.36.0.tar.gz", hash = "sha256:159898f51c2997a12541c0e02d6e5a8fe2993ddb307b9478fd9a339f98b57e00"}, + {file = "boto3-1.40.76-py3-none-any.whl", hash = "sha256:8df6df755727be40ad9e309cfda07f9a12c147e17b639430c55d4e4feee8a167"}, + {file = "boto3-1.40.76.tar.gz", hash = "sha256:16f4cf97f8dd8e0aae015f4dc66219bd7716a91a40d1e2daa0dafa241a4761c5"}, ] [package.dependencies] -botocore = ">=1.36.0,<1.37.0" +botocore = ">=1.40.76,<1.41.0" jmespath = ">=0.7.1,<2.0.0" -s3transfer = ">=0.11.0,<0.12.0" +s3transfer = ">=0.14.0,<0.15.0" [package.extras] crt = ["botocore[crt] (>=1.21.0,<2.0a0)"] [[package]] name = "botocore" -version = "1.36.26" +version = "1.40.76" description = "Low-level, data-driven core of boto 3." optional = true -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main"] markers = "extra == \"proxy\"" files = [ - {file = "botocore-1.36.26-py3-none-any.whl", hash = "sha256:4e3f19913887a58502e71ef8d696fe7eaa54de7813ff73390cd5883f837dfa6e"}, - {file = "botocore-1.36.26.tar.gz", hash = "sha256:4a63bcef7ecf6146fd3a61dc4f9b33b7473b49bdaf1770e9aaca6eee0c9eab62"}, + {file = "botocore-1.40.76-py3-none-any.whl", hash = "sha256:fe425d386e48ac64c81cbb4a7181688d813df2e2b4c78b95ebe833c9e868c6f4"}, + {file = "botocore-1.40.76.tar.gz", hash = "sha256:2b16024d68b29b973005adfb5039adfe9099ebe772d40a90ca89f2e165c495dc"}, ] [package.dependencies] jmespath = ">=0.7.1,<2.0.0" python-dateutil = ">=2.1,<3.0.0" urllib3 = [ - {version = ">=1.25.4,<2.2.0 || >2.2.0,<3", markers = "python_version >= \"3.10\""}, {version = ">=1.25.4,<1.27", markers = "python_version < \"3.10\""}, + {version = ">=1.25.4,<2.2.0 || >2.2.0,<3", markers = "python_version >= \"3.10\""}, ] [package.extras] -crt = ["awscrt (==0.23.8)"] +crt = ["awscrt (==0.28.4)"] [[package]] name = "cachetools" @@ -607,7 +607,7 @@ description = "Extensible memoizing collections and decorators" optional = true python-versions = ">=3.9" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\"" +markers = "(extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") and python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\"" files = [ {file = "cachetools-6.2.2-py3-none-any.whl", hash = "sha256:6c09c98183bf58560c97b2abfcedcbaf6a896a490f534b031b661d3723b45ace"}, {file = "cachetools-6.2.2.tar.gz", hash = "sha256:8e6d266b25e539df852251cfd6f990b4bc3a141db73b939058d809ebd2590fc6"}, @@ -1393,6 +1393,24 @@ docs = ["myst-parser (==0.18.0)", "sphinx (==5.1.1)"] ssh = ["paramiko (>=2.4.3)"] websockets = ["websocket-client (>=1.3.0)"] +[[package]] +name = "docstring-parser" +version = "0.17.0" +description = "Parse Python docstrings in reST, Google and Numpydoc format" +optional = true +python-versions = ">=3.8" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "docstring_parser-0.17.0-py3-none-any.whl", hash = "sha256:cf2569abd23dce8099b300f9b4fa8191e9582dda731fd533daf54c4551658708"}, + {file = "docstring_parser-0.17.0.tar.gz", hash = "sha256:583de4a309722b3315439bb31d64ba3eebada841f2e2cee23b99df001434c912"}, +] + +[package.extras] +dev = ["pre-commit (>=2.16.0) ; python_version >= \"3.9\"", "pydoctor (>=25.4.0)", "pytest"] +docs = ["pydoctor (>=25.4.0)"] +test = ["pytest"] + [[package]] name = "docutils" version = "0.21.2" @@ -1453,7 +1471,7 @@ files = [ {file = "fastapi-0.121.3-py3-none-any.whl", hash = "sha256:0c78fc87587fcd910ca1bbf5bc8ba37b80e119b388a7206b39f0ecc95ebf53e9"}, {file = "fastapi-0.121.3.tar.gz", hash = "sha256:0055bc24fe53e56a40e9e0ad1ae2baa81622c406e548e501e717634e2dfbc40b"}, ] -markers = {main = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\""} +markers = {main = "(extra == \"mlflow\" or extra == \"proxy\") and python_version >= \"3.10\" or extra == \"proxy\""} [package.dependencies] annotated-doc = ">=0.0.2" @@ -1982,7 +2000,7 @@ description = "Google API client core library" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "python_version >= \"3.14\" and extra == \"extra-proxy\"" +markers = "python_version >= \"3.14\" and (extra == \"extra-proxy\" or extra == \"google\")" files = [ {file = "google_api_core-2.25.2-py3-none-any.whl", hash = "sha256:e9a8f62d363dc8424a8497f4c2a47d6bcda6c16514c935629c257ab5d10210e7"}, {file = "google_api_core-2.25.2.tar.gz", hash = "sha256:1c63aa6af0d0d5e37966f157a77f9396d820fba59f9e43e9415bc3dc5baff300"}, @@ -2010,7 +2028,7 @@ description = "Google API client core library" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "extra == \"extra-proxy\" and python_version < \"3.14\"" +markers = "python_version == \"3.9\" and (extra == \"google\" or extra == \"extra-proxy\") or python_version < \"3.14\" and (extra == \"extra-proxy\" or extra == \"google\")" files = [ {file = "google_api_core-2.28.1-py3-none-any.whl", hash = "sha256:4021b0f8ceb77a6fb4de6fde4502cecab45062e66ff4f2895169e0b35bc9466c"}, {file = "google_api_core-2.28.1.tar.gz", hash = "sha256:2b405df02d68e68ce0fbc138559e6036559e685159d148ae5861013dc201baf8"}, @@ -2020,12 +2038,12 @@ files = [ google-auth = ">=2.14.1,<3.0.0" googleapis-common-protos = ">=1.56.2,<2.0.0" grpcio = [ + {version = ">=1.33.2,<2.0.0", optional = true, markers = "extra == \"grpc\""}, {version = ">=1.49.1,<2.0.0", optional = true, markers = "python_version >= \"3.11\" and extra == \"grpc\""}, - {version = ">=1.33.2,<2.0.0", optional = true, markers = "python_version < \"3.11\" and extra == \"grpc\""}, ] grpcio-status = [ - {version = ">=1.49.1,<2.0.0", optional = true, markers = "python_version >= \"3.11\" and extra == \"grpc\""}, {version = ">=1.33.2,<2.0.0", optional = true, markers = "extra == \"grpc\""}, + {version = ">=1.49.1,<2.0.0", optional = true, markers = "python_version >= \"3.11\" and extra == \"grpc\""}, ] proto-plus = [ {version = ">=1.22.3,<2.0.0"}, @@ -2047,7 +2065,7 @@ description = "Google Authentication Library" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\"" +markers = "(extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") and python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\"" files = [ {file = "google_auth-2.43.0-py2.py3-none-any.whl", hash = "sha256:af628ba6fa493f75c7e9dbe9373d148ca9f4399b5ea29976519e0a3848eddd16"}, {file = "google_auth-2.43.0.tar.gz", hash = "sha256:88228eee5fc21b62a1b5fe773ca15e67778cb07dc8363adcb4a8827b52d81483"}, @@ -2056,6 +2074,7 @@ files = [ [package.dependencies] cachetools = ">=2.0.0,<7.0" pyasn1-modules = ">=0.2.1" +requests = {version = ">=2.20.0,<3.0.0", optional = true, markers = "extra == \"requests\""} rsa = ">=3.1.4,<5" [package.extras] @@ -2068,6 +2087,120 @@ requests = ["requests (>=2.20.0,<3.0.0)"] testing = ["aiohttp (<3.10.0)", "aiohttp (>=3.6.2,<4.0.0)", "aioresponses", "cryptography (<39.0.0) ; python_version < \"3.8\"", "cryptography (<39.0.0) ; python_version < \"3.8\"", "cryptography (>=38.0.3)", "cryptography (>=38.0.3)", "flask", "freezegun", "grpcio", "mock", "oauth2client", "packaging", "pyjwt (>=2.0)", "pyopenssl (<24.3.0)", "pyopenssl (>=20.0.0)", "pytest", "pytest-asyncio", "pytest-cov", "pytest-localserver", "pyu2f (>=0.1.5)", "requests (>=2.20.0,<3.0.0)", "responses", "urllib3"] urllib3 = ["packaging", "urllib3"] +[[package]] +name = "google-cloud-aiplatform" +version = "1.130.0" +description = "Vertex AI API client library" +optional = true +python-versions = ">=3.9" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "google_cloud_aiplatform-1.130.0-py2.py3-none-any.whl", hash = "sha256:f578ccee55655dd9e2300cfcafb178e47c3dfdcf746ad465234b875d3e955929"}, + {file = "google_cloud_aiplatform-1.130.0.tar.gz", hash = "sha256:f66aeb23f0a6848fc2d5bbdf1b5777c3cf8e06056f73ef815317abf89d5a0262"}, +] + +[package.dependencies] +docstring_parser = "<1" +google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.8.dev0,<3.0.0", extras = ["grpc"]} +google-auth = ">=2.14.1,<3.0.0" +google-cloud-bigquery = ">=1.15.0,<3.20.0 || >3.20.0,<4.0.0" +google-cloud-resource-manager = ">=1.3.3,<3.0.0" +google-cloud-storage = [ + {version = ">=1.32.0,<4.0.0", markers = "python_version < \"3.13\""}, + {version = ">=2.10.0,<4.0.0", markers = "python_version >= \"3.13\""}, +] +google-genai = ">=1.37.0,<2.0.0" +packaging = ">=14.3" +proto-plus = ">=1.22.3,<2.0.0" +protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<7.0.0" +pydantic = "<3" +shapely = "<3.0.0" +typing_extensions = "*" + +[package.extras] +adk = ["google-adk (>=1.0.0,<2.0.0)", "opentelemetry-instrumentation-google-genai (>=0.3b0,<1.0.0)"] +ag2 = ["ag2[gemini]", "openinference-instrumentation-autogen (>=0.1.6,<0.2)"] +ag2-testing = ["absl-py", "ag2[gemini]", "cloudpickle (>=3.0,<4.0)", "google-cloud-trace (<2)", "openinference-instrumentation-autogen (>=0.1.6,<0.2)", "opentelemetry-exporter-gcp-logging (>=1.11.0a0,<2.0.0)", "opentelemetry-exporter-gcp-trace (<2)", "opentelemetry-exporter-otlp-proto-http (<2)", "opentelemetry-sdk (<2)", "pydantic (>=2.11.1,<3)", "pytest-xdist", "typing_extensions"] +agent-engines = ["cloudpickle (>=3.0,<4.0)", "google-cloud-logging (<4)", "google-cloud-trace (<2)", "opentelemetry-exporter-gcp-logging (>=1.11.0a0,<2.0.0)", "opentelemetry-exporter-gcp-trace (<2)", "opentelemetry-exporter-otlp-proto-http (<2)", "opentelemetry-sdk (<2)", "packaging (>=24.0)", "pydantic (>=2.11.1,<3)", "typing_extensions"] +autologging = ["mlflow (>=1.27.0) ; python_version >= \"3.13\"", "mlflow (>=1.27.0,<=2.16.0) ; python_version < \"3.13\""] +cloud-profiler = ["tensorboard-plugin-profile (>=2.4.0,<2.18.0)", "werkzeug (>=2.0.0,<4.0.0)"] +datasets = ["pyarrow (>=10.0.1) ; python_version == \"3.11\"", "pyarrow (>=14.0.0) ; python_version >= \"3.12\"", "pyarrow (>=3.0.0,<8.0.0) ; python_version < \"3.11\""] +endpoint = ["requests (>=2.28.1)", "requests-toolbelt (<=1.0.0)"] +evaluation = ["jsonschema", "litellm (>=1.72.4,!=1.77.2,!=1.77.3,!=1.77.4)", "pandas (>=1.0.0)", "pyyaml", "ruamel.yaml", "scikit-learn (<1.6.0) ; python_version <= \"3.10\"", "scikit-learn ; python_version > \"3.10\"", "tqdm (>=4.23.0)"] +full = ["docker (>=5.0.3)", "explainable-ai-sdk (>=1.0.0) ; python_version < \"3.13\"", "fastapi (>=0.71.0,<=0.114.0)", "google-cloud-bigquery", "google-cloud-bigquery-storage", "google-vizier (>=0.1.6)", "httpx (>=0.23.0,<=0.28.1)", "immutabledict", "jsonschema", "lit-nlp (==0.4.0) ; python_version < \"3.14\"", "litellm (>=1.72.4,!=1.77.2,!=1.77.3,!=1.77.4)", "mlflow (>=1.27.0) ; python_version >= \"3.13\"", "mlflow (>=1.27.0,<=2.16.0) ; python_version < \"3.13\"", "numpy (>=1.15.0)", "pandas (>=1.0.0)", "pyarrow (>=10.0.1) ; python_version == \"3.11\"", "pyarrow (>=14.0.0) ; python_version >= \"3.12\"", "pyarrow (>=3.0.0,<8.0.0) ; python_version < \"3.11\"", "pyarrow (>=6.0.1)", "pyyaml", "pyyaml (>=5.3.1,<7)", "ray[default] (>=2.4,<2.5.dev0 || >2.9.0,!=2.9.1,!=2.9.2,<2.10.dev0 || ==2.33.* || >=2.42.dev0,<=2.42.0) ; python_version < \"3.11\"", "ray[default] (>=2.5,<=2.47.1) ; python_version == \"3.11\"", "requests (>=2.28.1)", "requests-toolbelt (<=1.0.0)", "ruamel.yaml", "scikit-learn (<1.6.0) ; python_version <= \"3.10\"", "scikit-learn ; python_version > \"3.10\"", "starlette (>=0.17.1)", "tensorboard-plugin-profile (>=2.4.0,<2.18.0)", "tensorflow (>=2.3.0,<3.0.0) ; python_version < \"3.13\"", "tensorflow (>=2.3.0,<3.0.0) ; python_version < \"3.13\"", "tqdm (>=4.23.0)", "urllib3 (>=1.21.1,<1.27)", "uvicorn[standard] (>=0.16.0)", "werkzeug (>=2.0.0,<4.0.0)"] +langchain = ["langchain (>=0.3,<0.4)", "langchain-core (>=0.3,<0.4)", "langchain-google-vertexai (>=2.0.22,<3)", "langgraph (>=0.2.45,<0.4)", "openinference-instrumentation-langchain (>=0.1.19,<0.2)"] +langchain-testing = ["absl-py", "cloudpickle (>=3.0,<4.0)", "google-cloud-trace (<2)", "langchain (>=0.3,<0.4)", "langchain-core (>=0.3,<0.4)", "langchain-google-vertexai (>=2.0.22,<3)", "langgraph (>=0.2.45,<0.4)", "openinference-instrumentation-langchain (>=0.1.19,<0.2)", "opentelemetry-exporter-gcp-logging (>=1.11.0a0,<2.0.0)", "opentelemetry-exporter-gcp-trace (<2)", "opentelemetry-exporter-otlp-proto-http (<2)", "opentelemetry-sdk (<2)", "pydantic (>=2.11.1,<3)", "pytest-xdist", "typing_extensions"] +lit = ["explainable-ai-sdk (>=1.0.0) ; python_version < \"3.13\"", "lit-nlp (==0.4.0) ; python_version < \"3.14\"", "pandas (>=1.0.0)", "tensorflow (>=2.3.0,<3.0.0) ; python_version < \"3.13\""] +llama-index = ["llama-index", "llama-index-llms-google-genai", "openinference-instrumentation-llama-index (>=3.0,<4.0)"] +llama-index-testing = ["absl-py", "cloudpickle (>=3.0,<4.0)", "google-cloud-trace (<2)", "llama-index", "llama-index-llms-google-genai", "openinference-instrumentation-llama-index (>=3.0,<4.0)", "opentelemetry-exporter-gcp-logging (>=1.11.0a0,<2.0.0)", "opentelemetry-exporter-gcp-trace (<2)", "opentelemetry-exporter-otlp-proto-http (<2)", "opentelemetry-sdk (<2)", "pydantic (>=2.11.1,<3)", "pytest-xdist", "typing_extensions"] +metadata = ["numpy (>=1.15.0)", "pandas (>=1.0.0)"] +pipelines = ["pyyaml (>=5.3.1,<7)"] +prediction = ["docker (>=5.0.3)", "fastapi (>=0.71.0,<=0.114.0)", "httpx (>=0.23.0,<=0.28.1)", "starlette (>=0.17.1)", "uvicorn[standard] (>=0.16.0)"] +private-endpoints = ["requests (>=2.28.1)", "urllib3 (>=1.21.1,<1.27)"] +ray = ["google-cloud-bigquery", "google-cloud-bigquery-storage", "immutabledict", "pandas (>=1.0.0)", "pyarrow (>=6.0.1)", "ray[default] (>=2.4,<2.5.dev0 || >2.9.0,!=2.9.1,!=2.9.2,<2.10.dev0 || ==2.33.* || >=2.42.dev0,<=2.42.0) ; python_version < \"3.11\"", "ray[default] (>=2.5,<=2.47.1) ; python_version == \"3.11\""] +ray-testing = ["google-cloud-bigquery", "google-cloud-bigquery-storage", "immutabledict", "pandas (>=1.0.0)", "pyarrow (>=6.0.1)", "pytest-xdist", "ray[default] (>=2.4,<2.5.dev0 || >2.9.0,!=2.9.1,!=2.9.2,<2.10.dev0 || ==2.33.* || >=2.42.dev0,<=2.42.0) ; python_version < \"3.11\"", "ray[default] (>=2.5,<=2.47.1) ; python_version == \"3.11\"", "ray[train]", "scikit-learn (<1.6.0)", "tensorflow ; python_version < \"3.13\"", "torch (>=2.0.0,<2.1.0)", "xgboost", "xgboost_ray"] +reasoningengine = ["cloudpickle (>=3.0,<4.0)", "google-cloud-trace (<2)", "opentelemetry-exporter-gcp-logging (>=1.11.0a0,<2.0.0)", "opentelemetry-exporter-gcp-trace (<2)", "opentelemetry-exporter-otlp-proto-http (<2)", "opentelemetry-sdk (<2)", "pydantic (>=2.11.1,<3)", "typing_extensions"] +tensorboard = ["tensorboard-plugin-profile (>=2.4.0,<2.18.0)", "werkzeug (>=2.0.0,<4.0.0)"] +testing = ["Pillow", "aiohttp", "bigframes ; python_version >= \"3.10\" and python_version < \"3.14\"", "docker (>=5.0.3)", "explainable-ai-sdk (>=1.0.0) ; python_version < \"3.13\"", "fastapi (>=0.71.0,<=0.114.0)", "google-api-core (>=2.11,<3.0.0)", "google-cloud-bigquery", "google-cloud-bigquery-storage", "google-vizier (>=0.1.6)", "google-vizier (>=0.1.6)", "grpcio-testing", "grpcio-tools (>=1.63.0) ; python_version >= \"3.13\"", "httpx (>=0.23.0,<=0.28.1)", "immutabledict", "immutabledict", "ipython", "jsonschema", "kfp (>=2.6.0,<3.0.0) ; python_version < \"3.13\"", "lit-nlp (==0.4.0) ; python_version < \"3.14\"", "litellm (>=1.72.4,!=1.77.2,!=1.77.3,!=1.77.4)", "mlflow (>=1.27.0) ; python_version >= \"3.13\"", "mlflow (>=1.27.0,<=2.16.0) ; python_version < \"3.13\"", "mock", "nltk", "numpy (>=1.15.0)", "pandas (>=1.0.0)", "protobuf (<=5.29.4)", "pyarrow (>=10.0.1) ; python_version == \"3.11\"", "pyarrow (>=14.0.0) ; python_version >= \"3.12\"", "pyarrow (>=3.0.0,<8.0.0) ; python_version < \"3.11\"", "pyarrow (>=6.0.1)", "pytest-asyncio", "pytest-cov", "pytest-xdist", "pyyaml", "pyyaml (>=5.3.1,<7)", "ray[default] (>=2.4,<2.5.dev0 || >2.9.0,!=2.9.1,!=2.9.2,<2.10.dev0 || ==2.33.* || >=2.42.dev0,<=2.42.0) ; python_version < \"3.11\"", "ray[default] (>=2.5,<=2.47.1) ; python_version == \"3.11\"", "requests (>=2.28.1)", "requests-toolbelt (<=1.0.0)", "requests-toolbelt (<=1.0.0)", "ruamel.yaml", "scikit-learn (<1.6.0) ; python_version <= \"3.10\"", "scikit-learn (<1.6.0) ; python_version <= \"3.10\"", "scikit-learn ; python_version > \"3.10\"", "scikit-learn ; python_version > \"3.10\"", "sentencepiece (>=0.2.0)", "starlette (>=0.17.1)", "tensorboard-plugin-profile (>=2.4.0,<2.18.0)", "tensorboard-plugin-profile (>=2.4.0,<2.18.0)", "tensorflow (==2.14.1) ; python_version <= \"3.11\"", "tensorflow (==2.19.0) ; python_version > \"3.11\" and python_version < \"3.13\"", "tensorflow (>=2.3.0,<3.0.0) ; python_version < \"3.13\"", "tensorflow (>=2.3.0,<3.0.0) ; python_version < \"3.13\"", "torch (>=2.0.0,<2.1.0) ; python_version <= \"3.11\"", "torch (>=2.2.0) ; python_version > \"3.11\" and python_version < \"3.13\"", "tqdm (>=4.23.0)", "urllib3 (>=1.21.1,<1.27)", "uvicorn[standard] (>=0.16.0)", "werkzeug (>=2.0.0,<4.0.0)", "werkzeug (>=2.0.0,<4.0.0)", "xgboost"] +tokenization = ["sentencepiece (>=0.2.0)"] +vizier = ["google-vizier (>=0.1.6)"] +xai = ["tensorflow (>=2.3.0,<3.0.0) ; python_version < \"3.13\""] + +[[package]] +name = "google-cloud-bigquery" +version = "3.40.0" +description = "Google BigQuery API client library" +optional = true +python-versions = ">=3.9" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "google_cloud_bigquery-3.40.0-py3-none-any.whl", hash = "sha256:0469bcf9e3dad3cab65b67cce98180c8c0aacf3253d47f0f8e976f299b49b5ab"}, + {file = "google_cloud_bigquery-3.40.0.tar.gz", hash = "sha256:b3ccb11caf0029f15b29569518f667553fe08f6f1459b959020c83fbbd8f2e68"}, +] + +[package.dependencies] +google-api-core = {version = ">=2.11.1,<3.0.0", extras = ["grpc"]} +google-auth = ">=2.14.1,<3.0.0" +google-cloud-core = ">=2.4.1,<3.0.0" +google-resumable-media = ">=2.0.0,<3.0.0" +packaging = ">=24.2.0" +python-dateutil = ">=2.8.2,<3.0.0" +requests = ">=2.21.0,<3.0.0" + +[package.extras] +all = ["google-cloud-bigquery[bigquery-v2,bqstorage,geopandas,ipython,ipywidgets,matplotlib,opentelemetry,pandas,tqdm]"] +bigquery-v2 = ["proto-plus (>=1.22.3,<2.0.0)", "protobuf (>=3.20.2,!=4.21.0,!=4.21.1,!=4.21.2,!=4.21.3,!=4.21.4,!=4.21.5,<7.0.0)"] +bqstorage = ["google-cloud-bigquery-storage (>=2.18.0,<3.0.0)", "grpcio (>=1.47.0,<2.0.0)", "grpcio (>=1.49.1,<2.0.0) ; python_version >= \"3.11\"", "grpcio (>=1.75.1,<2.0.0) ; python_version >= \"3.14\"", "pyarrow (>=4.0.0)"] +geopandas = ["Shapely (>=1.8.4,<3.0.0)", "geopandas (>=0.9.0,<2.0.0)"] +ipython = ["bigquery-magics (>=0.6.0)", "ipython (>=7.23.1)"] +ipywidgets = ["ipykernel (>=6.2.0)", "ipywidgets (>=7.7.1)"] +matplotlib = ["matplotlib (>=3.10.3) ; python_version >= \"3.10\"", "matplotlib (>=3.7.1,<=3.9.2) ; python_version == \"3.9\""] +opentelemetry = ["opentelemetry-api (>=1.1.0)", "opentelemetry-instrumentation (>=0.20b0)", "opentelemetry-sdk (>=1.1.0)"] +pandas = ["db-dtypes (>=1.0.4,<2.0.0)", "grpcio (>=1.47.0,<2.0.0)", "grpcio (>=1.49.1,<2.0.0) ; python_version >= \"3.11\"", "grpcio (>=1.75.1,<2.0.0) ; python_version >= \"3.14\"", "pandas (>=1.3.0)", "pandas-gbq (>=0.26.1)", "pyarrow (>=3.0.0)"] +tqdm = ["tqdm (>=4.23.4,<5.0.0)"] + +[[package]] +name = "google-cloud-core" +version = "2.5.0" +description = "Google Cloud API client core library" +optional = true +python-versions = ">=3.7" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "google_cloud_core-2.5.0-py3-none-any.whl", hash = "sha256:67d977b41ae6c7211ee830c7912e41003ea8194bff15ae7d72fd6f51e57acabc"}, + {file = "google_cloud_core-2.5.0.tar.gz", hash = "sha256:7c1b7ef5c92311717bd05301aa1a91ffbc565673d3b0b4163a52d8413a186963"}, +] + +[package.dependencies] +google-api-core = ">=1.31.6,<2.0.dev0 || >2.3.0,<3.0.0" +google-auth = ">=1.25.0,<3.0.0" + +[package.extras] +grpc = ["grpcio (>=1.38.0,<2.0.0) ; python_version < \"3.14\"", "grpcio (>=1.75.1,<2.0.0) ; python_version >= \"3.14\"", "grpcio-status (>=1.38.0,<2.0.0)"] + [[package]] name = "google-cloud-iam" version = "2.20.0" @@ -2115,6 +2248,204 @@ grpc-google-iam-v1 = ">=0.12.4,<1.0.0dev" proto-plus = ">=1.22.3,<2.0.0dev" protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0dev" +[[package]] +name = "google-cloud-resource-manager" +version = "1.16.0" +description = "Google Cloud Resource Manager API client library" +optional = true +python-versions = ">=3.7" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "google_cloud_resource_manager-1.16.0-py3-none-any.whl", hash = "sha256:fb9a2ad2b5053c508e1c407ac31abfd1a22e91c32876c1892830724195819a28"}, + {file = "google_cloud_resource_manager-1.16.0.tar.gz", hash = "sha256:cc938f87cc36c2672f062b1e541650629e0d954c405a4dac35ceedee70c267c3"}, +] + +[package.dependencies] +google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0", extras = ["grpc"]} +google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0" +grpc-google-iam-v1 = ">=0.14.0,<1.0.0" +grpcio = [ + {version = ">=1.33.2,<2.0.0"}, + {version = ">=1.75.1,<2.0.0", markers = "python_version >= \"3.14\""}, +] +proto-plus = [ + {version = ">=1.22.3,<2.0.0"}, + {version = ">=1.25.0,<2.0.0", markers = "python_version >= \"3.13\""}, +] +protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<7.0.0" + +[[package]] +name = "google-cloud-storage" +version = "3.4.1" +description = "Google Cloud Storage API client library" +optional = true +python-versions = ">=3.7" +groups = ["main"] +markers = "extra == \"google\" and python_version >= \"3.14\"" +files = [ + {file = "google_cloud_storage-3.4.1-py3-none-any.whl", hash = "sha256:972764cc0392aa097be8f49a5354e22eb47c3f62370067fb1571ffff4a1c1189"}, + {file = "google_cloud_storage-3.4.1.tar.gz", hash = "sha256:6f041a297e23a4b485fad8c305a7a6e6831855c208bcbe74d00332a909f82268"}, +] + +[package.dependencies] +google-api-core = ">=2.15.0,<3.0.0" +google-auth = ">=2.26.1,<3.0.0" +google-cloud-core = ">=2.4.2,<3.0.0" +google-crc32c = ">=1.1.3,<2.0.0" +google-resumable-media = ">=2.7.2,<3.0.0" +requests = ">=2.22.0,<3.0.0" + +[package.extras] +protobuf = ["protobuf (>=3.20.2,<7.0.0)"] +tracing = ["opentelemetry-api (>=1.1.0,<2.0.0)"] + +[[package]] +name = "google-cloud-storage" +version = "3.8.0" +description = "Google Cloud Storage API client library" +optional = true +python-versions = ">=3.7" +groups = ["main"] +markers = "extra == \"google\" and python_version < \"3.14\"" +files = [ + {file = "google_cloud_storage-3.8.0-py3-none-any.whl", hash = "sha256:78cfeae7cac2ca9441d0d0271c2eb4ebfa21aa4c6944dd0ccac0389e81d955a7"}, + {file = "google_cloud_storage-3.8.0.tar.gz", hash = "sha256:cc67952dce84ebc9d44970e24647a58260630b7b64d72360cedaf422d6727f28"}, +] + +[package.dependencies] +google-api-core = ">=2.27.0,<3.0.0" +google-auth = ">=2.26.1,<3.0.0" +google-cloud-core = ">=2.4.2,<3.0.0" +google-crc32c = ">=1.1.3,<2.0.0" +google-resumable-media = ">=2.7.2,<3.0.0" +requests = ">=2.22.0,<3.0.0" + +[package.extras] +grpc = ["google-api-core[grpc] (>=2.27.0,<3.0.0)", "grpc-google-iam-v1 (>=0.14.0,<1.0.0)", "grpcio (>=1.33.2,<2.0.0) ; python_version < \"3.14\"", "grpcio (>=1.75.1,<2.0.0) ; python_version >= \"3.14\"", "grpcio-status (>=1.76.0,<2.0.0)", "proto-plus (>=1.22.3,<2.0.0) ; python_version < \"3.13\"", "proto-plus (>=1.25.0,<2.0.0) ; python_version >= \"3.13\"", "protobuf (>=3.20.2,!=4.21.0,!=4.21.1,!=4.21.2,!=4.21.3,!=4.21.4,!=4.21.5,<7.0.0)"] +protobuf = ["protobuf (>=3.20.2,<7.0.0)"] +tracing = ["opentelemetry-api (>=1.1.0,<2.0.0)"] + +[[package]] +name = "google-crc32c" +version = "1.8.0" +description = "A python wrapper of the C library 'Google CRC32C'" +optional = true +python-versions = ">=3.9" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "google_crc32c-1.8.0-cp310-cp310-macosx_12_0_arm64.whl", hash = "sha256:0470b8c3d73b5f4e3300165498e4cf25221c7eb37f1159e221d1825b6df8a7ff"}, + {file = "google_crc32c-1.8.0-cp310-cp310-macosx_12_0_x86_64.whl", hash = "sha256:119fcd90c57c89f30040b47c211acee231b25a45d225e3225294386f5d258288"}, + {file = "google_crc32c-1.8.0-cp310-cp310-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:6f35aaffc8ccd81ba3162443fabb920e65b1f20ab1952a31b13173a67811467d"}, + {file = "google_crc32c-1.8.0-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:864abafe7d6e2c4c66395c1eb0fe12dc891879769b52a3d56499612ca93b6092"}, + {file = "google_crc32c-1.8.0-cp310-cp310-win_amd64.whl", hash = "sha256:db3fe8eaf0612fc8b20fa21a5f25bd785bc3cd5be69f8f3412b0ac2ffd49e733"}, + {file = "google_crc32c-1.8.0-cp311-cp311-macosx_12_0_arm64.whl", hash = "sha256:014a7e68d623e9a4222d663931febc3033c5c7c9730785727de2a81f87d5bab8"}, + {file = "google_crc32c-1.8.0-cp311-cp311-macosx_12_0_x86_64.whl", hash = "sha256:86cfc00fe45a0ac7359e5214a1704e51a99e757d0272554874f419f79838c5f7"}, + {file = "google_crc32c-1.8.0-cp311-cp311-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:19b40d637a54cb71e0829179f6cb41835f0fbd9e8eb60552152a8b52c36cbe15"}, + {file = "google_crc32c-1.8.0-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:17446feb05abddc187e5441a45971b8394ea4c1b6efd88ab0af393fd9e0a156a"}, + {file = "google_crc32c-1.8.0-cp311-cp311-win_amd64.whl", hash = "sha256:71734788a88f551fbd6a97be9668a0020698e07b2bf5b3aa26a36c10cdfb27b2"}, + {file = "google_crc32c-1.8.0-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:4b8286b659c1335172e39563ab0a768b8015e88e08329fa5321f774275fc3113"}, + {file = "google_crc32c-1.8.0-cp312-cp312-macosx_12_0_x86_64.whl", hash = "sha256:2a3dc3318507de089c5384cc74d54318401410f82aa65b2d9cdde9d297aca7cb"}, + {file = "google_crc32c-1.8.0-cp312-cp312-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:14f87e04d613dfa218d6135e81b78272c3b904e2a7053b841481b38a7d901411"}, + {file = "google_crc32c-1.8.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb5c869c2923d56cb0c8e6bcdd73c009c36ae39b652dbe46a05eb4ef0ad01454"}, + {file = "google_crc32c-1.8.0-cp312-cp312-win_amd64.whl", hash = "sha256:3cc0c8912038065eafa603b238abf252e204accab2a704c63b9e14837a854962"}, + {file = "google_crc32c-1.8.0-cp313-cp313-macosx_12_0_arm64.whl", hash = "sha256:3ebb04528e83b2634857f43f9bb8ef5b2bbe7f10f140daeb01b58f972d04736b"}, + {file = "google_crc32c-1.8.0-cp313-cp313-macosx_12_0_x86_64.whl", hash = "sha256:450dc98429d3e33ed2926fc99ee81001928d63460f8538f21a5d6060912a8e27"}, + {file = "google_crc32c-1.8.0-cp313-cp313-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:3b9776774b24ba76831609ffbabce8cdf6fa2bd5e9df37b594221c7e333a81fa"}, + {file = "google_crc32c-1.8.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:89c17d53d75562edfff86679244830599ee0a48efc216200691de8b02ab6b2b8"}, + {file = "google_crc32c-1.8.0-cp313-cp313-win_amd64.whl", hash = "sha256:57a50a9035b75643996fbf224d6661e386c7162d1dfdab9bc4ca790947d1007f"}, + {file = "google_crc32c-1.8.0-cp314-cp314-macosx_12_0_arm64.whl", hash = "sha256:e6584b12cb06796d285d09e33f63309a09368b9d806a551d8036a4207ea43697"}, + {file = "google_crc32c-1.8.0-cp314-cp314-macosx_12_0_x86_64.whl", hash = "sha256:f4b51844ef67d6cf2e9425983274da75f18b1597bb2c998e1c0a0e8d46f8f651"}, + {file = "google_crc32c-1.8.0-cp314-cp314-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:b0d1a7afc6e8e4635564ba8aa5c0548e3173e41b6384d7711a9123165f582de2"}, + {file = "google_crc32c-1.8.0-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:8b3f68782f3cbd1bce027e48768293072813469af6a61a86f6bb4977a4380f21"}, + {file = "google_crc32c-1.8.0-cp314-cp314-win_amd64.whl", hash = "sha256:d511b3153e7011a27ab6ee6bb3a5404a55b994dc1a7322c0b87b29606d9790e2"}, + {file = "google_crc32c-1.8.0-cp39-cp39-macosx_12_0_arm64.whl", hash = "sha256:ba6aba18daf4d36ad4412feede6221414692f44d17e5428bdd81ad3fc1eee5dc"}, + {file = "google_crc32c-1.8.0-cp39-cp39-macosx_12_0_x86_64.whl", hash = "sha256:87b0072c4ecc9505cfa16ee734b00cd7721d20a0f595be4d40d3d21b41f65ae2"}, + {file = "google_crc32c-1.8.0-cp39-cp39-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:3d488e98b18809f5e322978d4506373599c0c13e6c5ad13e53bb44758e18d215"}, + {file = "google_crc32c-1.8.0-cp39-cp39-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:01f126a5cfddc378290de52095e2c7052be2ba7656a9f0caf4bcd1bfb1833f8a"}, + {file = "google_crc32c-1.8.0-cp39-cp39-win_amd64.whl", hash = "sha256:61f58b28e0b21fcb249a8247ad0db2e64114e201e2e9b4200af020f3b6242c9f"}, + {file = "google_crc32c-1.8.0-pp311-pypy311_pp73-manylinux1_x86_64.manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:87fa445064e7db928226b2e6f0d5304ab4cd0339e664a4e9a25029f384d9bb93"}, + {file = "google_crc32c-1.8.0-pp311-pypy311_pp73-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:f639065ea2042d5c034bf258a9f085eaa7af0cd250667c0635a3118e8f92c69c"}, + {file = "google_crc32c-1.8.0.tar.gz", hash = "sha256:a428e25fb7691024de47fecfbff7ff957214da51eddded0da0ae0e0f03a2cf79"}, +] + +[[package]] +name = "google-genai" +version = "1.47.0" +description = "GenAI Python SDK" +optional = true +python-versions = ">=3.9" +groups = ["main"] +markers = "python_version == \"3.9\" and extra == \"google\"" +files = [ + {file = "google_genai-1.47.0-py3-none-any.whl", hash = "sha256:e3851237556cbdec96007d8028b4b1f2425cdc5c099a8dc36b72a57e42821b60"}, + {file = "google_genai-1.47.0.tar.gz", hash = "sha256:ecece00d0a04e6739ea76cc8dad82ec9593d9380aaabef078990e60574e5bf59"}, +] + +[package.dependencies] +anyio = ">=4.8.0,<5.0.0" +google-auth = ">=2.14.1,<3.0.0" +httpx = ">=0.28.1,<1.0.0" +pydantic = ">=2.9.0,<3.0.0" +requests = ">=2.28.1,<3.0.0" +tenacity = ">=8.2.3,<9.2.0" +typing-extensions = ">=4.11.0,<5.0.0" +websockets = ">=13.0.0,<15.1.0" + +[package.extras] +aiohttp = ["aiohttp (<4.0.0)"] +local-tokenizer = ["protobuf", "sentencepiece (>=0.2.0)"] + +[[package]] +name = "google-genai" +version = "1.55.0" +description = "GenAI Python SDK" +optional = true +python-versions = ">=3.10" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"google\"" +files = [ + {file = "google_genai-1.55.0-py3-none-any.whl", hash = "sha256:98c422762b5ff6e16b8d9a1e4938e8e0ad910392a5422e47f5301498d7f373a1"}, + {file = "google_genai-1.55.0.tar.gz", hash = "sha256:ae9f1318fedb05c7c1b671a4148724751201e8908a87568364a309804064d986"}, +] + +[package.dependencies] +anyio = ">=4.8.0,<5.0.0" +distro = ">=1.7.0,<2" +google-auth = {version = ">=2.14.1,<3.0.0", extras = ["requests"]} +httpx = ">=0.28.1,<1.0.0" +pydantic = ">=2.9.0,<3.0.0" +requests = ">=2.28.1,<3.0.0" +sniffio = "*" +tenacity = ">=8.2.3,<9.2.0" +typing-extensions = ">=4.11.0,<5.0.0" +websockets = ">=13.0.0,<15.1.0" + +[package.extras] +aiohttp = ["aiohttp (<3.13.3)"] +local-tokenizer = ["protobuf", "sentencepiece (>=0.2.0)"] + +[[package]] +name = "google-resumable-media" +version = "2.8.0" +description = "Utilities for Google Media Downloads and Resumable Uploads" +optional = true +python-versions = ">=3.7" +groups = ["main"] +markers = "extra == \"google\"" +files = [ + {file = "google_resumable_media-2.8.0-py3-none-any.whl", hash = "sha256:dd14a116af303845a8d932ddae161a26e86cc229645bc98b39f026f9b1717582"}, + {file = "google_resumable_media-2.8.0.tar.gz", hash = "sha256:f1157ed8b46994d60a1bc432544db62352043113684d4e030ee02e77ebe9a1ae"}, +] + +[package.dependencies] +google-crc32c = ">=1.0.0,<2.0.0" + +[package.extras] +aiohttp = ["aiohttp (>=3.6.2,<4.0.0)", "google-auth (>=1.22.0,<2.0.0)"] +requests = ["requests (>=2.18.0,<3.0.0)"] + [[package]] name = "googleapis-common-protos" version = "1.72.0" @@ -2126,7 +2457,7 @@ files = [ {file = "googleapis_common_protos-1.72.0-py3-none-any.whl", hash = "sha256:4299c5a82d5ae1a9702ada957347726b167f9f8d1fc352477702a1e851ff4038"}, {file = "googleapis_common_protos-1.72.0.tar.gz", hash = "sha256:e55a601c1b32b52d7a3e65f43563e2aa61bcd737998ee672ac9b951cd49319f5"}, ] -markers = {main = "extra == \"extra-proxy\""} +markers = {main = "extra == \"extra-proxy\" or extra == \"google\""} [package.dependencies] grpcio = {version = ">=1.44.0,<2.0.0", optional = true, markers = "extra == \"grpc\""} @@ -2194,7 +2525,7 @@ description = "Lightweight in-process concurrent programming" optional = true python-versions = ">=3.9" groups = ["main"] -markers = "python_version >= \"3.10\" and extra == \"mlflow\" and (platform_machine == \"aarch64\" or platform_machine == \"ppc64le\" or platform_machine == \"x86_64\" or platform_machine == \"amd64\" or platform_machine == \"AMD64\" or platform_machine == \"win32\" or platform_machine == \"WIN32\")" +markers = "python_version >= \"3.10\" and (platform_machine == \"aarch64\" or platform_machine == \"ppc64le\" or platform_machine == \"x86_64\" or platform_machine == \"amd64\" or platform_machine == \"AMD64\" or platform_machine == \"win32\" or platform_machine == \"WIN32\") and extra == \"mlflow\"" files = [ {file = "greenlet-3.2.4-cp310-cp310-macosx_11_0_universal2.whl", hash = "sha256:8c68325b0d0acf8d91dde4e6f930967dd52a5302cd4062932a6b2e7c2969f47c"}, {file = "greenlet-3.2.4-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:94385f101946790ae13da500603491f04a76b6e4c059dab271b3ce2e283b2590"}, @@ -2275,7 +2606,7 @@ description = "IAM API client library" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "extra == \"extra-proxy\"" +markers = "extra == \"extra-proxy\" or extra == \"google\"" files = [ {file = "grpc_google_iam_v1-0.14.3-py3-none-any.whl", hash = "sha256:7a7f697e017a067206a3dfef44e4c634a34d3dee135fe7d7a4613fe3e59217e6"}, {file = "grpc_google_iam_v1-0.14.3.tar.gz", hash = "sha256:879ac4ef33136c5491a6300e27575a9ec760f6cdf9a2518798c1b8977a5dc389"}, @@ -2356,7 +2687,7 @@ files = [ {file = "grpcio-1.76.0-cp39-cp39-win_amd64.whl", hash = "sha256:acab0277c40eff7143c2323190ea57b9ee5fd353d8190ee9652369fae735668a"}, {file = "grpcio-1.76.0.tar.gz", hash = "sha256:7be78388d6da1a25c0d5ec506523db58b18be22d9c37d8d3a32c08be4987bd73"}, ] -markers = {main = "extra == \"extra-proxy\" or extra == \"grpc\""} +markers = {main = "extra == \"extra-proxy\" or extra == \"google\" or extra == \"grpc\""} [package.dependencies] typing-extensions = ">=4.12,<5.0" @@ -2371,7 +2702,7 @@ description = "Status proto mapping for gRPC" optional = true python-versions = ">=3.6" groups = ["main"] -markers = "extra == \"extra-proxy\"" +markers = "extra == \"extra-proxy\" or extra == \"google\"" files = [ {file = "grpcio-status-1.62.3.tar.gz", hash = "sha256:289bdd7b2459794a12cf95dc0cb727bd4a1742c37bd823f760236c937e53a485"}, {file = "grpcio_status-1.62.3-py3-none-any.whl", hash = "sha256:f9049b762ba8de6b1086789d8315846e094edac2c50beaf462338b301a8fd4b8"}, @@ -2389,7 +2720,7 @@ description = "WSGI HTTP Server for UNIX" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "extra == \"proxy\" or (extra == \"proxy\" or extra == \"mlflow\") and platform_system != \"Windows\" and python_version >= \"3.10\"" +markers = "(python_version < \"3.14\" or extra == \"mlflow\" or extra == \"proxy\") and (platform_system != \"Windows\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\") and (extra == \"proxy\" or extra == \"mlflow\")" files = [ {file = "gunicorn-23.0.0-py3-none-any.whl", hash = "sha256:ec400d38950de4dfd418cff8328b2c8faed0edb0d517d3394e457c317908ca4d"}, {file = "gunicorn-23.0.0.tar.gz", hash = "sha256:f014447a0101dc57e294f6c18ca6b40227a4c90e9bdb586042628030cba004ec"}, @@ -3095,15 +3426,15 @@ files = [ [[package]] name = "litellm-proxy-extras" -version = "0.4.23" +version = "0.4.26" 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.23-py3-none-any.whl", hash = "sha256:dfda21203dde9fd97cf364396a9b5be0cfdf00fa9846439ee33ce11b7a52f9ce"}, - {file = "litellm_proxy_extras-0.4.23.tar.gz", hash = "sha256:8e3f95576dc2a296e7f73d8c87e73628bd899b4644c45863960fe3c3762d8f64"}, + {file = "litellm_proxy_extras-0.4.26-py3-none-any.whl", hash = "sha256:6b8720c7ae0b51637a3155b6cb02783c316443fe002641061ed9162d2667aee4"}, + {file = "litellm_proxy_extras-0.4.26.tar.gz", hash = "sha256:cb9544703580172fc6436670e1a693a46ccb34a9057548c94f43e97b7b032754"}, ] [[package]] @@ -3446,8 +3777,8 @@ files = [ [package.dependencies] numpy = [ - {version = ">=1.23.3", markers = "python_version >= \"3.11\""}, {version = ">1.20"}, + {version = ">=1.23.3", markers = "python_version >= \"3.11\""}, {version = ">=1.21.2", markers = "python_version >= \"3.10\""}, {version = ">=1.26.0", markers = "python_version >= \"3.12\""}, ] @@ -3860,7 +4191,7 @@ description = "Fundamental package for array computing in Python" optional = true python-versions = ">=3.9" groups = ["main"] -markers = "(python_version >= \"3.10\" or extra == \"extra-proxy\" or extra == \"semantic-router\") and python_version < \"3.12\" and (extra == \"extra-proxy\" or extra == \"semantic-router\" or extra == \"mlflow\")" +markers = "python_version < \"3.12\" and (python_version >= \"3.10\" or extra == \"extra-proxy\" or extra == \"semantic-router\" or extra == \"google\") and (extra == \"extra-proxy\" or extra == \"semantic-router\" or extra == \"google\" or extra == \"mlflow\")" files = [ {file = "numpy-1.26.4-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:9ff0f4f29c51e2803569d7a51c2304de5554655a60c5d776e35b4a41413830d0"}, {file = "numpy-1.26.4-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:2e4ee3380d6de9c9ec04745830fd9e2eccb3e6cf790d39d7b98ffd19b0dd754a"}, @@ -3907,7 +4238,7 @@ description = "Fundamental package for array computing in Python" optional = true python-versions = ">=3.11" groups = ["main"] -markers = "python_version >= \"3.12\" and (extra == \"extra-proxy\" or extra == \"semantic-router\" or extra == \"mlflow\") and (python_version < \"3.14\" or extra == \"mlflow\")" +markers = "python_version >= \"3.12\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"semantic-router\" or extra == \"mlflow\") and (python_version < \"3.14\" or extra == \"mlflow\" or extra == \"google\")" files = [ {file = "numpy-2.3.5-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:de5672f4a7b200c15a4127042170a694d4df43c992948f5e1af57f0174beed10"}, {file = "numpy-2.3.5-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:acfd89508504a19ed06ef963ad544ec6664518c863436306153e13e94605c218"}, @@ -4838,7 +5169,7 @@ description = "Beautiful, Pythonic protocol buffers" optional = true python-versions = ">=3.7" groups = ["main"] -markers = "extra == \"extra-proxy\"" +markers = "extra == \"google\" or extra == \"extra-proxy\"" files = [ {file = "proto_plus-1.26.1-py3-none-any.whl", hash = "sha256:13285478c2dcf2abb829db158e1047e2f1e8d63a077d94263c2b88b043c75a66"}, {file = "proto_plus-1.26.1.tar.gz", hash = "sha256:21a515a4c4c0088a773899e23c7bbade3d18f9c66c73edd4c7ee3816bc96a012"}, @@ -4870,7 +5201,7 @@ files = [ {file = "protobuf-5.29.5-py3-none-any.whl", hash = "sha256:6cf42630262c59b2d8de33954443d94b746c952b01434fc58a417fdbd2e84bd5"}, {file = "protobuf-5.29.5.tar.gz", hash = "sha256:bc1463bafd4b0929216c35f437a8e28731a2b7fe3d98bb77a600efced5a15c84"}, ] -markers = {main = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\""} +markers = {main = "(extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") and python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\""} [[package]] name = "pyarrow" @@ -4940,7 +5271,7 @@ description = "Pure-Python implementation of ASN.1 types and DER/BER/CER codecs optional = true python-versions = ">=3.8" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\"" +markers = "(extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") and python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\"" files = [ {file = "pyasn1-0.6.1-py3-none-any.whl", hash = "sha256:0d632f46f2ba09143da3a8afe9e33fb6f92fa2320ab7e886e2d0f7672af84629"}, {file = "pyasn1-0.6.1.tar.gz", hash = "sha256:6f580d2bdd84365380830acf45550f2511469f673cb4a5ae3857a3170128b034"}, @@ -4953,7 +5284,7 @@ description = "A collection of ASN.1-based protocols modules" optional = true python-versions = ">=3.8" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\"" +markers = "(extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") and python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\"" files = [ {file = "pyasn1_modules-0.4.2-py3-none-any.whl", hash = "sha256:29253a9207ce32b64c3ac6600edc75368f98473906e8fd1043bd6b5b1de2c14a"}, {file = "pyasn1_modules-0.4.2.tar.gz", hash = "sha256:677091de870a80aae844b1ca6134f54652fa2c8c5a52aa396440ac3106e941e6"}, @@ -4985,7 +5316,7 @@ files = [ {file = "pycparser-2.23-py3-none-any.whl", hash = "sha256:e5c6e8d3fbad53479cab09ac03729e0a9faf2bee3db8208a550daf5af81a5934"}, {file = "pycparser-2.23.tar.gz", hash = "sha256:78816d4f24add8f10a06d6f05b4d424ad9e96cfebf68a4ddc99c65c0720d00c2"}, ] -markers = {main = "(platform_python_implementation != \"PyPy\" or extra == \"proxy\") and implementation_name != \"PyPy\"", dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\""} +markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""} [[package]] name = "pydantic" @@ -5362,7 +5693,7 @@ description = "Extensions to the standard Python datetime module" optional = true python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,>=2.7" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\"" +markers = "(extra == \"mlflow\" or extra == \"proxy\" or extra == \"google\") and python_version >= \"3.10\" or extra == \"proxy\" or extra == \"google\"" files = [ {file = "python-dateutil-2.9.0.post0.tar.gz", hash = "sha256:37dd54208da7e1cd875388217d5e00ebd4179249f90fb72437e91a35459a0ad3"}, {file = "python_dateutil-2.9.0.post0-py2.py3-none-any.whl", hash = "sha256:a8b2bc7bffae282281c8140a97d3aa9c14da0b136dfe83f850eea9a5f7470427"}, @@ -5422,7 +5753,7 @@ description = "World timezone definitions, modern and historical" optional = true python-versions = "*" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\"" +markers = "(extra == \"mlflow\" or extra == \"proxy\") and python_version >= \"3.10\" or extra == \"proxy\"" files = [ {file = "pytz-2025.2-py2.py3-none-any.whl", hash = "sha256:5ddf76296dd8c44c26eb8f4b6f35488f3ccbf6fbbd7adee0b7262d43f0ec2f00"}, {file = "pytz-2025.2.tar.gz", hash = "sha256:360b9e3dbb49a209c21ad61809c7fb453643e048b38924c765813546746e81c3"}, @@ -5435,7 +5766,7 @@ description = "Python for Window Extensions" optional = true python-versions = "*" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"mlflow\") and sys_platform == \"win32\"" +markers = "python_version >= \"3.10\" and sys_platform == \"win32\" and (extra == \"proxy\" or extra == \"mlflow\")" files = [ {file = "pywin32-311-cp310-cp310-win32.whl", hash = "sha256:d03ff496d2a0cd4a5893504789d4a15399133fe82517455e78bad62efbb7f0a3"}, {file = "pywin32-311-cp310-cp310-win_amd64.whl", hash = "sha256:797c2772017851984b97180b0bebe4b620bb86328e8a884bb626156295a63b3b"}, @@ -6241,7 +6572,7 @@ description = "Pure-Python RSA implementation" optional = true python-versions = "<4,>=3.6" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"extra-proxy\") or extra == \"extra-proxy\"" +markers = "(extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") and python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\"" files = [ {file = "rsa-4.9.1-py3-none-any.whl", hash = "sha256:68635866661c6836b8d39430f97a996acbd61bfa49406748ea243539fe239762"}, {file = "rsa-4.9.1.tar.gz", hash = "sha256:e7bdbfdb5497da4c07dfd35530e1a902659db6ff241e39d9953cad06ebd0ae75"}, @@ -6279,22 +6610,22 @@ files = [ [[package]] name = "s3transfer" -version = "0.11.3" +version = "0.14.0" description = "An Amazon S3 Transfer Manager" optional = true -python-versions = ">=3.8" +python-versions = ">=3.9" groups = ["main"] markers = "extra == \"proxy\"" files = [ - {file = "s3transfer-0.11.3-py3-none-any.whl", hash = "sha256:ca855bdeb885174b5ffa95b9913622459d4ad8e331fc98eb01e6d5eb6a30655d"}, - {file = "s3transfer-0.11.3.tar.gz", hash = "sha256:edae4977e3a122445660c7c114bba949f9d191bae3b34a096f18a1c8c354527a"}, + {file = "s3transfer-0.14.0-py3-none-any.whl", hash = "sha256:ea3b790c7077558ed1f02a3072fb3cb992bbbd253392f4b6e9e8976941c7d456"}, + {file = "s3transfer-0.14.0.tar.gz", hash = "sha256:eff12264e7c8b4985074ccce27a3b38a485bb7f7422cc8046fee9be4983e4125"}, ] [package.dependencies] -botocore = ">=1.36.0,<2.0a.0" +botocore = ">=1.37.4,<2.0a.0" [package.extras] -crt = ["botocore[crt] (>=1.36.0,<2.0a.0)"] +crt = ["botocore[crt] (>=1.37.4,<2.0a.0)"] [[package]] name = "scikit-learn" @@ -6542,6 +6873,141 @@ postgres = ["psycopg[binary] (>=3.1.0,<4)"] qdrant = ["qdrant-client (>=1.11.1,<2)"] vision = ["pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\""] +[[package]] +name = "shapely" +version = "2.0.7" +description = "Manipulation and analysis of geometric objects" +optional = true +python-versions = ">=3.7" +groups = ["main"] +markers = "python_version == \"3.9\" and extra == \"google\"" +files = [ + {file = "shapely-2.0.7-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:33fb10e50b16113714ae40adccf7670379e9ccf5b7a41d0002046ba2b8f0f691"}, + {file = "shapely-2.0.7-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:f44eda8bd7a4bccb0f281264b34bf3518d8c4c9a8ffe69a1a05dabf6e8461147"}, + {file = "shapely-2.0.7-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:cf6c50cd879831955ac47af9c907ce0310245f9d162e298703f82e1785e38c98"}, + {file = "shapely-2.0.7-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:04a65d882456e13c8b417562c36324c0cd1e5915f3c18ad516bb32ee3f5fc895"}, + {file = "shapely-2.0.7-cp310-cp310-win32.whl", hash = "sha256:7e97104d28e60b69f9b6a957c4d3a2a893b27525bc1fc96b47b3ccef46726bf2"}, + {file = "shapely-2.0.7-cp310-cp310-win_amd64.whl", hash = "sha256:35524cc8d40ee4752520819f9894b9f28ba339a42d4922e92c99b148bed3be39"}, + {file = "shapely-2.0.7-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:5cf23400cb25deccf48c56a7cdda8197ae66c0e9097fcdd122ac2007e320bc34"}, + {file = "shapely-2.0.7-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:d8f1da01c04527f7da59ee3755d8ee112cd8967c15fab9e43bba936b81e2a013"}, + {file = "shapely-2.0.7-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:8f623b64bb219d62014781120f47499a7adc30cf7787e24b659e56651ceebcb0"}, + {file = "shapely-2.0.7-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:e6d95703efaa64aaabf278ced641b888fc23d9c6dd71f8215091afd8a26a66e3"}, + {file = "shapely-2.0.7-cp311-cp311-win32.whl", hash = "sha256:2f6e4759cf680a0f00a54234902415f2fa5fe02f6b05546c662654001f0793a2"}, + {file = "shapely-2.0.7-cp311-cp311-win_amd64.whl", hash = "sha256:b52f3ab845d32dfd20afba86675c91919a622f4627182daec64974db9b0b4608"}, + {file = "shapely-2.0.7-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:4c2b9859424facbafa54f4a19b625a752ff958ab49e01bc695f254f7db1835fa"}, + {file = "shapely-2.0.7-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:5aed1c6764f51011d69a679fdf6b57e691371ae49ebe28c3edb5486537ffbd51"}, + {file = "shapely-2.0.7-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:73c9ae8cf443187d784d57202199bf9fd2d4bb7d5521fe8926ba40db1bc33e8e"}, + {file = "shapely-2.0.7-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:a9469f49ff873ef566864cb3516091881f217b5d231c8164f7883990eec88b73"}, + {file = "shapely-2.0.7-cp312-cp312-win32.whl", hash = "sha256:6bca5095e86be9d4ef3cb52d56bdd66df63ff111d580855cb8546f06c3c907cd"}, + {file = "shapely-2.0.7-cp312-cp312-win_amd64.whl", hash = "sha256:f86e2c0259fe598c4532acfcf638c1f520fa77c1275912bbc958faecbf00b108"}, + {file = "shapely-2.0.7-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:a0c09e3e02f948631c7763b4fd3dd175bc45303a0ae04b000856dedebefe13cb"}, + {file = "shapely-2.0.7-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:06ff6020949b44baa8fc2e5e57e0f3d09486cd5c33b47d669f847c54136e7027"}, + {file = "shapely-2.0.7-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:5d6dbf096f961ca6bec5640e22e65ccdec11e676344e8157fe7d636e7904fd36"}, + {file = "shapely-2.0.7-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:adeddfb1e22c20548e840403e5e0b3d9dc3daf66f05fa59f1fcf5b5f664f0e98"}, + {file = "shapely-2.0.7-cp313-cp313-win32.whl", hash = "sha256:a7f04691ce1c7ed974c2f8b34a1fe4c3c5dfe33128eae886aa32d730f1ec1913"}, + {file = "shapely-2.0.7-cp313-cp313-win_amd64.whl", hash = "sha256:aaaf5f7e6cc234c1793f2a2760da464b604584fb58c6b6d7d94144fd2692d67e"}, + {file = "shapely-2.0.7-cp37-cp37m-macosx_10_9_x86_64.whl", hash = "sha256:19cbc8808efe87a71150e785b71d8a0e614751464e21fb679d97e274eca7bd43"}, + {file = "shapely-2.0.7-cp37-cp37m-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:fc19b78cc966db195024d8011649b4e22812f805dd49264323980715ab80accc"}, + {file = "shapely-2.0.7-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:dd37d65519b3f8ed8976fa4302a2827cbb96e0a461a2e504db583b08a22f0b98"}, + {file = "shapely-2.0.7-cp37-cp37m-win32.whl", hash = "sha256:25085a30a2462cee4e850a6e3fb37431cbbe4ad51cbcc163af0cea1eaa9eb96d"}, + {file = "shapely-2.0.7-cp37-cp37m-win_amd64.whl", hash = "sha256:1a2e03277128e62f9a49a58eb7eb813fa9b343925fca5e7d631d50f4c0e8e0b8"}, + {file = "shapely-2.0.7-cp38-cp38-macosx_10_9_x86_64.whl", hash = "sha256:e1c4f1071fe9c09af077a69b6c75f17feb473caeea0c3579b3e94834efcbdc36"}, + {file = "shapely-2.0.7-cp38-cp38-macosx_11_0_arm64.whl", hash = "sha256:3697bd078b4459f5a1781015854ef5ea5d824dbf95282d0b60bfad6ff83ec8dc"}, + {file = "shapely-2.0.7-cp38-cp38-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:1e9fed9a7d6451979d914cb6ebbb218b4b4e77c0d50da23e23d8327948662611"}, + {file = "shapely-2.0.7-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:2934834c7f417aeb7cba3b0d9b4441a76ebcecf9ea6e80b455c33c7c62d96a24"}, + {file = "shapely-2.0.7-cp38-cp38-win32.whl", hash = "sha256:2e4a1749ad64bc6e7668c8f2f9479029f079991f4ae3cb9e6b25440e35a4b532"}, + {file = "shapely-2.0.7-cp38-cp38-win_amd64.whl", hash = "sha256:8ae5cb6b645ac3fba34ad84b32fbdccb2ab321facb461954925bde807a0d3b74"}, + {file = "shapely-2.0.7-cp39-cp39-macosx_10_9_x86_64.whl", hash = "sha256:4abeb44b3b946236e4e1a1b3d2a0987fb4d8a63bfb3fdefb8a19d142b72001e5"}, + {file = "shapely-2.0.7-cp39-cp39-macosx_11_0_arm64.whl", hash = "sha256:cd0e75d9124b73e06a42bf1615ad3d7d805f66871aa94538c3a9b7871d620013"}, + {file = "shapely-2.0.7-cp39-cp39-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7977d8a39c4cf0e06247cd2dca695ad4e020b81981d4c82152c996346cf1094b"}, + {file = "shapely-2.0.7-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0145387565fcf8f7c028b073c802956431308da933ef41d08b1693de49990d27"}, + {file = "shapely-2.0.7-cp39-cp39-win32.whl", hash = "sha256:98697c842d5c221408ba8aa573d4f49caef4831e9bc6b6e785ce38aca42d1999"}, + {file = "shapely-2.0.7-cp39-cp39-win_amd64.whl", hash = "sha256:a3fb7fbae257e1b042f440289ee7235d03f433ea880e73e687f108d044b24db5"}, + {file = "shapely-2.0.7.tar.gz", hash = "sha256:28fe2997aab9a9dc026dc6a355d04e85841546b2a5d232ed953e3321ab958ee5"}, +] + +[package.dependencies] +numpy = ">=1.14,<3" + +[package.extras] +docs = ["matplotlib", "numpydoc (==1.1.*)", "sphinx", "sphinx-book-theme", "sphinx-remove-toctrees"] +test = ["pytest", "pytest-cov"] + +[[package]] +name = "shapely" +version = "2.1.2" +description = "Manipulation and analysis of geometric objects" +optional = true +python-versions = ">=3.10" +groups = ["main"] +markers = "python_version >= \"3.10\" and extra == \"google\"" +files = [ + {file = "shapely-2.1.2-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:7ae48c236c0324b4e139bea88a306a04ca630f49be66741b340729d380d8f52f"}, + {file = "shapely-2.1.2-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:eba6710407f1daa8e7602c347dfc94adc02205ec27ed956346190d66579eb9ea"}, + {file = "shapely-2.1.2-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:ef4a456cc8b7b3d50ccec29642aa4aeda959e9da2fe9540a92754770d5f0cf1f"}, + {file = "shapely-2.1.2-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:e38a190442aacc67ff9f75ce60aec04893041f16f97d242209106d502486a142"}, + {file = "shapely-2.1.2-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:40d784101f5d06a1fd30b55fc11ea58a61be23f930d934d86f19a180909908a4"}, + {file = "shapely-2.1.2-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:f6f6cd5819c50d9bcf921882784586aab34a4bd53e7553e175dece6db513a6f0"}, + {file = "shapely-2.1.2-cp310-cp310-win32.whl", hash = "sha256:fe9627c39c59e553c90f5bc3128252cb85dc3b3be8189710666d2f8bc3a5503e"}, + {file = "shapely-2.1.2-cp310-cp310-win_amd64.whl", hash = "sha256:1d0bfb4b8f661b3b4ec3565fa36c340bfb1cda82087199711f86a88647d26b2f"}, + {file = "shapely-2.1.2-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:91121757b0a36c9aac3427a651a7e6567110a4a67c97edf04f8d55d4765f6618"}, + {file = "shapely-2.1.2-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:16a9c722ba774cf50b5d4541242b4cce05aafd44a015290c82ba8a16931ff63d"}, + {file = "shapely-2.1.2-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cc4f7397459b12c0b196c9efe1f9d7e92463cbba142632b4cc6d8bbbbd3e2b09"}, + {file = "shapely-2.1.2-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:136ab87b17e733e22f0961504d05e77e7be8c9b5a8184f685b4a91a84efe3c26"}, + {file = "shapely-2.1.2-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:16c5d0fc45d3aa0a69074979f4f1928ca2734fb2e0dde8af9611e134e46774e7"}, + {file = "shapely-2.1.2-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:6ddc759f72b5b2b0f54a7e7cde44acef680a55019eb52ac63a7af2cf17cb9cd2"}, + {file = "shapely-2.1.2-cp311-cp311-win32.whl", hash = "sha256:2fa78b49485391224755a856ed3b3bd91c8455f6121fee0db0e71cefb07d0ef6"}, + {file = "shapely-2.1.2-cp311-cp311-win_amd64.whl", hash = "sha256:c64d5c97b2f47e3cd9b712eaced3b061f2b71234b3fc263e0fcf7d889c6559dc"}, + {file = "shapely-2.1.2-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:fe2533caae6a91a543dec62e8360fe86ffcdc42a7c55f9dfd0128a977a896b94"}, + {file = "shapely-2.1.2-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:ba4d1333cc0bc94381d6d4308d2e4e008e0bd128bdcff5573199742ee3634359"}, + {file = "shapely-2.1.2-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:0bd308103340030feef6c111d3eb98d50dc13feea33affc8a6f9fa549e9458a3"}, + {file = "shapely-2.1.2-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:1e7d4d7ad262a48bb44277ca12c7c78cb1b0f56b32c10734ec9a1d30c0b0c54b"}, + {file = "shapely-2.1.2-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:e9eddfe513096a71896441a7c37db72da0687b34752c4e193577a145c71736fc"}, + {file = "shapely-2.1.2-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:980c777c612514c0cf99bc8a9de6d286f5e186dcaf9091252fcd444e5638193d"}, + {file = "shapely-2.1.2-cp312-cp312-win32.whl", hash = "sha256:9111274b88e4d7b54a95218e243282709b330ef52b7b86bc6aaf4f805306f454"}, + {file = "shapely-2.1.2-cp312-cp312-win_amd64.whl", hash = "sha256:743044b4cfb34f9a67205cee9279feaf60ba7d02e69febc2afc609047cb49179"}, + {file = "shapely-2.1.2-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:b510dda1a3672d6879beb319bc7c5fd302c6c354584690973c838f46ec3e0fa8"}, + {file = "shapely-2.1.2-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:8cff473e81017594d20ec55d86b54bc635544897e13a7cfc12e36909c5309a2a"}, + {file = "shapely-2.1.2-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:fe7b77dc63d707c09726b7908f575fc04ff1d1ad0f3fb92aec212396bc6cfe5e"}, + {file = "shapely-2.1.2-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:7ed1a5bbfb386ee8332713bf7508bc24e32d24b74fc9a7b9f8529a55db9f4ee6"}, + {file = "shapely-2.1.2-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:a84e0582858d841d54355246ddfcbd1fce3179f185da7470f41ce39d001ee1af"}, + {file = "shapely-2.1.2-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:dc3487447a43d42adcdf52d7ac73804f2312cbfa5d433a7d2c506dcab0033dfd"}, + {file = "shapely-2.1.2-cp313-cp313-win32.whl", hash = "sha256:9c3a3c648aedc9f99c09263b39f2d8252f199cb3ac154fadc173283d7d111350"}, + {file = "shapely-2.1.2-cp313-cp313-win_amd64.whl", hash = "sha256:ca2591bff6645c216695bdf1614fca9c82ea1144d4a7591a466fef64f28f0715"}, + {file = "shapely-2.1.2-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:2d93d23bdd2ed9dc157b46bc2f19b7da143ca8714464249bef6771c679d5ff40"}, + {file = "shapely-2.1.2-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:01d0d304b25634d60bd7cf291828119ab55a3bab87dc4af1e44b07fb225f188b"}, + {file = "shapely-2.1.2-cp313-cp313t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:8d8382dd120d64b03698b7298b89611a6ea6f55ada9d39942838b79c9bc89801"}, + {file = "shapely-2.1.2-cp313-cp313t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:19efa3611eef966e776183e338b2d7ea43569ae99ab34f8d17c2c054d3205cc0"}, + {file = "shapely-2.1.2-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:346ec0c1a0fcd32f57f00e4134d1200e14bf3f5ae12af87ba83ca275c502498c"}, + {file = "shapely-2.1.2-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:6305993a35989391bd3476ee538a5c9a845861462327efe00dd11a5c8c709a99"}, + {file = "shapely-2.1.2-cp313-cp313t-win32.whl", hash = "sha256:c8876673449f3401f278c86eb33224c5764582f72b653a415d0e6672fde887bf"}, + {file = "shapely-2.1.2-cp313-cp313t-win_amd64.whl", hash = "sha256:4a44bc62a10d84c11a7a3d7c1c4fe857f7477c3506e24c9062da0db0ae0c449c"}, + {file = "shapely-2.1.2-cp314-cp314-macosx_10_13_x86_64.whl", hash = "sha256:9a522f460d28e2bf4e12396240a5fc1518788b2fcd73535166d748399ef0c223"}, + {file = "shapely-2.1.2-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:1ff629e00818033b8d71139565527ced7d776c269a49bd78c9df84e8f852190c"}, + {file = "shapely-2.1.2-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:f67b34271dedc3c653eba4e3d7111aa421d5be9b4c4c7d38d30907f796cb30df"}, + {file = "shapely-2.1.2-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:21952dc00df38a2c28375659b07a3979d22641aeb104751e769c3ee825aadecf"}, + {file = "shapely-2.1.2-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:1f2f33f486777456586948e333a56ae21f35ae273be99255a191f5c1fa302eb4"}, + {file = "shapely-2.1.2-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:cf831a13e0d5a7eb519e96f58ec26e049b1fad411fc6fc23b162a7ce04d9cffc"}, + {file = "shapely-2.1.2-cp314-cp314-win32.whl", hash = "sha256:61edcd8d0d17dd99075d320a1dd39c0cb9616f7572f10ef91b4b5b00c4aeb566"}, + {file = "shapely-2.1.2-cp314-cp314-win_amd64.whl", hash = "sha256:a444e7afccdb0999e203b976adb37ea633725333e5b119ad40b1ca291ecf311c"}, + {file = "shapely-2.1.2-cp314-cp314t-macosx_10_13_x86_64.whl", hash = "sha256:5ebe3f84c6112ad3d4632b1fd2290665aa75d4cef5f6c5d77c4c95b324527c6a"}, + {file = "shapely-2.1.2-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:5860eb9f00a1d49ebb14e881f5caf6c2cf472c7fd38bd7f253bbd34f934eb076"}, + {file = "shapely-2.1.2-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:b705c99c76695702656327b819c9660768ec33f5ce01fa32b2af62b56ba400a1"}, + {file = "shapely-2.1.2-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:a1fd0ea855b2cf7c9cddaf25543e914dd75af9de08785f20ca3085f2c9ca60b0"}, + {file = "shapely-2.1.2-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:df90e2db118c3671a0754f38e36802db75fe0920d211a27481daf50a711fdf26"}, + {file = "shapely-2.1.2-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:361b6d45030b4ac64ddd0a26046906c8202eb60d0f9f53085f5179f1d23021a0"}, + {file = "shapely-2.1.2-cp314-cp314t-win32.whl", hash = "sha256:b54df60f1fbdecc8ebc2c5b11870461a6417b3d617f555e5033f1505d36e5735"}, + {file = "shapely-2.1.2-cp314-cp314t-win_amd64.whl", hash = "sha256:0036ac886e0923417932c2e6369b6c52e38e0ff5d9120b90eef5cd9a5fc5cae9"}, + {file = "shapely-2.1.2.tar.gz", hash = "sha256:2ed4ecb28320a433db18a5bf029986aa8afcfd740745e78847e330d5d94922a9"}, +] + +[package.dependencies] +numpy = ">=1.21" + +[package.extras] +docs = ["matplotlib", "numpydoc (==1.1.*)", "sphinx", "sphinx-book-theme", "sphinx-remove-toctrees"] +test = ["pytest", "pytest-cov", "scipy-doctest"] + [[package]] name = "shellingham" version = "1.5.4" @@ -6561,7 +7027,7 @@ description = "Python 2 and 3 compatibility utilities" optional = true python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,>=2.7" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\"" +markers = "(extra == \"mlflow\" or extra == \"proxy\" or extra == \"google\") and python_version >= \"3.10\" or extra == \"proxy\" or extra == \"google\"" files = [ {file = "six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274"}, {file = "six-1.17.0.tar.gz", hash = "sha256:ff70335d468e7eb6ec65b95b99d3a2836546063f63acc5171de367e834932a81"}, @@ -6984,6 +7450,26 @@ examples = ["aiosqlite (>=0.21.0)", "fastapi (>=0.115.12)", "sqlalchemy[asyncio] granian = ["granian (>=2.3.1)"] uvicorn = ["uvicorn (>=0.34.0)"] +[[package]] +name = "starlette" +version = "0.49.3" +description = "The little ASGI library that shines." +optional = false +python-versions = ">=3.9" +groups = ["main", "dev"] +files = [ + {file = "starlette-0.49.3-py3-none-any.whl", hash = "sha256:b579b99715fdc2980cf88c8ec96d3bf1ce16f5a8051a7c2b84ef9b1cdecaea2f"}, + {file = "starlette-0.49.3.tar.gz", hash = "sha256:1c14546f299b5901a1ea0e34410575bc33bbd741377a10484a54445588d00284"}, +] +markers = {main = "python_version == \"3.9\" and extra == \"proxy\"", dev = "python_version == \"3.9\""} + +[package.dependencies] +anyio = ">=3.6.2,<5" +typing-extensions = {version = ">=4.10.0", markers = "python_version < \"3.13\""} + +[package.extras] +full = ["httpx (>=0.27.0,<0.29.0)", "itsdangerous", "jinja2", "python-multipart (>=0.0.18)", "pyyaml"] + [[package]] name = "starlette" version = "0.50.0" @@ -7044,7 +7530,7 @@ description = "Retry code until it succeeds" optional = true python-versions = ">=3.9" groups = ["main"] -markers = "extra == \"extra-proxy\" and python_version < \"3.14\"" +markers = "(extra == \"extra-proxy\" or extra == \"google\") and (python_version < \"3.14\" or extra == \"google\")" files = [ {file = "tenacity-9.1.2-py3-none-any.whl", hash = "sha256:f77bf36710d8b73a50b2dd155c97b870017ad21afe6ab300326b0371b3b05138"}, {file = "tenacity-9.1.2.tar.gz", hash = "sha256:1169d376c297e7de388d18b4481760d478b0e99a777cad3a9c86e556f4b697cb"}, @@ -7522,7 +8008,7 @@ description = "The lightning-fast ASGI server." optional = true python-versions = ">=3.8" groups = ["main"] -markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\"" +markers = "(extra == \"mlflow\" or extra == \"proxy\") and python_version >= \"3.10\" or extra == \"proxy\"" files = [ {file = "uvicorn-0.31.1-py3-none-any.whl", hash = "sha256:adc42d9cac80cf3e51af97c1851648066841e7cfb6993a4ca8de29ac1548ed41"}, {file = "uvicorn-0.31.1.tar.gz", hash = "sha256:f5167919867b161b7bcaf32646c6a94cdbd4c3aa2eb5c17d36bb9aa5cfd8c493"}, @@ -7543,7 +8029,7 @@ description = "Fast implementation of asyncio event loop on top of libuv" optional = true python-versions = ">=3.8.0" groups = ["main"] -markers = "sys_platform != \"win32\" and extra == \"proxy\"" +markers = "extra == \"proxy\" and sys_platform != \"win32\"" files = [ {file = "uvloop-0.21.0-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:ec7e6b09a6fdded42403182ab6b832b71f4edaf7f37a9a0e371a01db5f0cb45f"}, {file = "uvloop-0.21.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:196274f2adb9689a289ad7d65700d37df0c0930fd8e4e743fa4834e850d7719d"}, @@ -7596,7 +8082,7 @@ description = "Waitress WSGI server" optional = true python-versions = ">=3.9.0" groups = ["main"] -markers = "python_version >= \"3.10\" and extra == \"mlflow\" and platform_system == \"Windows\"" +markers = "python_version >= \"3.10\" and platform_system == \"Windows\" and extra == \"mlflow\"" files = [ {file = "waitress-3.0.2-py3-none-any.whl", hash = "sha256:c56d67fd6e87c2ee598b76abdd4e96cfad1f24cacdea5078d382b1f9d7b5ed2e"}, {file = "waitress-3.0.2.tar.gz", hash = "sha256:682aaaf2af0c44ada4abfb70ded36393f0e307f4ab9456a215ce0020baefc31f"}, @@ -7613,7 +8099,7 @@ description = "An implementation of the WebSocket Protocol (RFC 6455 & 7692)" optional = true python-versions = ">=3.9" groups = ["main"] -markers = "extra == \"proxy\"" +markers = "extra == \"google\" or extra == \"proxy\"" files = [ {file = "websockets-15.0.1-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:d63efaa0cd96cf0c5fe4d581521d9fa87744540d4bc999ae6e08595a1014b45b"}, {file = "websockets-15.0.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:ac60e3b188ec7574cb761b08d50fcedf9d77f1530352db4eef1707fe9dee7205"}, @@ -7996,6 +8482,7 @@ type = ["pytest-mypy"] [extras] caching = ["diskcache"] extra-proxy = ["a2a-sdk", "azure-identity", "azure-keyvault-secrets", "google-cloud-iam", "google-cloud-kms", "prisma", "redisvl", "resend"] +google = ["google-cloud-aiplatform"] grpc = ["grpcio", "grpcio"] mlflow = ["mlflow"] proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "polars", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "soundfile", "uvicorn", "uvloop", "websockets"] @@ -8005,4 +8492,4 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.9,<4.0" -content-hash = "f6a98e687d478db6e30274a4cf70391960775cbf648da0783558444da3a662ea" +content-hash = "ee8411854731cb6cc038a6135d7589a81789a269e718a226101953ba24fb2033" diff --git a/pyproject.toml b/pyproject.toml index 1e2dfe43d8b..862611879a5 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.25", optional = true} +litellm-proxy-extras = {version = "0.4.26", optional = true} rich = {version = "13.7.1", optional = true} litellm-enterprise = {version = "0.1.27", optional = true} diskcache = {version = "^5.6.1", optional = true} diff --git a/requirements.txt b/requirements.txt index 7f662d9ef8e..9e32cd79507 100644 --- a/requirements.txt +++ b/requirements.txt @@ -50,7 +50,7 @@ sentry_sdk==2.21.0 # for sentry error handling detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests cryptography==44.0.1 tzdata==2025.1 # IANA time zone database -litellm-proxy-extras==0.4.25 # for proxy extras - e.g. prisma migrations +litellm-proxy-extras==0.4.26 # 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 52170f2f3e6..d7aa6e9f0d0 100644 --- a/schema.prisma +++ b/schema.prisma @@ -126,6 +126,7 @@ model LiteLLM_TeamTable { model_max_budget Json @default("{}") router_settings Json? @default("{}") team_member_permissions String[] @default([]) + policies String[] @default([]) model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) @@ -156,6 +157,7 @@ model LiteLLM_DeletedTeamTable { model_max_budget Json @default("{}") router_settings Json? @default("{}") team_member_permissions String[] @default([]) + policies String[] @default([]) model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases // Original timestamps from team creation/updates @@ -197,6 +199,7 @@ model LiteLLM_UserTable { budget_duration String? budget_reset_at DateTime? allowed_cache_controls String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") created_at DateTime? @default(now()) @map("created_at") @@ -283,6 +286,7 @@ model LiteLLM_VerificationToken { budget_reset_at DateTime? allowed_cache_controls String[] @default([]) allowed_routes String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") budget_id String? @@ -327,6 +331,7 @@ model LiteLLM_DeletedVerificationToken { budget_reset_at DateTime? allowed_cache_controls String[] @default([]) allowed_routes String[] @default([]) + policies String[] @default([]) model_spend Json @default("{}") model_max_budget Json @default("{}") router_settings Json? @default("{}") @@ -863,3 +868,32 @@ model LiteLLM_SkillsTable { updated_at DateTime @default(now()) @updatedAt updated_by String? } + +// Policy table for storing guardrail policies +model LiteLLM_PolicyTable { + policy_id String @id @default(uuid()) + policy_name String @unique + inherit String? // Name of parent policy to inherit from + description String? + guardrails_add String[] @default([]) + guardrails_remove String[] @default([]) + condition Json? @default("{}") // Policy conditions (e.g., model matching) + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? +} + +// Policy attachment table for defining where policies apply +model LiteLLM_PolicyAttachmentTable { + attachment_id String @id @default(uuid()) + policy_name String // Name of the policy to attach + scope String? // Use '*' for global scope + teams String[] @default([]) // Team aliases or patterns + keys String[] @default([]) // Key aliases or patterns + models String[] @default([]) // Model names or patterns + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? +} diff --git a/tests/local_testing/test_router_utils.py b/tests/local_testing/test_router_utils.py index 0e3835a7f9d..7ade0777093 100644 --- a/tests/local_testing/test_router_utils.py +++ b/tests/local_testing/test_router_utils.py @@ -265,7 +265,7 @@ async def test_call_router_callbacks_on_success(): ) assert increment["increment_value"] == 1 - +@pytest.mark.serial @pytest.mark.asyncio async def test_call_router_callbacks_on_failure(): router = Router( @@ -288,7 +288,7 @@ async def test_call_router_callbacks_on_failure(): mock_response="litellm.RateLimitError", num_retries=0, ) - await asyncio.sleep(1) + await asyncio.sleep(3) print(mock_callback.call_args_list) assert mock_callback.call_count == 1 diff --git a/tests/logging_callback_tests/test_otel_logging.py b/tests/logging_callback_tests/test_otel_logging.py index b53be44fe0f..a0c78305e60 100644 --- a/tests/logging_callback_tests/test_otel_logging.py +++ b/tests/logging_callback_tests/test_otel_logging.py @@ -10,11 +10,25 @@ sys.path.insert( import pytest import litellm -from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig, Span import asyncio import logging +from opentelemetry import trace from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from litellm._logging import verbose_logger +from litellm.integrations.arize.arize_phoenix import ArizePhoenixLogger +from litellm.integrations._types.open_inference import ( + OpenInferenceSpanKindValues, + SpanAttributes as OISpanAttributes, +) +from litellm.integrations.opentelemetry import ( + LITELLM_PROXY_REQUEST_SPAN_NAME, + LITELLM_TRACER_NAME, + LITELLM_REQUEST_SPAN_NAME, + OpenTelemetry, + OpenTelemetryConfig, + RAW_REQUEST_SPAN_NAME, + Span, +) from litellm.proxy._types import SpanAttributes verbose_logger.setLevel(logging.DEBUG) @@ -242,3 +256,51 @@ def validate_redacted_message_span_attributes(span): ), f"Non-metadata attribute found: {attr}" pass + +@pytest.mark.asyncio +async def test_arize_phoenix_adds_openinference_kind_and_avoids_duplicate_litellm_spans(): + """ + Ensure Arize Phoenix spans include OpenInference span kind and do not create + a duplicate litellm_request span when a proxy parent span is already active. + """ + + exporter.clear() + litellm.logging_callback_manager._reset_all_callbacks() + + otel_logger = ArizePhoenixLogger(config=OpenTelemetryConfig(exporter=exporter)) + litellm.callbacks = [otel_logger] + litellm.success_callback = [] + litellm.failure_callback = [] + + tracer = trace.get_tracer(LITELLM_TRACER_NAME) + parent_span = tracer.start_span(LITELLM_PROXY_REQUEST_SPAN_NAME) + + # Keep parent span active; OpenTelemetry logger will attach attributes and end it. + with trace.use_span(parent_span, end_on_exit=False): + await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "ping"}], + mock_response="pong", + ) + + # Flush span processing + await asyncio.sleep(1) + + if parent_span.is_recording(): + parent_span.end() + + spans = exporter.get_finished_spans() + + span_names = [span.name for span in spans] + assert LITELLM_REQUEST_SPAN_NAME not in span_names + assert span_names.count(LITELLM_PROXY_REQUEST_SPAN_NAME) == 1 + assert span_names.count(RAW_REQUEST_SPAN_NAME) == 1 + + # All spans should belong to the same trace (parent + raw child) + assert len({span.context.trace_id for span in spans}) == 1 + assert len(spans) == 2 + + proxy_span = next(span for span in spans if span.name == LITELLM_PROXY_REQUEST_SPAN_NAME) + assert proxy_span.attributes.get(OISpanAttributes.OPENINFERENCE_SPAN_KIND) == OpenInferenceSpanKindValues.LLM.value + + exporter.clear() diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 5f03ef18171..f52f14b860e 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -215,28 +215,21 @@ class TestProxyInitializationHelpers: assert "pool_timeout=60" in modified_url @patch("uvicorn.run") - @patch("builtins.print") - def test_skip_server_startup(self, mock_print, mock_uvicorn_run): - """Test that the skip_server_startup flag prevents server startup when True""" + @patch("atexit.register") # 🔥 critical + def test_skip_server_startup(self, mock_atexit_register, mock_uvicorn_run): from click.testing import CliRunner - from litellm.proxy.proxy_cli import run_server runner = CliRunner() - mock_app = MagicMock() - mock_proxy_config = MagicMock() - mock_key_mgmt = MagicMock() - mock_save_worker_config = MagicMock() - with patch.dict( "sys.modules", { "proxy_server": MagicMock( - app=mock_app, - ProxyConfig=mock_proxy_config, - KeyManagementSettings=mock_key_mgmt, - save_worker_config=mock_save_worker_config, + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), ) }, ), patch( @@ -248,16 +241,15 @@ class TestProxyInitializationHelpers: "port": 8000, } + # --- skip startup --- result = runner.invoke(run_server, ["--local", "--skip_server_startup"]) assert result.exit_code == 0 + assert "Skipping server startup" in result.output mock_uvicorn_run.assert_not_called() - mock_print.assert_any_call( - "LiteLLM: Setup complete. Skipping server startup as requested." - ) + # --- normal startup --- mock_uvicorn_run.reset_mock() - mock_print.reset_mock() result = runner.invoke(run_server, ["--local"]) diff --git a/tests/test_proxy_server_non_root.py b/tests/test_proxy_server_non_root.py index aedd3f9202b..6a73b509dfd 100644 --- a/tests/test_proxy_server_non_root.py +++ b/tests/test_proxy_server_non_root.py @@ -1,6 +1,6 @@ from unittest.mock import patch - - +import pytest +@pytest.mark.skip(reason="Very Flaky in CI, will debug later") def test_restructure_ui_html_files_skipped_in_non_root(monkeypatch): """ Test that _restructure_ui_html_files is SKIPPED when: @@ -36,7 +36,7 @@ def test_restructure_ui_html_files_skipped_in_non_root(monkeypatch): # Verify it was NOT called mock_restructure.assert_not_called() - +@pytest.mark.skip(reason="Very Flaky in CI, will debug later") def test_restructure_ui_html_files_NOT_skipped_locally(monkeypatch): """ Test that _restructure_ui_html_files is NOT skipped for local development diff --git a/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx b/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx index 260cac16e02..405f8329b67 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx @@ -19,6 +19,7 @@ import { ExperimentOutlined, ToolOutlined, TagsOutlined, + AuditOutlined, } from "@ant-design/icons"; // import { // all_admin_roles, @@ -102,6 +103,8 @@ const routeFor = (slug: string): string => { return "logs"; case "guardrails": return "guardrails"; + case "policies": + return "policies"; // tools case "mcp-servers": @@ -202,6 +205,13 @@ const menuItems: MenuItemCfg[] = [ icon: , roles: all_admin_roles, }, + { + key: "28", + page: "policies", + label: "Policies", + icon: , + roles: all_admin_roles, + }, { key: "26", page: "tools", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/page.tsx new file mode 100644 index 00000000000..86e063c7358 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/page.tsx @@ -0,0 +1,37 @@ +"use client"; + +import PoliciesPanel from "@/components/policies"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { + getPoliciesList, + createPolicyCall, + updatePolicyCall, + deletePolicyCall, + getPolicyInfo, + getPolicyAttachmentsList, + createPolicyAttachmentCall, + deletePolicyAttachmentCall, + getGuardrailsList, +} from "@/components/networking"; + +const PoliciesPage = () => { + const { accessToken, userRole } = useAuthorized(); + + return ( + + ); +}; + +export default PoliciesPage; diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index 8ca25eb9e4c..8ac1f756f96 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -14,6 +14,7 @@ import LoadingScreen from "@/components/common_components/LoadingScreen"; import { CostTrackingSettings } from "@/components/CostTrackingSettings"; import GeneralSettings from "@/components/general_settings"; import GuardrailsPanel from "@/components/guardrails"; +import PoliciesPanel from "@/components/policies"; import { Team } from "@/components/key_team_helpers/key_list"; import { MCPServers } from "@/components/mcp_tools"; import ModelHubTable from "@/components/AIHub/ModelHubTable"; @@ -472,6 +473,8 @@ export default function CreateKeyPage() { ) : page == "guardrails" ? ( + ) : page == "policies" ? ( + ) : page == "agents" ? ( ) : page == "prompts" ? ( diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.test.tsx new file mode 100644 index 00000000000..0e05c141570 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.test.tsx @@ -0,0 +1,309 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import AddFallbacks, { Fallbacks } from "./AddFallbacks"; +import * as fetchModelsModule from "../../../playground/llm_calls/fetch_models"; + +vi.mock("../../../playground/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn(), +})); + +vi.mock("antd", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + message: { + error: vi.fn(), + }, + }; +}); + +vi.mock("./FallbackSelectionForm", () => ({ + FallbackSelectionForm: ({ groups, onGroupsChange }: any) => { + const handleUpdateGroup = () => { + if (groups.length > 0) { + const updatedGroups = groups.map((group: any, index: number) => { + if (index === 0 && !group.primaryModel) { + return { + ...group, + primaryModel: "gpt-4", + fallbackModels: ["gpt-3.5-turbo"], + }; + } + return group; + }); + onGroupsChange(updatedGroups); + } + }; + + return ( +
+ +
{groups.length}
+ {groups.map((group: any) => ( +
+ Primary: {group.primaryModel || "None"}, Fallbacks: {group.fallbackModels.length} +
+ ))} +
+ ); + }, +})); + +describe("AddFallbacks", () => { + const mockOnChange = vi.fn(); + const mockAccessToken = "test-token"; + const mockModelGroups = [ + { model_group: "gpt-4", mode: "chat" }, + { model_group: "gpt-3.5-turbo", mode: "chat" }, + { model_group: "claude-3-opus", mode: "chat" }, + ]; + + const defaultProps = { + accessToken: mockAccessToken, + value: [] as Fallbacks, + onChange: mockOnChange, + }; + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValue(mockModelGroups); + }); + + it("should render the component", () => { + render(); + expect(screen.getByRole("button", { name: /add fallbacks/i })).toBeInTheDocument(); + }); + + it("should open modal when Add Fallbacks button is clicked", async () => { + const user = userEvent.setup(); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + }); + + it("should fetch available models when modal opens", async () => { + const user = userEvent.setup(); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(fetchModelsModule.fetchAvailableModels).toHaveBeenCalledWith(mockAccessToken); + }); + }); + + it("should close modal when Cancel button is clicked", async () => { + const user = userEvent.setup(); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + const cancelButton = screen.getByRole("button", { name: /cancel/i }); + await user.click(cancelButton); + + await waitFor(() => { + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + }); + }); + + it("should show error when saving incomplete groups", async () => { + const user = userEvent.setup(); + const antd = await import("antd"); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + await waitFor(() => { + const saveButton = screen.getByRole("button", { name: /save all configurations/i }); + expect(saveButton).toBeInTheDocument(); + }); + + const saveButton = screen.getByRole("button", { name: /save all configurations/i }); + await user.click(saveButton); + + await waitFor(() => { + expect(antd.message.error).toHaveBeenCalled(); + }); + }); + + it("should show error message when saving incomplete groups", async () => { + const user = userEvent.setup(); + const antd = await import("antd"); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + const saveButton = screen.getByRole("button", { name: /save all configurations/i }); + await user.click(saveButton); + + await waitFor(() => { + expect(antd.message.error).toHaveBeenCalled(); + }); + }); + + it("should call onChange with new fallbacks when Save is clicked with valid configuration", async () => { + const user = userEvent.setup(); + mockOnChange.mockResolvedValue(undefined); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + await waitFor(() => { + expect(screen.getByTestId("fallback-selection-form")).toBeInTheDocument(); + }); + + const updateGroupButton = screen.getByTestId("update-group-button"); + await user.click(updateGroupButton); + + await waitFor(() => { + expect(screen.getByText(/Primary: gpt-4/i)).toBeInTheDocument(); + }); + + const saveButton = screen.getByRole("button", { name: /save all configurations/i }); + expect(saveButton).not.toBeDisabled(); + await user.click(saveButton); + + await waitFor(() => { + expect(mockOnChange).toHaveBeenCalled(); + const callArgs = mockOnChange.mock.calls[0][0]; + expect(callArgs).toHaveLength(1); + expect(callArgs[0]).toHaveProperty("gpt-4"); + expect(callArgs[0]["gpt-4"]).toContain("gpt-3.5-turbo"); + }); + }); + + it("should append new fallbacks to existing value", async () => { + const user = userEvent.setup(); + const existingFallbacks: Fallbacks = [{ "existing-model": ["fallback-1"] }]; + mockOnChange.mockResolvedValue(undefined); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + await waitFor(() => { + expect(screen.getByTestId("fallback-selection-form")).toBeInTheDocument(); + }); + + const updateGroupButton = screen.getByTestId("update-group-button"); + await user.click(updateGroupButton); + + await waitFor(() => { + expect(screen.getByText(/Primary: gpt-4/i)).toBeInTheDocument(); + }); + + const saveButton = screen.getByRole("button", { name: /save all configurations/i }); + await user.click(saveButton); + + await waitFor(() => { + expect(mockOnChange).toHaveBeenCalled(); + const callArgs = mockOnChange.mock.calls[0][0]; + expect(callArgs).toHaveLength(2); + expect(callArgs[0]).toEqual({ "existing-model": ["fallback-1"] }); + expect(callArgs[1]).toHaveProperty("gpt-4"); + }); + }); + + it("should reset form state when modal is closed", async () => { + const user = userEvent.setup(); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + const cancelButton = screen.getByRole("button", { name: /cancel/i }); + await user.click(cancelButton); + + await waitFor(() => { + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + }); + + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + expect(screen.getByTestId("fallback-selection-form")).toBeInTheDocument(); + }); + }); + + it("should handle onChange error gracefully", async () => { + const user = userEvent.setup(); + const error = new Error("Save failed"); + mockOnChange.mockRejectedValue(error); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + + await waitFor(() => { + expect(screen.getByTestId("fallback-selection-form")).toBeInTheDocument(); + }); + + const updateGroupButton = screen.getByTestId("update-group-button"); + await user.click(updateGroupButton); + + await waitFor(() => { + expect(screen.getByText(/Primary: gpt-4/i)).toBeInTheDocument(); + }); + + const saveButton = screen.getByRole("button", { name: /save all configurations/i }); + await user.click(saveButton); + + await waitFor(() => { + expect(mockOnChange).toHaveBeenCalled(); + }); + }); + + it("should not call onChange when onChange prop is not provided", async () => { + const user = userEvent.setup(); + render(); + + const addButton = screen.getByRole("button", { name: /add fallbacks/i }); + await user.click(addButton); + + await waitFor(() => { + expect(screen.getByRole("dialog")).toBeInTheDocument(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.tsx new file mode 100644 index 00000000000..c0b626cbb12 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.tsx @@ -0,0 +1,169 @@ +/** + * Parent component for adding fallbacks to the proxy router config + * Handles value/onChange logic and form submission + * Works with forms - reads from and writes to router_settings.fallbacks + */ + +import { Button as TremorButton } from "@tremor/react"; +import { Button, message } from "antd"; +import React, { useEffect, useState } from "react"; +import NotificationManager from "../../../molecules/notifications_manager"; +import { fetchAvailableModels, ModelGroup } from "../../../playground/llm_calls/fetch_models"; +import { AddFallbacksModal } from "./AddFallbacksModal"; +import { FallbackGroup } from "./FallbackGroupConfig"; +import { FallbackSelectionForm } from "./FallbackSelectionForm"; + +export type FallbackEntry = { [modelName: string]: string[] }; +export type Fallbacks = FallbackEntry[]; + +interface AddFallbacksProps { + models?: string[]; + accessToken: string; + value?: Fallbacks; // Current fallbacks value from form + onChange?: (fallbacks: Fallbacks) => Promise; // Callback to update form value +} + +export default function AddFallbacks({ + models, + accessToken, + value = [], + onChange, +}: AddFallbacksProps) { + const [isModalVisible, setIsModalVisible] = useState(false); + const [modelInfo, setModelInfo] = useState([]); + const [modalKey, setModalKey] = useState(0); // Key to force remount of form when modal opens + const [isSaving, setIsSaving] = useState(false); + const [groups, setGroups] = useState([ + { + id: "1", + primaryModel: null, + fallbackModels: [], + }, + ]); + + // Reset groups state and increment modal key when modal opens + useEffect(() => { + if (isModalVisible) { + setGroups([ + { + id: "1", + primaryModel: null, + fallbackModels: [], + }, + ]); + setModalKey((prev) => prev + 1); // Force remount of form + } + }, [isModalVisible]); + + useEffect(() => { + const loadModels = async () => { + try { + const uniqueModels = await fetchAvailableModels(accessToken); + console.log("Fetched models for fallbacks:", uniqueModels); + setModelInfo(uniqueModels); + } catch (error) { + console.error("Error fetching model info for fallbacks:", error); + } + }; + if (isModalVisible) { + loadModels(); + } + }, [accessToken, isModalVisible]); + + const availableModels = Array.from(new Set(modelInfo.map((option) => option.model_group))).sort(); + + const handleCancel = () => { + setIsModalVisible(false); + // Reset to initial state + setGroups([ + { + id: "1", + primaryModel: null, + fallbackModels: [], + }, + ]); + }; + + const handleSaveAll = async () => { + // Validation + const invalidGroups = groups.filter( + (g) => !g.primaryModel || g.fallbackModels.length === 0, + ); + if (invalidGroups.length > 0) { + message.error( + `Please complete configuration for all groups. ${invalidGroups.length} group(s) incomplete.`, + ); + return; + } + + // Create fallback objects in the format expected by the API + const newFallbacks = groups.map((g) => ({ + [g.primaryModel!]: g.fallbackModels, + })); + + // Get current fallbacks from form value, or an empty array if it's null/undefined + const currentFallbacks = value || []; + + // Add new fallbacks to the current fallbacks + const updatedFallbacks = [...currentFallbacks, ...newFallbacks]; + + // Call onChange to update the form value and wait for it to complete + if (onChange) { + setIsSaving(true); + try { + await onChange(updatedFallbacks); + NotificationManager.success(`${groups.length} fallback configuration(s) added successfully!`); + handleCancel(); + } catch (error) { + // Error handling is done in handleFallbacksChange, so we don't need to show another notification here + console.error("Error saving fallbacks:", error); + } finally { + setIsSaving(false); + } + } else { + NotificationManager.fromBackend("onChange callback not provided"); + } + }; + + return ( +
+ setIsModalVisible(true)} + icon={() => +} + > + Add Fallbacks + + + + {/* Footer with Cancel and Save buttons */} + {groups.length > 0 && ( +
+ + +
+ )} +
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.test.tsx new file mode 100644 index 00000000000..c6be52f7c14 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.test.tsx @@ -0,0 +1,56 @@ +import { render, screen } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { AddFallbacksModal } from "./AddFallbacksModal"; + +describe("AddFallbacksModal", () => { + const mockOnCancel = vi.fn(); + + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render the modal when open is true", () => { + render( + +
Test Content
+
, + ); + + expect(screen.getByRole("dialog")).toBeInTheDocument(); + expect(screen.getByText("Configure Model Fallbacks")).toBeInTheDocument(); + expect(screen.getByText("Manage multiple fallback chains for different models (up to 5 groups at a time)")).toBeInTheDocument(); + expect(screen.getByText("Test Content")).toBeInTheDocument(); + }); + + it("should not render the modal when open is false", () => { + render( + +
Test Content
+
, + ); + + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + }); + + it("should render children content when modal is open", () => { + render( + +
Child Component
+
, + ); + + expect(screen.getByTestId("child-content")).toBeInTheDocument(); + expect(screen.getByText("Child Component")).toBeInTheDocument(); + }); + + it("should display the correct title and description", () => { + render( + +
Content
+
, + ); + + expect(screen.getByText("Configure Model Fallbacks")).toBeInTheDocument(); + expect(screen.getByText(/Manage multiple fallback chains/i)).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.tsx new file mode 100644 index 00000000000..5048c99d817 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.tsx @@ -0,0 +1,52 @@ +/** + * Modal wrapper for the fallback selection form + * Handles modal visibility and layout, but delegates content to children + */ + +import { Modal } from "antd"; +import { ArrowRight } from "lucide-react"; +import React from "react"; + +interface AddFallbacksModalProps { + open: boolean; + onCancel: () => void; + children: React.ReactNode; +} + +export function AddFallbacksModal({ + open, + onCancel, + children, +}: AddFallbacksModalProps) { + return ( + +
+
+ +
+
+

Configure Model Fallbacks

+

+ Manage multiple fallback chains for different models (up to 5 groups at a time) +

+
+
+ + } + open={open} + width={900} + footer={null} + onCancel={onCancel} + maskClosable={false} + className="top-8" + styles={{ + body: { padding: "24px" }, + header: { padding: "24px 24px 0 24px", border: "none" }, + }} + > +
{children}
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx new file mode 100644 index 00000000000..24ce80f0a0b --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx @@ -0,0 +1,208 @@ +/** + * Component for configuring a single fallback group + * Handles primary model selection and fallback chain configuration + */ + +import { Select, Tooltip } from "antd"; +import { AlertCircle, ArrowDown, X } from "lucide-react"; +import React from "react"; + +export interface FallbackGroup { + id: string; + primaryModel: string | null; + fallbackModels: string[]; +} + +interface FallbackGroupConfigProps { + group: FallbackGroup; + onChange: (updatedGroup: FallbackGroup) => void; + availableModels: string[]; + maxFallbacks: number; +} + +export function FallbackGroupConfig({ + group, + onChange, + availableModels, + maxFallbacks, +}: FallbackGroupConfigProps) { + // Filter available options for fallbacks (exclude primary only, allow already selected to be shown for deselection) + const availableFallbackOptions = availableModels.filter( + (m) => m !== group.primaryModel, + ); + + const handlePrimaryChange = (value: string) => { + let newFallbacks = [...group.fallbackModels]; + // Remove from fallbacks if it was there + if (newFallbacks.includes(value)) { + newFallbacks = newFallbacks.filter((m) => m !== value); + } + onChange({ + ...group, + primaryModel: value, + fallbackModels: newFallbacks, + }); + }; + + const handleFallbackSelect = (values: string[]) => { + // Limit to maxFallbacks + const limitedValues = values.slice(0, maxFallbacks); + + onChange({ + ...group, + fallbackModels: limitedValues, + }); + }; + + const removeFallback = (indexToRemove: number) => { + const newFallbacks = group.fallbackModels.filter((_, index) => index !== indexToRemove); + onChange({ + ...group, + fallbackModels: newFallbacks, + }); + }; + + const canAddMoreFallbacks = group.fallbackModels.length < maxFallbacks; + + return ( +
+ {/* Primary Model Section */} +
+ + ({ + label: m, + value: m, + }))} + optionRender={(option, info) => { + const isSelected = group.fallbackModels.includes(option.value as string); + const orderIndex = isSelected + ? group.fallbackModels.indexOf(option.value as string) + 1 + : null; + return ( +
+ {isSelected && orderIndex !== null && ( + + {orderIndex} + + )} + {option.label} +
+ ); + }} + maxTagCount="responsive" + maxTagPlaceholder={(omittedValues) => ( + value).join(", ")} + > + +{omittedValues.length} more + + )} + showSearch + filterOption={(input, option) => + (option?.label ?? "").toLowerCase().includes(input.toLowerCase()) + } + /> +

+ {canAddMoreFallbacks + ? `Search and select multiple models. Selected models will appear below in order. (${group.fallbackModels.length}/${maxFallbacks} used)` + : `Maximum ${maxFallbacks} fallbacks reached. Remove some to add more.`} +

+
+ + {/* Fallback List */} +
+ {group.fallbackModels.length === 0 ? ( +
+ No fallback models selected + Add models from the dropdown above +
+ ) : ( + group.fallbackModels.map((modelValue, index) => { + return ( +
+
+
+ {index + 1} +
+
+ {modelValue} +
+
+ + +
+ ); + }) + )} +
+
+ + + ); +} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx new file mode 100644 index 00000000000..bb9ccb312a6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx @@ -0,0 +1,132 @@ +/** + * Form component for selecting and configuring fallback groups + * Manages groups state internally, but does not handle submission + * Decoupled from form submission logic + */ + +import { Button } from "@tremor/react"; +import { message, Tabs } from "antd"; +import { Plus } from "lucide-react"; +import React, { useEffect, useState } from "react"; +import { FallbackGroup, FallbackGroupConfig } from "./FallbackGroupConfig"; + +interface FallbackSelectionFormProps { + groups: FallbackGroup[]; + onGroupsChange: (groups: FallbackGroup[]) => void; + availableModels: string[]; + maxFallbacks?: number; + maxGroups?: number; +} + +export function FallbackSelectionForm({ + groups, + onGroupsChange, + availableModels, + maxFallbacks = 5, + maxGroups = 5, +}: FallbackSelectionFormProps) { + const [activeKey, setActiveKey] = useState(groups.length > 0 ? groups[0].id : "1"); + + // Reset activeKey when groups change (e.g., when modal reopens) + useEffect(() => { + if (groups.length > 0) { + // If current activeKey doesn't exist in groups, reset to first group + const activeKeyExists = groups.some((g) => g.id === activeKey); + if (!activeKeyExists) { + setActiveKey(groups[0].id); + } + } else { + // If groups is empty, reset activeKey + setActiveKey("1"); + } + }, [groups]); + + const handleAddGroup = () => { + if (groups.length >= maxGroups) { + return; + } + const newId = Date.now().toString(); + const newGroups = [ + ...groups, + { + id: newId, + primaryModel: null, + fallbackModels: [], + }, + ]; + onGroupsChange(newGroups); + setActiveKey(newId); + }; + + const handleRemoveGroup = (targetId: string) => { + if (groups.length === 1) { + message.warning("At least one group is required"); + return; + } + const newGroups = groups.filter((g) => g.id !== targetId); + onGroupsChange(newGroups); + if (activeKey === targetId && newGroups.length > 0) { + setActiveKey(newGroups[newGroups.length - 1].id); + } + }; + + const handleGroupUpdate = (updatedGroup: FallbackGroup) => { + const newGroups = groups.map((g) => (g.id === updatedGroup.id ? updatedGroup : g)); + onGroupsChange(newGroups); + }; + + // Generate tab items + const items = groups.map((group, index) => { + const label = group.primaryModel + ? group.primaryModel + : `Group ${index + 1}`; + return { + key: group.id, + label: label, + closable: groups.length > 1, // Only allow closing if there's more than 1 group + children: ( + + ), + }; + }); + + if (groups.length === 0) { + return ( +
+

No fallback groups configured

+ +
+ ); + } + + return ( + { + if (action === "add") handleAddGroup(); + else if (action === "remove" && groups.length > 1) { + handleRemoveGroup(targetKey as string); + } + }} + items={items} + className="fallback-tabs" + tabBarStyle={{ + marginBottom: 0, + }} + hideAdd={groups.length >= maxGroups} + /> + ); +} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx new file mode 100644 index 00000000000..bc6f6a4856c --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.test.tsx @@ -0,0 +1,372 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import Fallbacks from "./fallbacks"; +import * as networkingModule from "../../../networking"; +import * as fetchModelsModule from "../../../playground/llm_calls/fetch_models"; + +vi.mock("../../../networking", () => ({ + getCallbacksCall: vi.fn(), + setCallbacksCall: vi.fn(), +})); + +vi.mock("../../../playground/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn(), +})); + +vi.mock("openai", () => ({ + default: { + OpenAI: vi.fn().mockImplementation(() => ({ + chat: { + completions: { + create: vi.fn(), + }, + }, + })), + }, +})); + +vi.mock("../../../common_components/DeleteResourceModal", () => ({ + __esModule: true, + default: ({ isOpen, onOk, onCancel, title, message, resourceInformation, confirmLoading }: any) => { + if (!isOpen) return null; + return ( +
+
{title}
+
{message}
+ {resourceInformation?.map((info: any, idx: number) => ( +
+ {info.label}: {info.value} +
+ ))} + + +
+ ); + }, +})); + +vi.mock("./AddFallbacks", () => ({ + __esModule: true, + default: ({ value, onChange }: any) => { + const handleClick = async () => { + if (onChange) { + try { + const newFallbacks = [...(value || []), { "test-model": ["test-fallback"] }]; + await onChange(newFallbacks); + } catch (error) { + // Error is handled by the component + } + } + }; + return ( + + ); + }, +})); + +describe("Fallbacks", () => { + const mockAccessToken = "test-token"; + const mockUserRole = "Admin"; + const mockUserID = "user-123"; + const mockModelData = { + data: [ + { model_name: "gpt-4" }, + { model_name: "gpt-3.5-turbo" }, + { model_name: "claude-3-opus" }, + ], + }; + + const mockRouterSettings = { + fallbacks: [ + { "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] }, + { "claude-3-opus": ["gpt-4"] }, + ], + }; + + const defaultProps = { + accessToken: mockAccessToken, + userRole: mockUserRole, + userID: mockUserID, + modelData: mockModelData, + }; + + const findDeleteButton = (container: HTMLElement) => { + const tableRows = container.querySelectorAll("tbody tr"); + if (tableRows.length === 0) return null; + const firstRow = tableRows[0]; + const actionCells = firstRow.querySelectorAll("td"); + const lastCell = actionCells[actionCells.length - 1]; + const buttons = lastCell.querySelectorAll("button"); + if (buttons.length >= 2) { + return buttons[buttons.length - 1]; + } + const clickableElements = lastCell.querySelectorAll("[class*='cursor-pointer'], button"); + return Array.from(clickableElements).find((el) => + el.className.includes("red") || el.className.includes("hover:text-red") + ) || clickableElements[clickableElements.length - 1]; + }; + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(networkingModule.getCallbacksCall).mockResolvedValue({ + router_settings: mockRouterSettings, + }); + vi.mocked(networkingModule.setCallbacksCall).mockResolvedValue(undefined); + vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValue([ + { model_group: "gpt-4", mode: "chat" }, + { model_group: "gpt-3.5-turbo", mode: "chat" }, + { model_group: "claude-3-opus", mode: "chat" }, + ]); + }); + + it("should render the component", async () => { + render(); + + await waitFor(() => { + expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); + }); + }); + + it("should not render when accessToken is null", () => { + const { container } = render(); + expect(container.firstChild).toBeNull(); + }); + + it("should fetch router settings on mount", async () => { + render(); + + await waitFor(() => { + expect(networkingModule.getCallbacksCall).toHaveBeenCalledWith( + mockAccessToken, + mockUserID, + mockUserRole, + ); + }); + }); + + it("should display fallback entries in table", async () => { + render(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + expect(screen.getByText("gpt-3.5-turbo, claude-3-opus")).toBeInTheDocument(); + expect(screen.getByText("claude-3-opus")).toBeInTheDocument(); + }); + }); + + it("should open delete modal when delete icon is clicked", async () => { + const user = userEvent.setup(); + const { container } = render(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + }); + + const deleteButton = findDeleteButton(container); + expect(deleteButton).not.toBeNull(); + + await user.click(deleteButton as HTMLElement); + + await waitFor(() => { + expect(screen.getByTestId("delete-modal")).toBeInTheDocument(); + expect(screen.getByText("Delete Fallback?")).toBeInTheDocument(); + }); + }); + + it("should delete fallback when confirmed", async () => { + const user = userEvent.setup(); + const { container } = render(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + }); + + const deleteButton = findDeleteButton(container); + expect(deleteButton).not.toBeNull(); + + await user.click(deleteButton as HTMLElement); + + await waitFor(() => { + expect(screen.getByTestId("delete-modal")).toBeInTheDocument(); + }); + + const confirmButton = screen.getByRole("button", { name: /delete/i }); + await user.click(confirmButton); + + await waitFor(() => { + expect(networkingModule.setCallbacksCall).toHaveBeenCalled(); + const callArgs = networkingModule.setCallbacksCall.mock.calls[0]; + expect(callArgs[0]).toBe(mockAccessToken); + expect(callArgs[1].router_settings.fallbacks).toHaveLength(1); + }); + }); + + it("should close delete modal when cancel is clicked", async () => { + const user = userEvent.setup(); + const { container } = render(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + }); + + const deleteButton = findDeleteButton(container); + expect(deleteButton).not.toBeNull(); + + await user.click(deleteButton as HTMLElement); + + await waitFor(() => { + expect(screen.getByTestId("delete-modal")).toBeInTheDocument(); + }); + + const cancelButton = screen.getByRole("button", { name: /cancel/i }); + await user.click(cancelButton); + + await waitFor(() => { + expect(screen.queryByTestId("delete-modal")).not.toBeInTheDocument(); + }); + }); + + it("should show error notification on delete failure", async () => { + const user = userEvent.setup(); + const error = new Error("Delete failed"); + vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error); + const { container } = render(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + }); + + const deleteButton = findDeleteButton(container); + expect(deleteButton).not.toBeNull(); + + await user.click(deleteButton as HTMLElement); + + await waitFor(() => { + expect(screen.getByTestId("delete-modal")).toBeInTheDocument(); + }); + + const confirmButton = screen.getByRole("button", { name: /delete/i }); + await user.click(confirmButton); + + await waitFor(() => { + expect(networkingModule.setCallbacksCall).toHaveBeenCalled(); + }); + }); + + it("should handle delete error gracefully", async () => { + const user = userEvent.setup(); + const error = new Error("Delete failed"); + vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error); + const { container } = render(); + + await waitFor(() => { + expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); + }); + + const deleteButton = findDeleteButton(container); + expect(deleteButton).not.toBeNull(); + + await user.click(deleteButton as HTMLElement); + + await waitFor(() => { + expect(screen.getByTestId("delete-modal")).toBeInTheDocument(); + }); + + const confirmButton = screen.getByRole("button", { name: /delete/i }); + await user.click(confirmButton); + + await waitFor(() => { + expect(networkingModule.setCallbacksCall).toHaveBeenCalled(); + expect(screen.queryByTestId("delete-modal")).not.toBeInTheDocument(); + }); + }); + + it("should handle empty fallbacks array", async () => { + vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({ + router_settings: { fallbacks: [] }, + }); + render(); + + await waitFor(() => { + expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); + }); + + expect(screen.queryByText("gpt-4")).not.toBeInTheDocument(); + }); + + it("should handle router settings without fallbacks property", async () => { + vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({ + router_settings: {}, + }); + render(); + + await waitFor(() => { + expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); + }); + }); + + it("should remove model_group_retry_policy from router settings", async () => { + vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({ + router_settings: { + ...mockRouterSettings, + model_group_retry_policy: { some: "policy" }, + }, + }); + render(); + + await waitFor(() => { + expect(networkingModule.getCallbacksCall).toHaveBeenCalled(); + }); + }); + + it("should update fallbacks when AddFallbacks onChange is called", async () => { + const user = userEvent.setup(); + render(); + + await waitFor(() => { + expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); + }); + + const addButton = screen.getByTestId("add-fallbacks-button"); + await user.click(addButton); + + await waitFor(() => { + expect(networkingModule.setCallbacksCall).toHaveBeenCalled(); + }); + }); + + it("should handle fallbacks change error and refetch", async () => { + const user = userEvent.setup(); + const error = new Error("Update failed"); + vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error); + vi.mocked(networkingModule.getCallbacksCall).mockResolvedValue({ + router_settings: mockRouterSettings, + }); + render(); + + await waitFor(() => { + expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument(); + }); + + const addButton = screen.getByTestId("add-fallbacks-button"); + await user.click(addButton); + + await waitFor(() => { + expect(networkingModule.setCallbacksCall).toHaveBeenCalled(); + }); + + await waitFor( + () => { + expect(networkingModule.getCallbacksCall).toHaveBeenCalledTimes(2); + }, + { timeout: 3000 }, + ); + }); +}); diff --git a/ui/litellm-dashboard/src/components/fallbacks.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx similarity index 81% rename from ui/litellm-dashboard/src/components/fallbacks.tsx rename to ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx index dc90ae79136..f493cc51323 100644 --- a/ui/litellm-dashboard/src/components/fallbacks.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx @@ -3,10 +3,10 @@ import { Icon, Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow import { Tooltip } from "antd"; import openai from "openai"; import React, { useEffect, useState } from "react"; -import AddFallbacks from "./add_fallbacks"; -import DeleteResourceModal from "./common_components/DeleteResourceModal"; -import NotificationsManager from "./molecules/notifications_manager"; -import { getCallbacksCall, setCallbacksCall } from "./networking"; +import DeleteResourceModal from "../../../common_components/DeleteResourceModal"; +import NotificationsManager from "../../../molecules/notifications_manager"; +import { getCallbacksCall, setCallbacksCall } from "../../../networking"; +import AddFallbacks from "./AddFallbacks"; type FallbackEntry = { [modelName: string]: string[] }; type Fallbacks = FallbackEntry[]; @@ -21,7 +21,7 @@ interface FallbacksProps { async function testFallbackModelResponse(selectedModel: string, accessToken: string) { const isLocal = process.env.NODE_ENV === "development"; if (isLocal != true) { - console.log = function () {}; + console.log = function () { }; } const proxyBaseUrl = isLocal ? "http://localhost:4000" : window.location.origin; const client = new openai.OpenAI({ @@ -142,13 +142,48 @@ const Fallbacks: React.FC = ({ accessToken, userRole, userID, mo return null; } + const handleFallbacksChange = async (fallbacks: Fallbacks): Promise => { + if (!accessToken) { + return; + } + + const updatedSettings = { + ...routerSettings, + fallbacks: fallbacks, + }; + + const payload = { + router_settings: updatedSettings, + }; + + try { + await setCallbacksCall(accessToken, payload); + // Update UI only after successful API call + setRouterSettings(updatedSettings); + } catch (error) { + // Revert on error by refetching from server + NotificationsManager.fromBackend("Failed to update router settings: " + error); + if (accessToken && userRole && userID) { + getCallbacksCall(accessToken, userID, userRole).then((data) => { + let router_settings = data.router_settings; + if ("model_group_retry_policy" in router_settings) { + delete router_settings["model_group_retry_policy"]; + } + setRouterSettings(router_settings); + }); + } + // Re-throw error so caller can handle it + throw error; + } + }; + return ( <> data.model_name) : []} - accessToken={accessToken} - routerSettings={routerSettings} - setRouterSettings={setRouterSettings} + accessToken={accessToken || ""} + value={routerSettings.fallbacks || []} + onChange={handleFallbacksChange} /> diff --git a/ui/litellm-dashboard/src/components/add_fallbacks.test.tsx b/ui/litellm-dashboard/src/components/add_fallbacks.test.tsx deleted file mode 100644 index 102b55af8af..00000000000 --- a/ui/litellm-dashboard/src/components/add_fallbacks.test.tsx +++ /dev/null @@ -1,47 +0,0 @@ -import { render, screen } from "@testing-library/react"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import AddFallbacks from "./add_fallbacks"; - -vi.mock("./networking", () => ({ - setCallbacksCall: vi.fn(), -})); - -vi.mock("./playground/llm_calls/fetch_models", () => ({ - fetchAvailableModels: vi.fn(() => - Promise.resolve([ - { model_group: "gpt-4" }, - { model_group: "gpt-3.5-turbo" }, - { model_group: "claude-3-opus" }, - { model_group: "claude-3-sonnet" }, - ]), - ), -})); - -vi.mock("./molecules/notifications_manager", () => ({ - default: { - success: vi.fn(), - fromBackend: vi.fn(), - }, -})); - -describe("AddFallbacks", () => { - const mockAccessToken = "test-token"; - const mockRouterSettings = { fallbacks: [] }; - const mockSetRouterSettings = vi.fn(); - - beforeEach(() => { - vi.clearAllMocks(); - }); - - it("should render the component", () => { - render( - , - ); - - expect(screen.getByRole("button", { name: /Add Fallbacks/i })).toBeInTheDocument(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/add_fallbacks.tsx b/ui/litellm-dashboard/src/components/add_fallbacks.tsx deleted file mode 100644 index 97b020005ae..00000000000 --- a/ui/litellm-dashboard/src/components/add_fallbacks.tsx +++ /dev/null @@ -1,250 +0,0 @@ -/** - * Modal to add fallbacks to the proxy router config - */ - -import { Button } from "@tremor/react"; -import { Form, Modal, Select } from "antd"; -import React, { useEffect, useState } from "react"; -import NotificationManager from "./molecules/notifications_manager"; -import { setCallbacksCall } from "./networking"; -import { fetchAvailableModels, ModelGroup } from "./playground/llm_calls/fetch_models"; - -interface AddFallbacksProps { - models?: string[]; - accessToken: string; - routerSettings: { [key: string]: any }; - setRouterSettings: React.Dispatch>; -} - -const AddFallbacks: React.FC = ({ models, accessToken, routerSettings, setRouterSettings }) => { - const [form] = Form.useForm(); - const [isModalVisible, setIsModalVisible] = useState(false); - const [selectedModel, setSelectedModel] = useState(""); - const [modelInfo, setModelInfo] = useState([]); - const [selectedFallbacks, setSelectedFallbacks] = useState([]); - - useEffect(() => { - const loadModels = async () => { - try { - const uniqueModels = await fetchAvailableModels(accessToken); - console.log("Fetched models for fallbacks:", uniqueModels); - setModelInfo(uniqueModels); - } catch (error) { - console.error("Error fetching model info for fallbacks:", error); - } - }; - loadModels(); - }, [accessToken]); - const handleOk = () => { - setIsModalVisible(false); - form.resetFields(); - setSelectedFallbacks([]); - setSelectedModel(""); - }; - - const handleCancel = () => { - setIsModalVisible(false); - form.resetFields(); - setSelectedFallbacks([]); - setSelectedModel(""); - }; - - const updateFallbacks = (formValues: Record) => { - // Print the received value - console.log(formValues); - - // Extract model_name and models from formValues - const { model_name, models } = formValues; - - // Create new fallback - const newFallback = { [model_name]: models }; - - // Get current fallbacks, or an empty array if it's null - const currentFallbacks = routerSettings.fallbacks || []; - - // Add new fallback to the current fallbacks - const updatedFallbacks = [...currentFallbacks, newFallback]; - - // Create a new routerSettings object with updated fallbacks - const updatedRouterSettings = { ...routerSettings, fallbacks: updatedFallbacks }; - - // Print updated routerSettings - console.log(updatedRouterSettings); - - const payload = { - router_settings: updatedRouterSettings, - }; - - try { - setCallbacksCall(accessToken, payload); - // Update routerSettings state - setRouterSettings(updatedRouterSettings); - } catch (error) { - NotificationManager.fromBackend("Failed to update router settings: " + error); - } - - NotificationManager.success("router settings updated successfully"); - - setIsModalVisible(false); - form.resetFields(); - setSelectedFallbacks([]); - setSelectedModel(""); - }; - - return ( -
- - -

Add Fallbacks

-
- } - open={isModalVisible} - width={900} - footer={null} - onOk={handleOk} - onCancel={handleCancel} - className="top-8" - styles={{ - body: { padding: "24px" }, - header: { padding: "24px 24px 0 24px", border: "none" }, - }} - > -
-
-

- Configure fallback models to improve reliability. When the primary model fails or is unavailable, requests - will automatically route to the specified fallback models in order. -

-
- -
-
- - Primary Model * - - } - name="model_name" - rules={[{ required: true, message: "Please select the primary model that needs fallbacks" }]} - className="!mb-0" - > - -

This is the primary model that users will request

-
- -
- - - Fallback Models (select multiple) * - - } - name="models" - rules={[{ required: true, message: "Please select at least one fallback model" }]} - className="!mb-0" - > -
- {/* Show selected models in order */} - {selectedFallbacks.length > 0 && ( -
-

Fallback Order:

-
- {selectedFallbacks.map((model, index) => ( -
- {index + 1}. - {model} - -
- ))} -
-
- )} - - {/* Model selector */} - -
-

- Order matters: Models will be tried in the order shown above (1st, 2nd, 3rd, etc.) -

-
-
- -
- - -
- -
- - - ); -}; - -export default AddFallbacks; diff --git a/ui/litellm-dashboard/src/components/fallbacks.test.tsx b/ui/litellm-dashboard/src/components/fallbacks.test.tsx deleted file mode 100644 index e6db11270cd..00000000000 --- a/ui/litellm-dashboard/src/components/fallbacks.test.tsx +++ /dev/null @@ -1,95 +0,0 @@ -import { render, screen, waitFor } from "@testing-library/react"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import Fallbacks from "./fallbacks"; -import { getCallbacksCall, setCallbacksCall } from "./networking"; - -vi.mock("./networking", () => ({ - getCallbacksCall: vi.fn(), - setCallbacksCall: vi.fn(), -})); - -vi.mock("./add_fallbacks", () => ({ - __esModule: true, - default: () =>
Mock Add Fallbacks
, -})); - -vi.mock("openai", () => ({ - default: { - OpenAI: vi.fn().mockImplementation(() => ({ - chat: { - completions: { - create: vi.fn().mockResolvedValue({ - model: "test-model", - }), - }, - }, - })), - }, -})); - -describe("Fallbacks", () => { - const defaultProps = { - accessToken: "token", - userRole: "admin", - userID: "user-123", - modelData: { data: [] }, - }; - const mockGetCallbacksCall = vi.mocked(getCallbacksCall); - const mockSetCallbacksCall = vi.mocked(setCallbacksCall); - - beforeEach(() => { - vi.clearAllMocks(); - mockGetCallbacksCall.mockResolvedValue({ - router_settings: { - fallbacks: [], - }, - }); - mockSetCallbacksCall.mockResolvedValue({}); - }); - - it("should render", async () => { - render(); - - await waitFor(() => { - expect(screen.getByRole("columnheader", { name: "Model Name" })).toBeInTheDocument(); - }); - }); - - it("should render fallback data when callback data is returned from network call", async () => { - const mockFallbackData = { - router_settings: { - fallbacks: [{ "xai/grok-2": ["xai/grok-4", "gpt-4"] }, { "gpt-3.5-turbo": ["gpt-4"] }], - }, - }; - - mockGetCallbacksCall.mockResolvedValue(mockFallbackData); - - render(); - - await waitFor(() => { - expect(screen.getByText("xai/grok-2")).toBeInTheDocument(); - expect(screen.getByText("xai/grok-4, gpt-4")).toBeInTheDocument(); - expect(screen.getByText("gpt-3.5-turbo")).toBeInTheDocument(); - expect(screen.getByText("gpt-4")).toBeInTheDocument(); - }); - - expect(mockGetCallbacksCall).toHaveBeenCalledWith( - defaultProps.accessToken, - defaultProps.userID, - defaultProps.userRole, - ); - }); - - it("should render AddFallbacks component", async () => { - render(); - - await waitFor(() => { - expect(screen.getByText("Mock Add Fallbacks")).toBeInTheDocument(); - }); - }); - - it("should not render when access token is not provided", () => { - const { container } = render(); - expect(container.firstChild).toBeNull(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/general_settings.tsx b/ui/litellm-dashboard/src/components/general_settings.tsx index 1dc35550558..18f891705e1 100644 --- a/ui/litellm-dashboard/src/components/general_settings.tsx +++ b/ui/litellm-dashboard/src/components/general_settings.tsx @@ -22,7 +22,7 @@ import { InputNumber } from "antd"; import { TrashIcon, CheckCircleIcon } from "@heroicons/react/outline"; import RouterSettings from "./router_settings"; -import Fallbacks from "./fallbacks"; +import Fallbacks from "./Settings/RouterSettings/Fallbacks/Fallbacks"; interface GeneralSettingsPageProps { accessToken: string | null; userRole: string | null; diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 1aa1df14a10..543d95d2ccd 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -3,6 +3,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { ApiOutlined, AppstoreOutlined, + AuditOutlined, BankOutlined, BarChartOutlined, BgColorsOutlined, @@ -123,6 +124,13 @@ const Sidebar: React.FC = ({ setPage, defaultSelectedKey, collapse icon: , roles: all_admin_roles, }, + { + key: "policies", + page: "policies", + label: "Policies", + icon: , + roles: all_admin_roles, + }, { key: "tools", page: "tools", diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a4f2dfce3bd..ffc0f80911d 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5347,6 +5347,249 @@ export const getGuardrailsList = async (accessToken: string) => { } }; +// ───────────────────────────────────────────────────────────────────────────── +// Policy CRUD API Calls +// ───────────────────────────────────────────────────────────────────────────── + +export const getPoliciesList = async (accessToken: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/list` : `/policies/list`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to get policies list:", error); + throw error; + } +}; + +export const createPolicyCall = async (accessToken: string, policyData: any) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies` : `/policies`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify(policyData), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to create policy:", error); + throw error; + } +}; + +export const updatePolicyCall = async (accessToken: string, policyId: string, policyData: any) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/${policyId}` : `/policies/${policyId}`; + const response = await fetch(url, { + method: "PUT", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify(policyData), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to update policy:", error); + throw error; + } +}; + +export const deletePolicyCall = async (accessToken: string, policyId: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/${policyId}` : `/policies/${policyId}`; + const response = await fetch(url, { + method: "DELETE", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to delete policy:", error); + throw error; + } +}; + +export const getPolicyInfo = async (accessToken: string, policyId: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/${policyId}` : `/policies/${policyId}`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to get policy info:", error); + throw error; + } +}; + +// Policy Attachments API Calls + +export const getPolicyAttachmentsList = async (accessToken: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/attachments/list` : `/policies/attachments/list`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to get policy attachments list:", error); + throw error; + } +}; + +export const createPolicyAttachmentCall = async (accessToken: string, attachmentData: any) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/attachments` : `/policies/attachments`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify(attachmentData), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to create policy attachment:", error); + throw error; + } +}; + +export const deletePolicyAttachmentCall = async (accessToken: string, attachmentId: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/attachments/${attachmentId}` : `/policies/attachments/${attachmentId}`; + const response = await fetch(url, { + method: "DELETE", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to delete policy attachment:", error); + throw error; + } +}; + +export const getResolvedGuardrails = async (accessToken: string, policyId: string) => { + try { + const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/${policyId}/resolved-guardrails` : `/policies/${policyId}/resolved-guardrails`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + const data = await response.json(); + return data; + } catch (error) { + console.error("Failed to get resolved guardrails:", error); + throw error; + } +}; + export const getPromptsList = async (accessToken: string): Promise => { try { const url = proxyBaseUrl ? `${proxyBaseUrl}/prompts/list` : `/prompts/list`; diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx index 3a7f8fec650..ff93194c3f1 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/ChatUI.tsx @@ -30,6 +30,7 @@ import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; import { v4 as uuidv4 } from "uuid"; import { truncateString } from "../../../utils/textUtils"; import GuardrailSelector from "../../guardrails/GuardrailSelector"; +import PolicySelector from "../../policies/PolicySelector"; import { MCPServer } from "../../mcp_tools/types"; import NotificationsManager from "../../molecules/notifications_manager"; import { fetchMCPServers, listMCPTools } from "../../networking"; @@ -188,6 +189,15 @@ const ChatUI: React.FC = ({ return []; } }); + const [selectedPolicies, setSelectedPolicies] = useState(() => { + const saved = sessionStorage.getItem("selectedPolicies"); + try { + return saved ? JSON.parse(saved) : []; + } catch (error) { + console.error("Error parsing selectedPolicies from sessionStorage", error); + return []; + } + }); const [messageTraceId, setMessageTraceId] = useState( () => sessionStorage.getItem("messageTraceId") || null, ); @@ -261,6 +271,7 @@ const ChatUI: React.FC = ({ selectedTags, selectedVectorStores, selectedGuardrails, + selectedPolicies, selectedMCPServers, mcpServers, mcpServerToolRestrictions, @@ -283,6 +294,7 @@ const ChatUI: React.FC = ({ selectedTags, selectedVectorStores, selectedGuardrails, + selectedPolicies, selectedMCPServers, mcpServers, mcpServerToolRestrictions, @@ -308,6 +320,7 @@ const ChatUI: React.FC = ({ sessionStorage.setItem("selectedTags", JSON.stringify(selectedTags)); sessionStorage.setItem("selectedVectorStores", JSON.stringify(selectedVectorStores)); sessionStorage.setItem("selectedGuardrails", JSON.stringify(selectedGuardrails)); + sessionStorage.setItem("selectedPolicies", JSON.stringify(selectedPolicies)); sessionStorage.setItem("selectedMCPServers", JSON.stringify(selectedMCPServers)); sessionStorage.setItem("mcpServerToolRestrictions", JSON.stringify(mcpServerToolRestrictions)); sessionStorage.setItem("selectedVoice", selectedVoice); @@ -338,6 +351,7 @@ const ChatUI: React.FC = ({ selectedTags, selectedVectorStores, selectedGuardrails, + selectedPolicies, messageTraceId, responsesSessionId, useApiSessionManagement, @@ -897,6 +911,7 @@ const ChatUI: React.FC = ({ traceId, selectedVectorStores.length > 0 ? selectedVectorStores : undefined, selectedGuardrails.length > 0 ? selectedGuardrails : undefined, + selectedPolicies.length > 0 ? selectedPolicies : undefined, selectedMCPServers, updateChatImageUI, updateSearchResults, @@ -977,6 +992,7 @@ const ChatUI: React.FC = ({ traceId, selectedVectorStores.length > 0 ? selectedVectorStores : undefined, selectedGuardrails.length > 0 ? selectedGuardrails : undefined, + selectedPolicies.length > 0 ? selectedPolicies : undefined, selectedMCPServers, // Pass the selected servers array useApiSessionManagement ? responsesSessionId : null, // Only pass session ID if API mode is enabled handleResponseId, // Pass callback to capture new response ID @@ -1008,6 +1024,7 @@ const ChatUI: React.FC = ({ traceId, selectedVectorStores.length > 0 ? selectedVectorStores : undefined, selectedGuardrails.length > 0 ? selectedGuardrails : undefined, + selectedPolicies.length > 0 ? selectedPolicies : undefined, selectedMCPServers, // Pass the selected tools array customProxyBaseUrl || undefined, ); @@ -1587,6 +1604,32 @@ const ChatUI: React.FC = ({ /> +
+ + Policies + + Select policy/policies to apply to this LLM API call. Policies define which guardrails are applied based on conditions. You can set up your policies{" "} + + here + + . + + } + > + + + + +
+ {/* Code Interpreter Toggle - Only for Responses endpoint */} {endpointType === EndpointType.RESPONSES && (
diff --git a/ui/litellm-dashboard/src/components/playground/chat_ui/CodeSnippets.tsx b/ui/litellm-dashboard/src/components/playground/chat_ui/CodeSnippets.tsx index 84f7f6b9e1c..6998d542401 100644 --- a/ui/litellm-dashboard/src/components/playground/chat_ui/CodeSnippets.tsx +++ b/ui/litellm-dashboard/src/components/playground/chat_ui/CodeSnippets.tsx @@ -6,6 +6,7 @@ interface CodeGenMetadata { tags?: string[]; vector_stores?: string[]; guardrails?: string[]; + policies?: string[]; } interface GenerateCodeParams { @@ -17,6 +18,7 @@ interface GenerateCodeParams { selectedTags: string[]; selectedVectorStores: string[]; selectedGuardrails: string[]; + selectedPolicies: string[]; selectedMCPServers: string[]; mcpServers?: MCPServer[]; mcpServerToolRestrictions?: Record; @@ -40,6 +42,7 @@ export const generateCodeSnippet = (params: GenerateCodeParams): string => { selectedTags, selectedVectorStores, selectedGuardrails, + selectedPolicies, selectedMCPServers, mcpServers, mcpServerToolRestrictions, @@ -72,6 +75,7 @@ export const generateCodeSnippet = (params: GenerateCodeParams): string => { if (selectedTags.length > 0) metadata.tags = selectedTags; if (selectedVectorStores.length > 0) metadata.vector_stores = selectedVectorStores; if (selectedGuardrails.length > 0) metadata.guardrails = selectedGuardrails; + if (selectedPolicies.length > 0) metadata.policies = selectedPolicies; const modelNameForCode = selectedModel || "your-model-name"; diff --git a/ui/litellm-dashboard/src/components/playground/llm_calls/anthropic_messages.tsx b/ui/litellm-dashboard/src/components/playground/llm_calls/anthropic_messages.tsx index 2e8b0be88bb..5570c7408fa 100644 --- a/ui/litellm-dashboard/src/components/playground/llm_calls/anthropic_messages.tsx +++ b/ui/litellm-dashboard/src/components/playground/llm_calls/anthropic_messages.tsx @@ -17,6 +17,7 @@ export async function makeAnthropicMessagesRequest( traceId?: string, vector_store_ids?: string[], guardrails?: string[], + policies?: string[], selectedMCPTools?: string[], customBaseUrl?: string, ) { @@ -59,6 +60,7 @@ export async function makeAnthropicMessagesRequest( if (vector_store_ids) requestBody.vector_store_ids = vector_store_ids; if (guardrails) requestBody.guardrails = guardrails; + if (policies) requestBody.policies = policies; // Use the streaming helper method for cleaner async iteration // @ts-ignore - The SDK types might not include all litellm-specific parameters const stream = client.messages.stream(requestBody, { signal }); diff --git a/ui/litellm-dashboard/src/components/playground/llm_calls/chat_completion.tsx b/ui/litellm-dashboard/src/components/playground/llm_calls/chat_completion.tsx index c3c623c25c3..61d232082e0 100644 --- a/ui/litellm-dashboard/src/components/playground/llm_calls/chat_completion.tsx +++ b/ui/litellm-dashboard/src/components/playground/llm_calls/chat_completion.tsx @@ -19,6 +19,7 @@ export async function makeOpenAIChatCompletionRequest( traceId?: string, vector_store_ids?: string[], guardrails?: string[], + policies?: string[], selectedMCPServers?: string[], onImageGenerated?: (imageUrl: string, model?: string) => void, onSearchResults?: (searchResults: VectorStoreSearchResponse[]) => void, @@ -110,6 +111,7 @@ export async function makeOpenAIChatCompletionRequest( messages: chatHistory as ChatCompletionMessageParam[], ...(vector_store_ids ? { vector_store_ids } : {}), ...(guardrails ? { guardrails } : {}), + ...(policies ? { policies } : {}), ...(tools.length > 0 ? { tools, tool_choice: "auto" } : {}), ...(temperature !== undefined ? { temperature } : {}), ...(max_tokens !== undefined ? { max_tokens } : {}), diff --git a/ui/litellm-dashboard/src/components/playground/llm_calls/responses_api.tsx b/ui/litellm-dashboard/src/components/playground/llm_calls/responses_api.tsx index b658610a21e..c69f82a37bf 100644 --- a/ui/litellm-dashboard/src/components/playground/llm_calls/responses_api.tsx +++ b/ui/litellm-dashboard/src/components/playground/llm_calls/responses_api.tsx @@ -27,6 +27,7 @@ export async function makeOpenAIResponsesRequest( traceId?: string, vector_store_ids?: string[], guardrails?: string[], + policies?: string[], selectedMCPServers?: string[], previousResponseId?: string | null, onResponseId?: (responseId: string) => void, @@ -137,6 +138,7 @@ export async function makeOpenAIResponsesRequest( ...(previousResponseId ? { previous_response_id: previousResponseId } : {}), ...(vector_store_ids ? { vector_store_ids } : {}), ...(guardrails ? { guardrails } : {}), + ...(policies ? { policies } : {}), ...(tools.length > 0 ? { tools, tool_choice: "auto" } : {}), }, { signal }, diff --git a/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx b/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx new file mode 100644 index 00000000000..fd20d9330a3 --- /dev/null +++ b/ui/litellm-dashboard/src/components/policies/PolicySelector.tsx @@ -0,0 +1,77 @@ +import React, { useEffect, useState } from "react"; +import { Select } from "antd"; +import { Policy } from "./types"; +import { getPoliciesList } from "../networking"; + +interface PolicySelectorProps { + onChange: (selectedPolicies: string[]) => void; + value?: string[]; + className?: string; + accessToken: string; + disabled?: boolean; +} + +const PolicySelector: React.FC = ({ + onChange, + value, + className, + accessToken, + disabled +}) => { + const [policies, setPolicies] = useState([]); + const [loading, setLoading] = useState(false); + + useEffect(() => { + const fetchPolicies = async () => { + if (!accessToken) return; + + setLoading(true); + try { + const response = await getPoliciesList(accessToken); + console.log("Policies response:", response); + if (response.policies) { + console.log("Policies data:", response.policies); + setPolicies(response.policies); + } + } catch (error) { + console.error("Error fetching policies:", error); + } finally { + setLoading(false); + } + }; + + fetchPolicies(); + }, [accessToken]); + + const handlePolicyChange = (selectedValues: string[]) => { + console.log("Selected policies:", selectedValues); + onChange(selectedValues); + }; + + return ( +
+ + (option?.label ?? "").toLowerCase().includes(input.toLowerCase()) + } + style={{ width: "100%" }} + /> + + + + Scope + + + + setScopeType(e.target.value)} + > + Global (applies to all requests) + Specific (teams, keys, or models) + + + + {scopeType === "specific" && ( + <> + + ({ + label: key, + value: key, + }))} + tokenSeparators={[","]} + showSearch + filterOption={(input, option) => + (option?.label ?? "").toLowerCase().includes(input.toLowerCase()) + } + style={{ width: "100%" }} + /> + + + +