diff --git a/.circleci/config.yml b/.circleci/config.yml index 7a982d74cbe..a8ccfcf7103 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -657,7 +657,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -2108,7 +2108,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -2250,7 +2250,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -2390,7 +2390,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -2551,7 +2551,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -2664,7 +2664,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -2800,7 +2800,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -3032,7 +3032,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 @@ -3549,7 +3549,7 @@ jobs: docker run -d \ --name postgres-db \ -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ + -e POSTGRES_PASSWORD=test-postgres \ -e POSTGRES_DB=circle_test \ -p 5432:5432 \ postgres:14 diff --git a/.gitguardian.yaml b/.gitguardian.yaml new file mode 100644 index 00000000000..861dd6e6d68 --- /dev/null +++ b/.gitguardian.yaml @@ -0,0 +1,84 @@ +version: 2 + +secret: + # Exclude files and paths by globbing + ignored_paths: + - "**/*.whl" + - "**/*.pyc" + - "**/__pycache__/**" + - "**/node_modules/**" + - "**/dist/**" + - "**/build/**" + - "**/.git/**" + - "**/venv/**" + - "**/.venv/**" + + # Large data/metadata files that don't need scanning + - "**/model_prices_and_context_window*.json" + - "**/*_metadata/*.txt" + - "**/tokenizers/*.json" + - "**/tokenizers/*" + - "miniconda.sh" + + # Build outputs and static assets + - "litellm/proxy/_experimental/out/**" + - "ui/litellm-dashboard/public/**" + - "**/swagger/*.js" + - "**/*.woff" + - "**/*.woff2" + - "**/*.avif" + - "**/*.webp" + + # Test data files + - "**/tests/**/data_map.txt" + - "tests/**/*.txt" + + # Documentation and other non-code files + - "docs/**" + - "**/*.md" + - "**/*.lock" + - "poetry.lock" + - "package-lock.json" + + # Ignore security incidents with the SHA256 of the occurrence (false positives) + ignored_matches: + # === Current detected false positives (SHA-based) === + + # gcs_pub_sub_body - folder name, not a password + - name: GCS pub/sub test folder name + match: 75f377c456eede69e5f6e47399ccee6016a2a93cc5dd11db09cc5b1359ae569a + + # os.environ/APORIA_API_KEY_1 - environment variable reference + - name: Environment variable reference APORIA_API_KEY_1 + match: e2ddeb8b88eca97a402559a2be2117764e11c074d86159ef9ad2375dea188094 + + # os.environ/APORIA_API_KEY_2 - environment variable reference + - name: Environment variable reference APORIA_API_KEY_2 + match: 09aa39a29e050b86603aa55138af1ff08fb86a4582aa965c1bd0672e1575e052 + + # oidc/circleci_v2/ - test authentication path, not a secret + - name: OIDC CircleCI test path + match: feb3475e1f89a65b7b7815ac4ec597e18a9ec1847742ad445c36ca617b536e15 + + # text-davinci-003 - OpenAI model identifier, not a secret + - name: OpenAI model identifier text-davinci-003 + match: c489000cf6c7600cee0eefb80ad0965f82921cfb47ece880930eb7e7635cf1f1 + + # === Preventive patterns for test keys (pattern-based) === + + # Test API keys (124 instances across 45 files) + - name: Test API keys with sk-test prefix + match: sk-test- + + # Mock API keys + - name: Mock API keys with sk-mock prefix + match: sk-mock- + + # Fake API keys + - name: Fake API keys with sk-fake prefix + match: sk-fake- + + # Generic test API key patterns + - name: Test API key patterns + match: test-api-key + diff --git a/.github/workflows/locustfile.py b/.github/workflows/locustfile.py index 36dbeee9c48..65d0d56b3a6 100644 --- a/.github/workflows/locustfile.py +++ b/.github/workflows/locustfile.py @@ -8,7 +8,7 @@ class MyUser(HttpUser): def chat_completion(self): headers = { "Content-Type": "application/json", - "Authorization": "Bearer sk-8N1tLOOyH8TIxwOLahhIVg", + "Authorization": "Bearer sk-test-load-test-key-123", # Include any additional headers you may need for authentication, etc. } diff --git a/.github/workflows/publish-migrations.yml b/.github/workflows/publish-migrations.yml index 8e5a67bcf85..a81a64ab46a 100644 --- a/.github/workflows/publish-migrations.yml +++ b/.github/workflows/publish-migrations.yml @@ -20,7 +20,7 @@ jobs: env: POSTGRES_DB: temp_db POSTGRES_USER: postgres - POSTGRES_PASSWORD: postgres + POSTGRES_PASSWORD: test-postgres ports: - 5432:5432 options: >- @@ -35,7 +35,7 @@ jobs: env: POSTGRES_DB: shadow_db POSTGRES_USER: postgres - POSTGRES_PASSWORD: postgres + POSTGRES_PASSWORD: test-postgres ports: - 5433:5432 options: >- diff --git a/README.md b/README.md index 2a67588706b..a020bd80898 100644 --- a/README.md +++ b/README.md @@ -266,6 +266,7 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature | [AI21 (`ai21`)](https://docs.litellm.ai/docs/providers/ai21) | ✅ | ✅ | ✅ | | | | | | | | | [AI21 Chat (`ai21_chat`)](https://docs.litellm.ai/docs/providers/ai21) | ✅ | ✅ | ✅ | | | | | | | | | [Aleph Alpha](https://docs.litellm.ai/docs/providers/aleph_alpha) | ✅ | ✅ | ✅ | | | | | | | | +| [Amazon Nova](https://docs.litellm.ai/docs/providers/amazon_nova) | ✅ | ✅ | ✅ | | | | | | | | | [Anthropic (`anthropic`)](https://docs.litellm.ai/docs/providers/anthropic) | ✅ | ✅ | ✅ | | | | | | ✅ | | | [Anthropic Text (`anthropic_text`)](https://docs.litellm.ai/docs/providers/anthropic) | ✅ | ✅ | ✅ | | | | | | ✅ | | | [Anyscale](https://docs.litellm.ai/docs/providers/anyscale) | ✅ | ✅ | ✅ | | | | | | | | diff --git a/ci_cd/TEST_KEY_PATTERNS.md b/ci_cd/TEST_KEY_PATTERNS.md new file mode 100644 index 00000000000..bd59f582839 --- /dev/null +++ b/ci_cd/TEST_KEY_PATTERNS.md @@ -0,0 +1,40 @@ +# Test Key Patterns Standard + +Standard patterns for test/mock keys and credentials in the LiteLLM codebase to avoid triggering secret detection. + +## How GitGuardian Works + +GitGuardian uses **machine learning and entropy analysis**, not just pattern matching: +- **Low entropy** values (like `sk-1234`, `postgres`) are automatically ignored +- **High entropy** values (realistic-looking secrets) trigger detection +- **Context-aware** detection understands code syntax like `os.environ["KEY"]` + +## Recommended Test Key Patterns + +### Option 1: Low Entropy Values (Simplest) +These won't trigger GitGuardian's ML detector: + +```python +api_key = "sk-1234" +api_key = "sk-12345" +database_password = "postgres" +token = "test123" +``` + +### Option 2: High Entropy with Test Prefixes +If you need realistic-looking test keys with high entropy, use these prefixes: + +```python +api_key = "sk-test-abc123def456ghi789..." # OpenAI-style test key +api_key = "sk-mock-1234567890abcdef1234..." # Mock key +api_key = "sk-fake-xyz789uvw456rst123..." # Fake key +token = "test-api-key-with-high-entropy" +``` + +## Configured Ignore Patterns + +These patterns are in `.gitguardian.yaml` for high-entropy test keys: +- `sk-test-*` - OpenAI-style test keys +- `sk-mock-*` - Mock API keys +- `sk-fake-*` - Fake API keys +- `test-api-key` - Generic test tokens diff --git a/cookbook/LiteLLM_PromptLayer.ipynb b/cookbook/LiteLLM_PromptLayer.ipynb index 3552636011a..8fd54941027 100644 --- a/cookbook/LiteLLM_PromptLayer.ipynb +++ b/cookbook/LiteLLM_PromptLayer.ipynb @@ -39,7 +39,7 @@ "import os\n", "os.environ['OPENAI_API_KEY'] = \"\"\n", "os.environ['REPLICATE_API_TOKEN'] = \"\"\n", - "os.environ['PROMPTLAYER_API_KEY'] = \"pl_4ea2bb00a4dca1b8a70cebf2e9e11564\"\n", + "os.environ['PROMPTLAYER_API_KEY'] = \"test-promptlayer-key-123\"\n", "\n", "# Set Promptlayer as a success callback\n", "litellm.success_callback =['promptlayer']\n", diff --git a/cookbook/Migrating_to_LiteLLM_Proxy_from_OpenAI_Azure_OpenAI.ipynb b/cookbook/Migrating_to_LiteLLM_Proxy_from_OpenAI_Azure_OpenAI.ipynb index 39677ed2a8a..740e7c7a4c8 100644 --- a/cookbook/Migrating_to_LiteLLM_Proxy_from_OpenAI_Azure_OpenAI.ipynb +++ b/cookbook/Migrating_to_LiteLLM_Proxy_from_OpenAI_Azure_OpenAI.ipynb @@ -1,21 +1,10 @@ { - "nbformat": 4, - "nbformat_minor": 0, - "metadata": { - "colab": { - "provenance": [] - }, - "kernelspec": { - "name": "python3", - "display_name": "Python 3" - }, - "language_info": { - "name": "python" - } - }, "cells": [ { "cell_type": "markdown", + "metadata": { + "id": "kccfk0mHZ4Ad" + }, "source": [ "# Migrating to LiteLLM Proxy from OpenAI/Azure OpenAI\n", "\n", @@ -32,29 +21,26 @@ "To pass provider-specific args, [go here](https://docs.litellm.ai/docs/completion/provider_specific_params#proxy-usage)\n", "\n", "To drop unsupported params (E.g. frequency_penalty for bedrock with librechat), [go here](https://docs.litellm.ai/docs/completion/drop_params#openai-proxy-usage)\n" - ], - "metadata": { - "id": "kccfk0mHZ4Ad" - } + ] }, { "cell_type": "markdown", + "metadata": { + "id": "nmSClzCPaGH6" + }, "source": [ "## /chat/completion\n", "\n" - ], - "metadata": { - "id": "nmSClzCPaGH6" - } + ] }, { "cell_type": "markdown", - "source": [ - "### OpenAI Python SDK" - ], "metadata": { "id": "_vqcjwOVaKpO" - } + }, + "source": [ + "### OpenAI Python SDK" + ] }, { "cell_type": "code", @@ -94,15 +80,20 @@ }, { "cell_type": "markdown", - "source": [ - "## Function Calling" - ], "metadata": { "id": "AqkyKk9Scxgj" - } + }, + "source": [ + "## Function Calling" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "wDg10VqLczE1" + }, + "outputs": [], "source": [ "from openai import OpenAI\n", "client = OpenAI(\n", @@ -139,24 +130,24 @@ ")\n", "\n", "print(completion)\n" - ], - "metadata": { - "id": "wDg10VqLczE1" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Azure OpenAI Python SDK" - ], "metadata": { "id": "YYoxLloSaNWW" - } + }, + "source": [ + "### Azure OpenAI Python SDK" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "yA1XcgowaSRy" + }, + "outputs": [], "source": [ "import openai\n", "client = openai.AzureOpenAI(\n", @@ -184,24 +175,24 @@ ")\n", "\n", "print(response)" - ], - "metadata": { - "id": "yA1XcgowaSRy" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Langchain Python" - ], "metadata": { "id": "yl9qhDvnaTpL" - } + }, + "source": [ + "### Langchain Python" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "5MUZgSquaW5t" + }, + "outputs": [], "source": [ "from langchain.chat_models import ChatOpenAI\n", "from langchain.prompts.chat import (\n", @@ -239,24 +230,22 @@ "response = chat(messages)\n", "\n", "print(response)" - ], - "metadata": { - "id": "5MUZgSquaW5t" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Curl" - ], "metadata": { "id": "B9eMgnULbRaz" - } + }, + "source": [ + "### Curl" + ] }, { "cell_type": "markdown", + "metadata": { + "id": "VWCCk5PFcmhS" + }, "source": [ "\n", "\n", @@ -280,22 +269,24 @@ "}'\n", "```\n", "\n" - ], - "metadata": { - "id": "VWCCk5PFcmhS" - } + ] }, { "cell_type": "markdown", - "source": [ - "### LlamaIndex" - ], "metadata": { "id": "drBAm2e1b6xe" - } + }, + "source": [ + "### LlamaIndex" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "d0bZcv8fb9mL" + }, + "outputs": [], "source": [ "import os, dotenv\n", "\n", @@ -326,24 +317,24 @@ "query_engine = index.as_query_engine()\n", "response = query_engine.query(\"What did the author do growing up?\")\n", "print(response)\n" - ], - "metadata": { - "id": "d0bZcv8fb9mL" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Langchain JS" - ], "metadata": { "id": "xypvNdHnb-Yy" - } + }, + "source": [ + "### Langchain JS" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "R55mK2vCcBN2" + }, + "outputs": [], "source": [ "import { ChatOpenAI } from \"@langchain/openai\";\n", "\n", @@ -359,24 +350,24 @@ "const message = await model.invoke(\"Hi there!\");\n", "\n", "console.log(message);\n" - ], - "metadata": { - "id": "R55mK2vCcBN2" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### OpenAI JS" - ], "metadata": { "id": "nC4bLifCcCiW" - } + }, + "source": [ + "### OpenAI JS" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "MICH8kIMcFpg" + }, + "outputs": [], "source": [ "const { OpenAI } = require('openai');\n", "\n", @@ -398,24 +389,24 @@ "}\n", "\n", "main();\n" - ], - "metadata": { - "id": "MICH8kIMcFpg" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Anthropic SDK" - ], "metadata": { "id": "D1Q07pEAcGTb" - } + }, + "source": [ + "### Anthropic SDK" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "qBjFcAvgcI3t" + }, + "outputs": [], "source": [ "import os\n", "\n", @@ -423,7 +414,7 @@ "\n", "client = Anthropic(\n", " base_url=\"http://localhost:4000\", # proxy endpoint\n", - " api_key=\"sk-s4xN1IiLTCytwtZFJaYQrA\", # litellm proxy virtual key\n", + " api_key=\"sk-test-proxy-key-123\", # litellm proxy virtual key (example)\n", ")\n", "\n", "message = client.messages.create(\n", @@ -437,33 +428,33 @@ " model=\"claude-3-opus-20240229\",\n", ")\n", "print(message.content)" - ], - "metadata": { - "id": "qBjFcAvgcI3t" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "## /embeddings" - ], "metadata": { "id": "dFAR4AJGcONI" - } + }, + "source": [ + "## /embeddings" + ] }, { "cell_type": "markdown", - "source": [ - "### OpenAI Python SDK" - ], "metadata": { "id": "lgNoM281cRzR" - } + }, + "source": [ + "### OpenAI Python SDK" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "NY3DJhPfcQhA" + }, + "outputs": [], "source": [ "import openai\n", "from openai import OpenAI\n", @@ -478,24 +469,24 @@ ")\n", "\n", "print(response)\n" - ], - "metadata": { - "id": "NY3DJhPfcQhA" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Langchain Embeddings" - ], "metadata": { "id": "hmbg-DW6cUZs" - } + }, + "source": [ + "### Langchain Embeddings" + ] }, { "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "lX2S8Nl1cWVP" + }, + "outputs": [], "source": [ "from langchain.embeddings import OpenAIEmbeddings\n", "\n", @@ -526,24 +517,22 @@ "\n", "print(f\"TITAN EMBEDDINGS\")\n", "print(query_result[:5])" - ], - "metadata": { - "id": "lX2S8Nl1cWVP" - }, - "execution_count": null, - "outputs": [] + ] }, { "cell_type": "markdown", - "source": [ - "### Curl Request" - ], "metadata": { "id": "oqGbWBCQcYfd" - } + }, + "source": [ + "### Curl Request" + ] }, { "cell_type": "markdown", + "metadata": { + "id": "7rkIMV9LcdwQ" + }, "source": [ "\n", "\n", @@ -556,10 +545,21 @@ " }'\n", "```\n", "\n" - ], - "metadata": { - "id": "7rkIMV9LcdwQ" - } + ] } - ] -} \ No newline at end of file + ], + "metadata": { + "colab": { + "provenance": [] + }, + "kernelspec": { + "display_name": "Python 3", + "name": "python3" + }, + "language_info": { + "name": "python" + } + }, + "nbformat": 4, + "nbformat_minor": 0 +} diff --git a/docs/my-website/blog/gemini_3_flash/index.md b/docs/my-website/blog/gemini_3_flash/index.md index c6135a09513..6cb8ddad992 100644 --- a/docs/my-website/blog/gemini_3_flash/index.md +++ b/docs/my-website/blog/gemini_3_flash/index.md @@ -27,6 +27,10 @@ import TabItem from '@theme/TabItem'; LiteLLM now supports `gemini-3-flash-preview` and all the new API changes along with it. +:::note +If you only want cost tracking, you need no change in your current Litellm version. But if you want the support for new features introduced along with it like thinking levels, you will need to use v1.80.8-stable.1 or above. +::: + ## Deploy this version @@ -232,6 +236,11 @@ response = completion( print(response) ``` +:::note +If using this model via vertex_ai, keep the location as global as this is the only supported location as of now. +::: + + ## `reasoning_effort` Mapping for Gemini 3+ | reasoning_effort | thinking_level | diff --git a/docs/my-website/docs/image_edits.md b/docs/my-website/docs/image_edits.md index 5a108aabf3a..a8438334542 100644 --- a/docs/my-website/docs/image_edits.md +++ b/docs/my-website/docs/image_edits.md @@ -16,7 +16,7 @@ LiteLLM provides image editing functionality that maps to OpenAI's `/images/edit | Supported operations | Create image edits | Single and multiple images supported | | Supported LiteLLM SDK Versions | 1.63.8+ | Gemini support requires 1.79.3+ | | Supported LiteLLM Proxy Versions | 1.71.1+ | Gemini support requires 1.79.3+ | -| Supported LLM providers | **OpenAI**, **Gemini (Google AI Studio)**, **Vertex AI** | Gemini supports the new `gemini-2.5-flash-image` family. Vertex AI supports both Gemini and Imagen models. | +| Supported LLM providers | **OpenAI**, **Gemini (Google AI Studio)**, **Vertex AI**, **Stability AI**, **AWS Bedrock (Stability)** | Gemini supports the new `gemini-2.5-flash-image` family. Vertex AI supports both Gemini and Imagen models. Stability AI and Bedrock Stability support various image editing operations. | #### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/) diff --git a/docs/my-website/docs/oidc.md b/docs/my-website/docs/oidc.md index 3db4b6ecdc5..b541329aa38 100644 --- a/docs/my-website/docs/oidc.md +++ b/docs/my-website/docs/oidc.md @@ -106,7 +106,7 @@ model_list: aws_region_name: us-west-2 aws_session_name: "my-test-session" aws_role_name: "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci" - aws_web_identity_token: "oidc/circleci_v2/" + aws_web_identity_token: "oidc/example-provider/" ``` #### Amazon IAM Role Configuration for CircleCI v2 -> Bedrock diff --git a/docs/my-website/docs/providers/openai/responses_api.md b/docs/my-website/docs/providers/openai/responses_api.md index 8d91ca674b7..75eab1afac5 100644 --- a/docs/my-website/docs/providers/openai/responses_api.md +++ b/docs/my-website/docs/providers/openai/responses_api.md @@ -623,6 +623,58 @@ display(styled_df) +## Function Calling + +```python showLineNumbers title="Function Calling with Parallel Tool Calls" +import litellm +import json + +tools = [ + { + "type": "function", + "name": "get_weather", + "description": "Get current weather for a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"} + }, + "required": ["location"] + } + } +] + +# Step 1: Request with tools (parallel_tool_calls=True allows multiple calls) +response = litellm.responses( + model="openai/gpt-4o", + input=[{"role": "user", "content": "What's the weather in Paris and Tokyo?"}], + tools=tools, + parallel_tool_calls=True, # Defaults = True +) + +# Step 2: Execute tool calls and collect results +tool_results = [] +for output in response.output: + if output.type == "function_call": + result = {"temperature": 15, "condition": "sunny"} # Your function logic here + tool_results.append({ + "type": "function_call_output", + "call_id": output.call_id, + "output": json.dumps(result) + }) + +# Step 3: Send results back +final_response = litellm.responses( + model="openai/gpt-4o", + input=tool_results, + tools=tools, +) + +print(final_response.output) +``` + +Set `parallel_tool_calls=False` to ensure zero or one tool is called per turn. [More details](https://platform.openai.com/docs/guides/function-calling#parallel-function-calling). + ## Free-form Function Calling @@ -633,7 +685,6 @@ display(styled_df) import litellm response = litellm.responses( - response = client.responses.create( model="gpt-5-mini", input="Please use the code_exec tool to calculate the area of a circle with radius equal to the number of 'r's in strawberry", text={"format": {"type": "text"}}, diff --git a/docs/my-website/docs/providers/stability.md b/docs/my-website/docs/providers/stability.md index 49773fffdb3..6b340267e69 100644 --- a/docs/my-website/docs/providers/stability.md +++ b/docs/my-website/docs/providers/stability.md @@ -8,7 +8,7 @@ https://stability.ai/ | Description | Stability AI creates open AI models for image, video, audio, and 3D generation. Known for Stable Diffusion. | | Provider Route on LiteLLM | `stability/` | | Link to Provider Doc | [Stability AI API ↗](https://platform.stability.ai/docs/api-reference) | -| Supported Operations | [`/images/generations`](#image-generation) | +| Supported Operations | [`/images/generations`](#image-generation), [`/images/edits`](#image-editing) | LiteLLM supports Stability AI Image Generation calls via the Stability AI REST API (not via Bedrock). @@ -169,13 +169,285 @@ Stability AI returns images in base64 format. The response is OpenAI-compatible: } ``` -## Comparing with Bedrock +## Image Editing + +Stability AI supports various image editing operations including inpainting, upscaling, outpainting, background removal, and more. + +### Usage - LiteLLM Python SDK + +#### Inpainting (Edit with Mask) + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Inpainting - edit specific areas using a mask +response = image_edit( + model="stability/stable-image-inpaint-v1:0", + image=open("original_image.png", "rb"), + mask=open("mask_image.png", "rb"), + prompt="Add a beautiful sunset in the masked area", + size="1024x1024", +) +print(response) +``` + +#### Image Upscaling + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Conservative upscaling - preserves details +response = image_edit( + model="stability/stable-conservative-upscale-v1:0", + image=open("low_res_image.png", "rb"), + prompt="Upscale this image while preserving details", +) + +# Creative upscaling - adds creative details +response = image_edit( + model="stability/stable-creative-upscale-v1:0", + image=open("low_res_image.png", "rb"), + prompt="Upscale and enhance with creative details", + creativity=0.3, # 0-0.35, higher = more creative +) + +# Fast upscaling - quick upscaling +response = image_edit( + model="stability/stable-fast-upscale-v1:0", + image=open("low_res_image.png", "rb"), + prompt="Quickly upscale this image", +) +print(response) +``` + +#### Image Outpainting + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Extend image beyond its borders +response = image_edit( + model="stability/stable-outpaint-v1:0", + image=open("original_image.png", "rb"), + prompt="Extend this landscape with mountains", + left=100, # Pixels to extend on the left + right=100, # Pixels to extend on the right + up=50, # Pixels to extend on top + down=50, # Pixels to extend on bottom +) +print(response) +``` + +#### Background Removal + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Remove background from image +response = image_edit( + model="stability/stable-image-remove-background-v1:0", + image=open("portrait.png", "rb"), + prompt="Remove the background", +) +print(response) +``` + +#### Search and Replace + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Search and replace objects in image +response = image_edit( + model="stability/stable-image-search-replace-v1:0", + image=open("scene.png", "rb"), + prompt="A red sports car", + search_prompt="blue sedan", # What to replace +) + +# Search and recolor +response = image_edit( + model="stability/stable-image-search-recolor-v1:0", + image=open("scene.png", "rb"), + prompt="Make it golden yellow", + select_prompt="the car", # What to recolor +) +print(response) +``` + +#### Image Control (Sketch/Structure) + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Control with sketch +response = image_edit( + model="stability/stable-image-control-sketch-v1:0", + image=open("sketch.png", "rb"), + prompt="Turn this sketch into a realistic photo", + control_strength=0.7, # 0-1, higher = more control +) + +# Control with structure +response = image_edit( + model="stability/stable-image-control-structure-v1:0", + image=open("structure_reference.png", "rb"), + prompt="Generate image following this structure", + control_strength=0.7, +) +print(response) +``` + +#### Erase Objects + +```python showLineNumbers +from litellm import image_edit +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Erase objects from image +response = image_edit( + model="stability/stable-image-erase-object-v1:0", + image=open("scene.png", "rb"), + mask=open("object_mask.png", "rb"), # Mask the object to erase + prompt="Remove the object", +) +print(response) +``` + +### Supported Image Edit Models + +| Model Name | Function Call | Description | +|------------|---------------|-------------| +| stable-image-inpaint-v1:0 | `image_edit(model="stability/stable-image-inpaint-v1:0", ...)` | Inpainting with mask | +| stable-conservative-upscale-v1:0 | `image_edit(model="stability/stable-conservative-upscale-v1:0", ...)` | Conservative upscaling | +| stable-creative-upscale-v1:0 | `image_edit(model="stability/stable-creative-upscale-v1:0", ...)` | Creative upscaling | +| stable-fast-upscale-v1:0 | `image_edit(model="stability/stable-fast-upscale-v1:0", ...)` | Fast upscaling | +| stable-outpaint-v1:0 | `image_edit(model="stability/stable-outpaint-v1:0", ...)` | Extend image borders | +| stable-image-remove-background-v1:0 | `image_edit(model="stability/stable-image-remove-background-v1:0", ...)` | Remove background | +| stable-image-search-replace-v1:0 | `image_edit(model="stability/stable-image-search-replace-v1:0", ...)` | Search and replace objects | +| stable-image-search-recolor-v1:0 | `image_edit(model="stability/stable-image-search-recolor-v1:0", ...)` | Search and recolor | +| stable-image-control-sketch-v1:0 | `image_edit(model="stability/stable-image-control-sketch-v1:0", ...)` | Control with sketch | +| stable-image-control-structure-v1:0 | `image_edit(model="stability/stable-image-control-structure-v1:0", ...)` | Control with structure | +| stable-image-erase-object-v1:0 | `image_edit(model="stability/stable-image-erase-object-v1:0", ...)` | Erase objects | +| stable-image-style-guide-v1:0 | `image_edit(model="stability/stable-image-style-guide-v1:0", ...)` | Apply style guide | +| stable-style-transfer-v1:0 | `image_edit(model="stability/stable-style-transfer-v1:0", ...)` | Transfer style | + +### Usage - LiteLLM Proxy Server + +#### 1. Setup config.yaml + +```yaml showLineNumbers +model_list: + - model_name: stability-inpaint + litellm_params: + model: stability/stable-image-inpaint-v1:0 + api_key: os.environ/STABILITY_API_KEY + model_info: + mode: image_edit + + - model_name: stability-upscale + litellm_params: + model: stability/stable-conservative-upscale-v1:0 + api_key: os.environ/STABILITY_API_KEY + model_info: + mode: image_edit + +general_settings: + master_key: sk-1234 +``` + +#### 2. Start the proxy + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +#### 3. Test it + +```bash showLineNumbers +curl -X POST "http://0.0.0.0:4000/v1/images/edits" \ + -H "Authorization: Bearer sk-1234" \ + -F "model=stability-inpaint" \ + -F "image=@original_image.png" \ + -F "mask=@mask_image.png" \ + -F "prompt=Add a beautiful garden in the masked area" +``` + +## AWS Bedrock (Stability) + +LiteLLM also supports Stability AI models via AWS Bedrock. This is useful if you're already using AWS infrastructure. + +### Usage - Bedrock Stability + +```python showLineNumbers +from litellm import image_edit +import os + +# Set AWS credentials +os.environ["AWS_ACCESS_KEY_ID"] = "your-access-key" +os.environ["AWS_SECRET_ACCESS_KEY"] = "your-secret-key" +os.environ["AWS_REGION_NAME"] = "us-east-1" + +# Bedrock Stability inpainting +response = image_edit( + model="bedrock/us.stability.stable-image-inpaint-v1:0", + image=open("original_image.png", "rb"), + mask=open("mask_image.png", "rb"), + prompt="Add flowers in the masked area", + size="1024x1024", +) +print(response) +``` + +### Supported Bedrock Stability Models + +All Stability AI image edit models are available via Bedrock with the `bedrock/` prefix: + +| Direct API Model | Bedrock Model | Description | +|------------------|---------------|-------------| +| stability/stable-image-inpaint-v1:0 | bedrock/us.stability.stable-image-inpaint-v1:0 | Inpainting | +| stability/stable-conservative-upscale-v1:0 | bedrock/stability.stable-conservative-upscale-v1:0 | Conservative upscaling | +| stability/stable-creative-upscale-v1:0 | bedrock/stability.stable-creative-upscale-v1:0 | Creative upscaling | +| stability/stable-fast-upscale-v1:0 | bedrock/stability.stable-fast-upscale-v1:0 | Fast upscaling | +| stability/stable-outpaint-v1:0 | bedrock/stability.stable-outpaint-v1:0 | Outpainting | +| stability/stable-image-remove-background-v1:0 | bedrock/stability.stable-image-remove-background-v1:0 | Remove background | +| stability/stable-image-search-replace-v1:0 | bedrock/stability.stable-image-search-replace-v1:0 | Search and replace | +| stability/stable-image-search-recolor-v1:0 | bedrock/stability.stable-image-search-recolor-v1:0 | Search and recolor | +| stability/stable-image-control-sketch-v1:0 | bedrock/stability.stable-image-control-sketch-v1:0 | Control with sketch | +| stability/stable-image-control-structure-v1:0 | bedrock/stability.stable-image-control-structure-v1:0 | Control with structure | +| stability/stable-image-erase-object-v1:0 | bedrock/stability.stable-image-erase-object-v1:0 | Erase objects | + +**Note:** Bedrock model IDs may use `us.stability.*` or `stability.*` prefix depending on the region and model. + +## Comparing Routes LiteLLM supports Stability AI models via two routes: -| Route | Provider | Use Case | -|-------|----------|----------| -| `stability/` | Stability AI Direct API | Direct access, all latest models | -| `bedrock/stability.*` | AWS Bedrock | AWS integration, enterprise features | +| Route | Provider | Use Case | Image Generation | Image Editing | +|-------|----------|----------|------------------|---------------| +| `stability/` | Stability AI Direct API | Direct access, all latest models | ✅ | ✅ | +| `bedrock/stability.*` | AWS Bedrock | AWS integration, enterprise features | ✅ | ✅ | Use `stability/` for direct API access. Use `bedrock/stability.*` if you're already using AWS Bedrock. diff --git a/docs/my-website/docs/providers/vertex_ocr.md b/docs/my-website/docs/providers/vertex_ocr.md index 4e3d4b0a063..9ff22a03775 100644 --- a/docs/my-website/docs/providers/vertex_ocr.md +++ b/docs/my-website/docs/providers/vertex_ocr.md @@ -140,7 +140,7 @@ with open("document.pdf", "rb") as f: pdf_base64 = base64.b64encode(f.read()).decode() response = litellm.ocr( - model="vertex_ai/mistral-ocr-2505", + model="vertex_ai/mistral-ocr-2505", # This doesn't work for deepseek document={ "type": "document_url", "document_url": f"data:application/pdf;base64,{pdf_base64}" @@ -219,7 +219,7 @@ print(f"Cost: ${response._hidden_params.get('response_cost', 0)}") ## Important Notes :::info URL Conversion -Vertex AI OCR endpoints don't have internet access. LiteLLM automatically converts public URLs to base64 data URIs before sending requests to Vertex AI. +Vertex AI Mistral OCR endpoints don't have internet access. LiteLLM automatically converts public URLs to base64 data URIs before sending requests to Vertex AI. ::: :::tip Regional Availability @@ -227,11 +227,14 @@ Mistral OCR is available in multiple regions. Specify `vertex_location` to use a - `us-central1` (default) - `europe-west1` - `asia-southeast1` + +Deepseek OCR is only available in global region. ::: ## Supported Models - `mistral-ocr-2505` - Latest Mistral OCR model on Vertex AI +- `deepseek-ocr-maas` - Lates Deepseek OCR model on Vertex AI Use the Vertex AI provider prefix: `vertex_ai/` diff --git a/docs/my-website/docs/proxy/alerting.md b/docs/my-website/docs/proxy/alerting.md index 4cbcd0cffce..38d6d47be44 100644 --- a/docs/my-website/docs/proxy/alerting.md +++ b/docs/my-website/docs/proxy/alerting.md @@ -215,16 +215,16 @@ general_settings: alerting: ["slack"] alerting_threshold: 0.0001 # (Seconds) set an artificially low threshold for testing alerting alert_to_webhook_url: { - "llm_exceptions": "https://hooks.slack.com/services/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", - "llm_too_slow": "https://hooks.slack.com/services/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", - "llm_requests_hanging": "https://hooks.slack.com/services/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", - "budget_alerts": "https://hooks.slack.com/services/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", - "db_exceptions": "https://hooks.slack.com/services/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", - "daily_reports": "https://hooks.slack.com/services/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", - "spend_reports": "https://hooks.slack.com/services/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", - "cooldown_deployment": "https://hooks.slack.com/services/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", - "new_model_added": "https://hooks.slack.com/services/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", - "outage_alerts": "https://hooks.slack.com/services/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", + "llm_exceptions": "example-slack-webhook-url", + "llm_too_slow": "example-slack-webhook-url", + "llm_requests_hanging": "example-slack-webhook-url", + "budget_alerts": "example-slack-webhook-url", + "db_exceptions": "example-slack-webhook-url", + "daily_reports": "example-slack-webhook-url", + "spend_reports": "example-slack-webhook-url", + "cooldown_deployment": "example-slack-webhook-url", + "new_model_added": "example-slack-webhook-url", + "outage_alerts": "example-slack-webhook-url", } litellm_settings: @@ -399,7 +399,7 @@ curl -X GET --location 'http://0.0.0.0:4000/health/services?service=webhook' \ { "spend": 1, # the spend for the 'event_group' "max_budget": 0, # the 'max_budget' set for the 'event_group' - "token": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "token": "example-api-key-123", "user_id": "default_user_id", "team_id": null, "user_email": null, diff --git a/docs/my-website/docs/proxy/cost_tracking.md b/docs/my-website/docs/proxy/cost_tracking.md index 019cd62c620..26a4920c093 100644 --- a/docs/my-website/docs/proxy/cost_tracking.md +++ b/docs/my-website/docs/proxy/cost_tracking.md @@ -722,7 +722,7 @@ curl -X GET 'http://localhost:4000/global/spend/report?start_date=2024-04-01&end ```shell [ { - "api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "api_key": "example-api-key-123", "total_cost": 0.3201286305151999, "total_input_tokens": 36.0, "total_output_tokens": 1593.0, @@ -766,7 +766,7 @@ curl -X GET 'http://localhost:4000/global/spend/report?start_date=2024-04-01&end ```shell [ { - "api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "api_key": "example-api-key-123", "total_cost": 0.00013132, "total_input_tokens": 105.0, "total_output_tokens": 872.0, @@ -1151,7 +1151,7 @@ curl -X GET "http://0.0.0.0:4000/spend/logs?request_id= -### Team-specific overrides (proxy) +### Team-specific overrides -When running the LiteLLM proxy you can override the Vault location per team. Set -`Secret Manager Settings` on the team with the following structure: +When running the LiteLLM proxy you can override the Vault location per team. Use the [Team-Level Secret Manager Settings](./overview.md#team-level-secret-manager-settings) flow in the dashboard and configure the panel shown below: + + + +Use the following structure for the JSON payload: ```json { diff --git a/docs/my-website/docs/secret_managers/overview.md b/docs/my-website/docs/secret_managers/overview.md index 957e7dc0a0b..a987c72d767 100644 --- a/docs/my-website/docs/secret_managers/overview.md +++ b/docs/my-website/docs/secret_managers/overview.md @@ -49,9 +49,28 @@ general_settings: ## Team-Level Secret Manager Settings -From the **Teams** page in the LiteLLM dashboard you can configure a secret manager per team. Open the team (or the “Create New Team” modal), find the **Secret Manager Settings** panel, and enter the provider-specific JSON configuration (e.g. `{"namespace": "admin", "mount": "secret", "path_prefix": "litellm"}`). This configuration is applied whenever LiteLLM writes secrets (e.g., storing virtual keys) on behalf of that team. +Team-level secret manager settings let every team bring their own key-management configuration. These settings are used when creating virtual keys tied to the team. - +Follow these steps to configure it: +1. **Create a team** + Open the Teams page and click `Create Team` to launch the modal. -Refer to each provider’s documentation (AWS, Azure, Google, Hashicorp, etc.) for the supported keys/values you can place inside `secret_manager_settings`. + + +2. **Expand Additional Settings** + Use the `Additional Settings` toggle to reveal the advanced configuration panel. + + + +3. **Configure the Secret Manager** + In the `Secret Manager Settings` panel, paste the provider-specific JSON. Refer to each provider page (AWS, Azure, Google, Hashicorp, etc.) for the supported keys/values. JSON is required today, but we plan to add a more UI-friendly editor. + + + +4. **Create the team** + Review the inputs and click `Create Team` to save. + + + +Once saved, LiteLLM will use this configuration. diff --git a/docs/my-website/docusaurus.config.js b/docs/my-website/docusaurus.config.js index 32d5d800b71..f6e61895e6a 100644 --- a/docs/my-website/docusaurus.config.js +++ b/docs/my-website/docusaurus.config.js @@ -8,7 +8,7 @@ const darkCodeTheme = require('prism-react-renderer/themes/dracula'); const inkeepConfig = { baseSettings: { - apiKey: "0cb9c9916ec71bfe0e53c9d7f83ff046daee3fa9ef318f6a", + apiKey: "test-inkeep-api-key-123", organizationDisplayName: 'liteLLM', primaryBrandColor: '#4965f5', theme: { diff --git a/docs/my-website/img/secret_manager_hashicorp_vault_settings.png b/docs/my-website/img/secret_manager_hashicorp_vault_settings.png new file mode 100644 index 00000000000..c471480a3b6 Binary files /dev/null and b/docs/my-website/img/secret_manager_hashicorp_vault_settings.png differ diff --git a/docs/my-website/img/secret_manager_settings.png b/docs/my-website/img/secret_manager_settings.png index ce13a60ee9e..4b01dd43206 100644 Binary files a/docs/my-website/img/secret_manager_settings.png and b/docs/my-website/img/secret_manager_settings.png differ diff --git a/docs/my-website/img/secret_manager_settings_additional_settings.png b/docs/my-website/img/secret_manager_settings_additional_settings.png new file mode 100644 index 00000000000..713031cb5c5 Binary files /dev/null and b/docs/my-website/img/secret_manager_settings_additional_settings.png differ diff --git a/docs/my-website/img/secret_manager_settings_create_button.png b/docs/my-website/img/secret_manager_settings_create_button.png new file mode 100644 index 00000000000..5c08eae8938 Binary files /dev/null and b/docs/my-website/img/secret_manager_settings_create_button.png differ diff --git a/docs/my-website/img/secret_manager_settings_create_team.png b/docs/my-website/img/secret_manager_settings_create_team.png new file mode 100644 index 00000000000..b6bd18e4287 Binary files /dev/null and b/docs/my-website/img/secret_manager_settings_create_team.png differ diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index fc019e033da..b4bf1293f98 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -671,6 +671,7 @@ const sidebars = { "providers/ai21", "providers/aiml", "providers/aleph_alpha", + "providers/amazon_nova", "providers/anyscale", "providers/baseten", "providers/bytez", diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py new file mode 100644 index 00000000000..4ee6a89cc98 --- /dev/null +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -0,0 +1,110 @@ +""" +Polls LiteLLM_ManagedObjectTable to check if the response is complete. +Cost tracking is handled automatically by litellm.aget_responses(). +""" + +from typing import TYPE_CHECKING + +import litellm +from litellm._logging import verbose_proxy_logger + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.router import Router + + +class CheckResponsesCost: + def __init__( + self, + proxy_logging_obj: "ProxyLogging", + prisma_client: "PrismaClient", + llm_router: "Router", + ): + from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.router import Router + + self.proxy_logging_obj: ProxyLogging = proxy_logging_obj + self.prisma_client: PrismaClient = prisma_client + self.llm_router: Router = llm_router + + async def check_responses_cost(self): + """ + Check if background responses are complete and track their cost. + - Get all status="queued" or "in_progress" and file_purpose="response" jobs + - Query the provider to check if response is complete + - Cost is automatically tracked by litellm.aget_responses() + - Mark completed/failed/cancelled responses as complete in the database + """ + jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( + where={ + "status": {"in": ["queued", "in_progress"]}, + "file_purpose": "response", + } + ) + + verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check") + completed_jobs = [] + + for job in jobs: + unified_object_id = job.unified_object_id + + try: + from litellm.proxy.hooks.responses_id_security import ( + ResponsesIDSecurity, + ) + + # Get the stored response object to extract model information + stored_response = job.file_object + model_name = stored_response.get("model", None) + + # Decrypt the response ID + responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id) + + # Prepare metadata with model information for cost tracking + litellm_metadata = { + "user_api_key_user_id": job.created_by or "default-user-id", + } + + # Add model information if available + if model_name: + litellm_metadata["model"] = model_name + litellm_metadata["model_group"] = model_name # Use same value for model_group + + response = await litellm.aget_responses( + response_id=responses_id_security, + litellm_metadata=litellm_metadata, + ) + + verbose_proxy_logger.debug( + f"Response {unified_object_id} status: {response.status}, model: {model_name}" + ) + + except Exception as e: + verbose_proxy_logger.info( + f"Skipping job {unified_object_id} due to error: {e}" + ) + continue + + # Check if response is in a terminal state + if response.status == "completed": + verbose_proxy_logger.info( + f"Response {unified_object_id} is complete. Cost automatically tracked by aget_responses." + ) + completed_jobs.append(job) + + elif response.status in ["failed", "cancelled"]: + verbose_proxy_logger.info( + f"Response {unified_object_id} has status {response.status}, marking as complete" + ) + completed_jobs.append(job) + + # Mark completed jobs in the database + if len(completed_jobs) > 0: + await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": {"in": [job.id for job in completed_jobs]}}, + data={"status": "completed"}, + ) + verbose_proxy_logger.info( + f"Marked {len(completed_jobs)} response jobs as completed" + ) + diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index e12be6baf5d..a83d7e224b5 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -23,7 +23,9 @@ from litellm.proxy._types import ( from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, get_batch_id_from_unified_batch_id, + get_content_type_from_file_object, get_model_id_from_unified_batch_id, + normalize_mime_type_for_provider, ) from litellm.types.llms.openai import ( AllMessageValues, @@ -33,6 +35,7 @@ from litellm.types.llms.openai import ( FileObject, OpenAIFileObject, OpenAIFilesPurpose, + ResponsesAPIResponse, ) from litellm.types.utils import ( CallTypesLiteral, @@ -41,10 +44,6 @@ from litellm.types.utils import ( LLMResponseTypes, SpecialEnums, ) -from litellm.proxy.openai_files_endpoints.common_utils import ( - get_content_type_from_file_object, - normalize_mime_type_for_provider, -) if TYPE_CHECKING: from litellm.types.llms.openai import HttpxBinaryResponseContent @@ -133,10 +132,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): async def store_unified_object_id( self, unified_object_id: str, - file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob], + file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, "ResponsesAPIResponse"], litellm_parent_otel_span: Optional[Span], model_object_id: str, - file_purpose: Literal["batch", "fine-tune"], + file_purpose: Literal["batch", "fine-tune", "response"], user_api_key_dict: UserAPIKeyAuth, ) -> None: verbose_logger.info( @@ -946,7 +945,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # File is stored in a storage backend, download and convert to base64 try: - from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend + from litellm.llms.base_llm.files.storage_backend_factory import ( + get_storage_backend, + ) storage_backend_name = db_file.storage_backend storage_url = db_file.storage_url diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index d5fb82808c1..2dd0c5e7556 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -824,4 +824,22 @@ model LiteLLM_UISettings { ui_settings Json created_at DateTime @default(now()) updated_at DateTime @updatedAt +} + +// Skills table for storing LiteLLM-managed skills +model LiteLLM_SkillsTable { + skill_id String @id @default(uuid()) + display_title String? + description String? + instructions String? // The skill instructions/prompt (from SKILL.md) + source String @default("custom") // "custom" or "anthropic" + latest_version String? + file_content Bytes? // Binary content of the skill files (zip) + file_name String? // Original filename + file_type String? // MIME type (e.g., "application/zip") + metadata Json? @default("{}") + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? } \ No newline at end of file diff --git a/litellm/__init__.py b/litellm/__init__.py index a1a12fc9313..9b69beccd79 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1203,9 +1203,9 @@ from .llms.bedrock.chat.invoke_transformations.amazon_openai_transformation impo AmazonBedrockOpenAIConfig, ) -from .llms.bedrock.image.amazon_stability1_transformation import AmazonStabilityConfig -from .llms.bedrock.image.amazon_stability3_transformation import AmazonStability3Config -from .llms.bedrock.image.amazon_nova_canvas_transformation import AmazonNovaCanvasConfig +from .llms.bedrock.image_generation.amazon_stability1_transformation import AmazonStabilityConfig +from .llms.bedrock.image_generation.amazon_stability3_transformation import AmazonStability3Config +from .llms.bedrock.image_generation.amazon_nova_canvas_transformation import AmazonNovaCanvasConfig from .llms.bedrock.embed.amazon_titan_g1_transformation import AmazonTitanG1Config from .llms.bedrock.embed.amazon_titan_multimodal_transformation import ( AmazonTitanMultimodalEmbeddingG1Config, diff --git a/litellm/anthropic_interface/messages/__init__.py b/litellm/anthropic_interface/messages/__init__.py index 16bb5f3d462..d7ff53a1763 100644 --- a/litellm/anthropic_interface/messages/__init__.py +++ b/litellm/anthropic_interface/messages/__init__.py @@ -37,6 +37,7 @@ async def acreate( tools: Optional[List[Dict]] = None, top_k: Optional[int] = None, top_p: Optional[float] = None, + container: Optional[Dict] = None, **kwargs ) -> Union[AnthropicMessagesResponse, AsyncIterator]: """ @@ -56,6 +57,7 @@ async def acreate( tools (List[Dict], optional): List of tool definitions top_k (int, optional): Top K sampling parameter top_p (float, optional): Nucleus sampling parameter + container (Dict, optional): Container config with skills for code execution **kwargs: Additional arguments Returns: @@ -75,6 +77,7 @@ async def acreate( tools=tools, top_k=top_k, top_p=top_p, + container=container, **kwargs, ) @@ -93,6 +96,7 @@ def create( tools: Optional[List[Dict]] = None, top_k: Optional[int] = None, top_p: Optional[float] = None, + container: Optional[Dict] = None, **kwargs ) -> Union[ AnthropicMessagesResponse, @@ -135,5 +139,6 @@ def create( tools=tools, top_k=top_k, top_p=top_p, + container=container, **kwargs, ) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 612bec239ba..6f9aa192f9b 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -167,24 +167,28 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) elif role == "tool": # Convert tool message to function call output format - # Transform content to responses format (handles str, list, and other types) - # _convert_content_to_responses_format always returns List[Dict[str, Any]] + # The Responses API expects 'output' to be a string, not a list if content is None: - transformed_output: list[dict[str, Any]] = [] - elif isinstance(content, (str, list)): - transformed_output = self._convert_content_to_responses_format( - content, "tool" - ) + output_str = "" + elif isinstance(content, str): + output_str = content + elif isinstance(content, list): + # If content is a list, extract text parts and join them + text_parts = [] + for item in content: + if isinstance(item, str): + text_parts.append(item) + elif isinstance(item, dict) and item.get("type") == "text": + text_parts.append(item.get("text", "")) + output_str = " ".join(text_parts) if text_parts else str(content) else: - # Fallback: convert unexpected types to string first - transformed_output = self._convert_content_to_responses_format( - str(content), "tool" - ) + # Fallback: convert unexpected types to string + output_str = str(content) input_items.append( { "type": "function_call_output", "call_id": tool_call_id, - "output": transformed_output, + "output": output_str, } ) elif role == "assistant" and tool_calls and isinstance(tool_calls, list): @@ -345,6 +349,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): index = 0 reasoning_content: Optional[str] = None + # Collect all tool calls to put them in a single choice + # (Chat Completions API expects all tool calls in one message) + accumulated_tool_calls: List[Dict[str, Any]] = [] + tool_call_index = 0 + for item in output_items: if isinstance(item, ResponseReasoningItem): for summary_item in item.summary: @@ -378,20 +387,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call( tool_call_item=item, - index=index, + index=tool_call_index, ) - - msg = Message( - content=None, - tool_calls=[tool_call_dict], - reasoning_content=reasoning_content, - ) - - choices.append( - Choices(message=msg, finish_reason="tool_calls", index=index) - ) - reasoning_content = None # flush reasoning content - index += 1 + accumulated_tool_calls.append(tool_call_dict) + tool_call_index += 1 elif isinstance(item, dict) and handle_raw_dict_callback is not None: # Handle raw dict responses (e.g., from GPT-5 Codex) @@ -401,6 +400,18 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): else: pass # don't fail request if item in list is not supported + # If we accumulated tool calls, create a single choice with all of them + if accumulated_tool_calls: + msg = Message( + content=None, + tool_calls=accumulated_tool_calls, + reasoning_content=reasoning_content, + ) + choices.append( + Choices(message=msg, finish_reason="tool_calls", index=index) + ) + reasoning_content = None + return choices def transform_response( # noqa: PLR0915 @@ -492,7 +503,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def _convert_content_str_to_input_text( self, content: str, role: str ) -> Dict[str, Any]: - if role == "user" or role == "system": + if role == "user" or role == "system" or role == "tool": return {"type": "input_text", "text": content} else: return {"type": "output_text", "text": content} diff --git a/litellm/constants.py b/litellm/constants.py index 23766874841..511cbafc748 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -892,6 +892,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[ "qwen2", "twelvelabs", "openai", + "stability", ] BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[ diff --git a/litellm/images/main.py b/litellm/images/main.py index b711aa31c05..03c0e36ad93 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -33,6 +33,7 @@ from litellm.main import ( base_llm_aiohttp_handler, base_llm_http_handler, bedrock_image_generation, + bedrock_image_edit, openai_chat_completions, openai_image_variations, ) @@ -670,7 +671,7 @@ def image_variation( @client -def image_edit( +def image_edit( # noqa: PLR0915 image: Union[FileTypes, List[FileTypes]], prompt: str, model: Optional[str] = None, @@ -695,6 +696,29 @@ def image_edit( """ local_vars = locals() try: + openai_params = [ + "user", + "request_timeout", + "api_base", + "api_version", + "api_key", + "deployment_id", + "organization", + "base_url", + "default_headers", + "timeout", + "max_retries", + "n", + "quality", + "size", + "style", + "async_call", + ] + litellm_params_list = all_litellm_params + default_params = openai_params + litellm_params_list + non_default_params = { + k: v for k, v in kwargs.items() if k not in default_params + } # model-specific params - pass them straight to the model/provider litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("async_call", False) is True @@ -788,7 +812,6 @@ def image_edit( image_edit_optional_params: ImageEditOptionalRequestParams = ( _get_ImageEditRequestUtils().get_requested_image_edit_optional_param(local_vars) ) - # Get optional parameters for the responses API image_edit_request_params: Dict = ( _get_ImageEditRequestUtils().get_optional_params_image_edit( @@ -812,6 +835,42 @@ def image_edit( custom_llm_provider=custom_llm_provider, ) + # Route bedrock to its specific handler (AWS signing required) + if custom_llm_provider == "bedrock": + if model is None: + raise Exception("Model needs to be set for bedrock") + image_edit_request_params.update(non_default_params) + return bedrock_image_edit.image_edit( # type: ignore + model=model, + image=images, + prompt=prompt, + timeout=timeout, + logging_obj=litellm_logging_obj, + optional_params=image_edit_request_params, + model_response=ImageResponse(), + aimage_edit=_is_async, + client=kwargs.get("client"), + api_base=kwargs.get("api_base"), + extra_headers=extra_headers, + api_key=kwargs.get("api_key"), + ) + elif custom_llm_provider == "stability": + image_edit_request_params.update(non_default_params) + return base_llm_http_handler.image_edit_handler( + model=model, + image=images, + prompt=prompt, + image_edit_provider_config=image_edit_provider_config, + image_edit_optional_request_params=image_edit_request_params, + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + logging_obj=litellm_logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout or DEFAULT_REQUEST_TIMEOUT, + _is_async=_is_async, + client=kwargs.get("client"), + ) # Call the handler with _is_async flag instead of directly calling the async handler return base_llm_http_handler.image_edit_handler( model=model, diff --git a/litellm/images/utils.py b/litellm/images/utils.py index fdf240ba2af..fa271b61b6a 100644 --- a/litellm/images/utils.py +++ b/litellm/images/utils.py @@ -82,7 +82,6 @@ class ImageEditRequestUtils: filtered_params = { k: v for k, v in params.items() if k in valid_keys and v is not None } - return cast(ImageEditOptionalRequestParams, filtered_params) @staticmethod diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 6892ba3426a..fe0ce208ee6 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -240,6 +240,28 @@ class CustomGuardrail(CustomLogger): return metadata["disable_global_guardrail"] return False + def _is_valid_response_type(self, result: Any) -> bool: + """ + Check if result is a valid LLMResponseTypes instance. + + Safely handles TypedDict types which don't support isinstance checks. + For non-LiteLLM responses (like passthrough httpx.Response), returns True + to allow them through. + """ + if result is None: + return False + + try: + # Try isinstance check on valid types that support it + response_types = get_args(LLMResponseTypes) + return isinstance(result, response_types) + except TypeError as e: + # TypedDict types don't support isinstance checks + # In this case, we can't validate the type, so we allow it through + if "TypedDict" in str(e): + return True + raise + def get_guardrail_from_metadata( self, data: dict ) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]: @@ -342,7 +364,7 @@ class CustomGuardrail(CustomLogger): response=response, ) - if result is None or not isinstance(result, get_args(LLMResponseTypes)): + if not self._is_valid_response_type(result): return response return result diff --git a/litellm/integrations/langfuse/langfuse_prompt_management.py b/litellm/integrations/langfuse/langfuse_prompt_management.py index adc8ae61d01..8f73eabad44 100644 --- a/litellm/integrations/langfuse/langfuse_prompt_management.py +++ b/litellm/integrations/langfuse/langfuse_prompt_management.py @@ -294,6 +294,11 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge self.async_log_success_event, kwargs, response_obj, start_time, end_time ) + def log_failure_event(self, kwargs, response_obj, start_time, end_time): + return run_async_function( + self.async_log_failure_event, kwargs, response_obj, start_time, end_time + ) + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): standard_callback_dynamic_params = kwargs.get( "standard_callback_dynamic_params" diff --git a/litellm/litellm_core_utils/api_route_to_call_types.py b/litellm/litellm_core_utils/api_route_to_call_types.py index 35f83de1dd7..4146ff6d6a6 100644 --- a/litellm/litellm_core_utils/api_route_to_call_types.py +++ b/litellm/litellm_core_utils/api_route_to_call_types.py @@ -5,10 +5,12 @@ This dictionary maps each API endpoint to the CallTypes that can be used for tha Each route can have both async (prefixed with 'a') and sync call types. """ +from typing import List, Optional + from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes -def get_call_types_for_route(route: str) -> list: +def get_call_types_for_route(route: str) -> Optional[List[CallTypes]]: """ Get the list of CallTypes for a given API route. @@ -16,9 +18,9 @@ def get_call_types_for_route(route: str) -> list: route: API route path (e.g., "/chat/completions") Returns: - List of CallTypes for that route, or empty list if route not found + List of CallTypes for that route, or None if route not found """ - return API_ROUTE_TO_CALL_TYPES.get(route, []) + return API_ROUTE_TO_CALL_TYPES.get(route, None) def get_routes_for_call_type(call_type: CallTypes) -> list: diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index ef2183a4556..232d9bfc5d1 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -674,7 +674,7 @@ class CostCalculatorUtils: from litellm.llms.azure_ai.image_generation.cost_calculator import ( cost_calculator as azure_ai_image_cost_calculator, ) - from litellm.llms.bedrock.image.cost_calculator import ( + from litellm.llms.bedrock.image_generation.cost_calculator import ( cost_calculator as bedrock_image_cost_calculator, ) from litellm.llms.gemini.image_generation.cost_calculator import ( diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 4c202b9eec0..9cfbf1b6d8d 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -169,7 +169,7 @@ class LiteLLMAnthropicMessagesAdapter: """ Which anthropic params, we need to translate to the openai format. """ - return ["messages", "metadata", "system", "tool_choice", "tools"] + return ["messages", "metadata", "system", "tool_choice", "tools", "thinking"] def translate_anthropic_messages_to_openai( # noqa: PLR0915 self, @@ -420,6 +420,35 @@ class LiteLLMAnthropicMessagesAdapter: return new_messages + def translate_anthropic_thinking_to_openai( + self, thinking: Dict[str, Any] + ) -> Optional[str]: + """ + Translate Anthropic's thinking parameter to OpenAI's reasoning_effort. + + Anthropic thinking format: {'type': 'enabled'|'disabled', 'budget_tokens': int} + OpenAI reasoning_effort: 'none' | 'minimal' | 'low' | 'medium' | 'high' | 'xhigh' | 'default' + """ + if not isinstance(thinking, dict): + return None + + thinking_type = thinking.get("type", "disabled") + + if thinking_type == "disabled": + return None + elif thinking_type == "enabled": + budget_tokens = thinking.get("budget_tokens", 0) + if budget_tokens >= 10000: + return "high" + elif budget_tokens >= 5000: + return "medium" + elif budget_tokens >= 2000: + return "low" + else: + return "minimal" + + return None + def translate_anthropic_tool_choice_to_openai( self, tool_choice: AnthropicMessagesToolChoice ) -> ChatCompletionToolChoiceValues: @@ -529,6 +558,16 @@ class LiteLLMAnthropicMessagesAdapter: tools=cast(List[AllAnthropicToolsValues], tools) ) + ## CONVERT THINKING + if "thinking" in anthropic_message_request: + thinking = anthropic_message_request["thinking"] + if thinking: + reasoning_effort = self.translate_anthropic_thinking_to_openai( + thinking=cast(Dict[str, Any], thinking) + ) + if reasoning_effort: + new_kwargs["reasoning_effort"] = reasoning_effort + translatable_params = self.translatable_anthropic_params() for k, v in anthropic_message_request.items(): if k not in translatable_params: # pass remaining params as is diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index cc9334ae68b..908b46c11e2 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -119,6 +119,7 @@ def anthropic_messages_handler( tools: Optional[List[Dict]] = None, top_k: Optional[int] = None, top_p: Optional[float] = None, + container: Optional[Dict] = None, api_key: Optional[str] = None, api_base: Optional[str] = None, client: Optional[AsyncHTTPHandler] = None, @@ -131,6 +132,9 @@ def anthropic_messages_handler( ]: """ Makes Anthropic `/v1/messages` API calls In the Anthropic API Spec + + Args: + container: Container config with skills for code execution """ from litellm.types.utils import LlmProviders diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index e53ac36a00d..71d21001cc3 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -365,6 +365,10 @@ class BaseAWSLLM: model_id = BaseAWSLLM._get_model_id_from_model_with_spec( model_id, spec="qwen3" ) + elif provider == "stability" and "stability/" in model_id: + model_id = BaseAWSLLM._get_model_id_from_model_with_spec( + model_id, spec="stability" + ) return model_id @staticmethod diff --git a/litellm/llms/bedrock/image_edit/__init__.py b/litellm/llms/bedrock/image_edit/__init__.py new file mode 100644 index 00000000000..f3a0e61067d --- /dev/null +++ b/litellm/llms/bedrock/image_edit/__init__.py @@ -0,0 +1,10 @@ +""" +Bedrock Image Edit Module + +Handles image edit operations for Bedrock stability models. +""" + +from .handler import BedrockImageEdit + +__all__ = ["BedrockImageEdit"] + diff --git a/litellm/llms/bedrock/image_edit/handler.py b/litellm/llms/bedrock/image_edit/handler.py new file mode 100644 index 00000000000..b4b6c8d7622 --- /dev/null +++ b/litellm/llms/bedrock/image_edit/handler.py @@ -0,0 +1,310 @@ +""" +Bedrock Image Edit Handler + +Handles image edit requests for Bedrock stability models. +""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING, Any, Optional, Union + +import httpx +from pydantic import BaseModel + +import litellm +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.llms.bedrock.image_edit.stability_transformation import ( + BedrockStabilityImageEditConfig, +) +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.types.utils import ImageResponse + +from ..base_aws_llm import BaseAWSLLM +from ..common_utils import BedrockError + +if TYPE_CHECKING: + from botocore.awsrequest import AWSPreparedRequest +else: + AWSPreparedRequest = Any + + +class BedrockImageEditPreparedRequest(BaseModel): + """ + Internal/Helper class for preparing the request for bedrock image edit + """ + + endpoint_url: str + prepped: AWSPreparedRequest + body: bytes + data: dict + + +class BedrockImageEdit(BaseAWSLLM): + """ + Bedrock Image Edit handler + """ + + @classmethod + def get_config_class(cls, model: str | None): + if BedrockStabilityImageEditConfig._is_stability_edit_model(model): + return BedrockStabilityImageEditConfig + else: + raise ValueError(f"Unsupported model for bedrock image edit: {model}") + + def image_edit( + self, + model: str, + image: list, + prompt: str, + model_response: ImageResponse, + optional_params: dict, + logging_obj: LitellmLogging, + timeout: Optional[Union[float, httpx.Timeout]], + aimage_edit: bool = False, + api_base: Optional[str] = None, + extra_headers: Optional[dict] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + api_key: Optional[str] = None, + ): + prepared_request = self._prepare_request( + model=model, + image=image, + prompt=prompt, + optional_params=optional_params, + api_base=api_base, + extra_headers=extra_headers, + logging_obj=logging_obj, + api_key=api_key, + ) + + if aimage_edit is True: + return self.async_image_edit( + prepared_request=prepared_request, + timeout=timeout, + model=model, + logging_obj=logging_obj, + prompt=prompt, + model_response=model_response, + client=( + client + if client is not None and isinstance(client, AsyncHTTPHandler) + else None + ), + ) + + if client is None or not isinstance(client, HTTPHandler): + client = _get_httpx_client() + try: + response = client.post(url=prepared_request.endpoint_url, headers=prepared_request.prepped.headers, data=prepared_request.body) # type: ignore + response.raise_for_status() + except httpx.HTTPStatusError as err: + error_code = err.response.status_code + raise BedrockError(status_code=error_code, message=err.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + + ### FORMAT RESPONSE TO OPENAI FORMAT ### + model_response = self._transform_response_dict_to_openai_response( + model_response=model_response, + model=model, + logging_obj=logging_obj, + prompt=prompt, + response=response, + data=prepared_request.data, + ) + return model_response + + async def async_image_edit( + self, + prepared_request: BedrockImageEditPreparedRequest, + timeout: Optional[Union[float, httpx.Timeout]], + model: str, + logging_obj: LitellmLogging, + prompt: str, + model_response: ImageResponse, + client: Optional[AsyncHTTPHandler] = None, + ) -> ImageResponse: + """ + Asynchronous handler for bedrock image edit + """ + async_client = client or get_async_httpx_client( + llm_provider=litellm.LlmProviders.BEDROCK, + params={"timeout": timeout}, + ) + + try: + response = await async_client.post(url=prepared_request.endpoint_url, headers=prepared_request.prepped.headers, data=prepared_request.body) # type: ignore + response.raise_for_status() + except httpx.HTTPStatusError as err: + error_code = err.response.status_code + raise BedrockError(status_code=error_code, message=err.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + + ### FORMAT RESPONSE TO OPENAI FORMAT ### + model_response = self._transform_response_dict_to_openai_response( + model=model, + logging_obj=logging_obj, + prompt=prompt, + response=response, + data=prepared_request.data, + model_response=model_response, + ) + return model_response + + def _prepare_request( + self, + model: str, + image: list, + prompt: str, + optional_params: dict, + api_base: Optional[str], + extra_headers: Optional[dict], + logging_obj: LitellmLogging, + api_key: Optional[str], + ) -> BedrockImageEditPreparedRequest: + """ + Prepare the request body, headers, and endpoint URL for the Bedrock Image Edit API + + Args: + model (str): The model to use for the image edit + image (list): The images to edit + prompt (str): The prompt for the edit + optional_params (dict): The optional parameters for the image edit + api_base (Optional[str]): The base URL for the Bedrock API + extra_headers (Optional[dict]): The extra headers to include in the request + logging_obj (LitellmLogging): The logging object to use for logging + api_key (Optional[str]): The API key to use + + Returns: + BedrockImageEditPreparedRequest: The prepared request object + """ + boto3_credentials_info = self._get_boto_credentials_from_optional_params( + optional_params, model + ) + + # Use the existing ARN-aware provider detection method + bedrock_provider = self.get_bedrock_invoke_provider(model) + ### SET RUNTIME ENDPOINT ### + modelId = self.get_bedrock_model_id( + model=model, + provider=bedrock_provider, + optional_params=optional_params, + ) + _, proxy_endpoint_url = self.get_runtime_endpoint( + api_base=api_base, + aws_bedrock_runtime_endpoint=boto3_credentials_info.aws_bedrock_runtime_endpoint, + aws_region_name=boto3_credentials_info.aws_region_name, + ) + proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke" + data = self._get_request_body( + model=model, + image=image, + prompt=prompt, + optional_params=optional_params, + ) + + # Make POST Request + body = json.dumps(data).encode("utf-8") + headers = {"Content-Type": "application/json"} + if extra_headers is not None: + headers = {"Content-Type": "application/json", **extra_headers} + + prepped = self.get_request_headers( + credentials=boto3_credentials_info.credentials, + aws_region_name=boto3_credentials_info.aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=body, + headers=headers, + api_key=api_key, + ) + + ## LOGGING + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, + ) + return BedrockImageEditPreparedRequest( + endpoint_url=proxy_endpoint_url, + prepped=prepped, + body=body, + data=data, + ) + + def _get_request_body( + self, + model: str, + image: list, + prompt: str, + optional_params: dict, + ) -> dict: + """ + Get the request body for the Bedrock Image Edit API + + Checks the model/provider and transforms the request body accordingly + + Returns: + dict: The request body to use for the Bedrock Image Edit API + """ + config_class = self.get_config_class(model=model) + config_instance = config_class() + request_body = config_instance.transform_image_edit_request( + model=model, + prompt=prompt, + image=image[0] if image else None, + image_edit_optional_request_params=optional_params, + litellm_params={}, + headers={}, + ) + return dict(request_body) + + def _transform_response_dict_to_openai_response( + self, + model_response: ImageResponse, + model: str, + logging_obj: LitellmLogging, + prompt: str, + response: httpx.Response, + data: dict, + ) -> ImageResponse: + """ + Transforms the Image Edit response from Bedrock to OpenAI format + """ + + ## LOGGING + if logging_obj is not None: + logging_obj.post_call( + input=prompt, + api_key="", + original_response=response.text, + additional_args={"complete_input_dict": data}, + ) + verbose_logger.debug("raw model_response: %s", response.text) + response_dict = response.json() + if response_dict is None: + raise ValueError("Error in response object format, got None") + + config_class = self.get_config_class(model=model) + config_instance = config_class() + + model_response = config_instance.transform_image_edit_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) + + return model_response + diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py new file mode 100644 index 00000000000..bcaf0923f69 --- /dev/null +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -0,0 +1,377 @@ +""" +Bedrock Stability AI Image Edit Transformation + +Handles transformation between OpenAI-compatible format and Bedrock Stability AI Image Edit API format. + +Supported models: +- stability.stable-conservative-upscale-v1:0 +- stability.stable-creative-upscale-v1:0 +- stability.stable-fast-upscale-v1:0 +- stability.stable-outpaint-v1:0 +- stability.stable-image-control-sketch-v1:0 +- stability.stable-image-control-structure-v1:0 +- stability.stable-image-erase-object-v1:0 +- stability.stable-image-inpaint-v1:0 +- stability.stable-image-remove-background-v1:0 +- stability.stable-image-search-recolor-v1:0 +- stability.stable-image-search-replace-v1:0 +- stability.stable-image-style-guide-v1:0 +- stability.stable-style-transfer-v1:0 + +API Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters.html +""" + +import json +import base64 +from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple + +import httpx + +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams +from litellm.types.llms.stability import ( + OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO, +) +from litellm.types.utils import FileTypes, ImageObject, ImageResponse +from litellm.utils import get_model_info + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class BedrockStabilityImageEditConfig(BaseImageEditConfig): + """ + Configuration for Bedrock Stability AI image edit. + + Supports all Stability image edit operations through Bedrock. + """ + + @classmethod + def _is_stability_edit_model(cls, model: Optional[str] = None) -> bool: + """ + Returns True if the model is a Bedrock Stability edit model. + + Bedrock Stability edit models follow this pattern: + stability.stable-conservative-upscale-v1:0 + stability.stable-creative-upscale-v1:0 + stability.stable-fast-upscale-v1:0 + stability.stable-outpaint-v1:0 + stability.stable-image-inpaint-v1:0 + stability.stable-image-erase-object-v1:0 + etc. + """ + if model: + model_lower = model.lower() + if "stability." in model_lower and any([ + "upscale" in model_lower, + "outpaint" in model_lower, + "inpaint" in model_lower, + "erase" in model_lower, + "remove-background" in model_lower, + "search-recolor" in model_lower, + "search-replace" in model_lower, + "control-sketch" in model_lower, + "control-structure" in model_lower, + "style-guide" in model_lower, + "style-transfer" in model_lower, + ]): + return True + return False + + def get_supported_openai_params( + self, model: str + ) -> list: + """ + Return list of OpenAI params supported by Bedrock Stability. + """ + return [ + "n", # Number of images (Stability always returns 1, we can loop) + "size", # Maps to aspect_ratio + "response_format", # b64_json or url (Stability only returns b64) + "mask", + ] + + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """ + Map OpenAI parameters to Bedrock Stability parameters. + + OpenAI -> Stability mappings: + - size -> aspect_ratio + - n -> (handled separately, Stability returns 1 image per request) + """ + supported_params = self.get_supported_openai_params(model) + # Define mapping from OpenAI params to Stability params + param_mapping = { + "size": "aspect_ratio", + # "n" and "response_format" are handled separately + } + + # Create a copy to not mutate original - convert TypedDict to regular dict + mapped_params: Dict[str, Any] = dict(image_edit_optional_params) + + for k, v in image_edit_optional_params.items(): + if k in param_mapping: + # Map param if mapping exists and value is valid + if k == "size" and v in OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO: + mapped_params[param_mapping[k]] = OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO[v] # type: ignore + # Don't copy "size" itself to final dict + elif k == "n": + # Store for logic but do not add to outgoing params + mapped_params["_n"] = v + elif k == "response_format": + # Only b64 supported at Stability; store for postprocessing + mapped_params["_response_format"] = v + elif k not in supported_params: + if not drop_params: + raise ValueError( + f"Parameter {k} is not supported for model {model}. " + f"Supported parameters are {supported_params}. " + f"Set drop_params=True to drop unsupported parameters." + ) + # Otherwise, param will simply be dropped + else: + # param is supported and not mapped, keep as-is + continue + + # Remove OpenAI params that have been mapped unless they're in stability + for mapped in ["size", "n", "response_format"]: + if mapped in mapped_params: + del mapped_params[mapped] + + return mapped_params + + def transform_image_edit_request( + self, + model: str, + prompt: str, + image: FileTypes, + image_edit_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[Dict, Any]: + """ + Transform OpenAI-style request to Bedrock Stability request format. + + Returns the request body dict that will be JSON-encoded by the handler. + """ + # Build Bedrock Stability request + data: Dict[str, Any] = { + "prompt": prompt, + "output_format": "png", # Default to PNG + } + + # Convert image to base64 + image_b64: str + if hasattr(image, 'read') and callable(getattr(image, 'read', None)): + # File-like object (e.g., BufferedReader from open()) + image_bytes = image.read() # type: ignore + image_b64 = base64.b64encode(image_bytes).decode('utf-8') # type: ignore + elif isinstance(image, bytes): + # Raw bytes + image_b64 = base64.b64encode(image).decode('utf-8') + elif isinstance(image, str): + # Already a base64 string + image_b64 = image + else: + # Try to handle as bytes + image_b64 = base64.b64encode(bytes(image)).decode('utf-8') # type: ignore + + data["image"] = image_b64 + + # Add optional params (already mapped in map_openai_params) + for key, value in image_edit_optional_request_params.items(): # type: ignore + # Skip internal params (prefixed with _) + if key.startswith("_") or value is None: + continue + + # File-like optional params (mask, init_image, style_image, etc.) + if key in ["mask", "init_image", "style_image"]: + # Handle case where value might be in a list + file_value = value + if isinstance(value, list) and len(value) > 0: + file_value = value[0] + + if hasattr(file_value, 'read') and callable(getattr(file_value, 'read', None)): + file_bytes = file_value.read() # type: ignore + elif isinstance(file_value, bytes): + file_bytes = file_value + elif isinstance(file_value, str): + # Already a base64 string + data[key] = file_value + continue + else: + file_bytes = file_value # type: ignore + + if isinstance(file_bytes, bytes): + file_b64 = base64.b64encode(file_bytes).decode('utf-8') + else: + file_b64 = str(file_bytes) + data[key] = file_b64 + continue + + # Supported text fields + if key in [ + "negative_prompt", + "aspect_ratio", + "seed", + "output_format", + "model", + "mode", + "strength", + "style_preset", + "creativity", + "control_strength", + "grow_mask", + "left", + "right", + "up", + "down", + "select_prompt", + "search_prompt", + "fidelity", + "composition_fidelity", + "style_strength", + "change_strength", + ]: + data[key] = value # type: ignore + + return data, {} + + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> ImageResponse: + """ + Transform Bedrock Stability response to OpenAI-compatible ImageResponse. + + Bedrock returns: {"images": ["base64..."], "finish_reasons": [null], "seeds": [123]} + OpenAI expects: {"data": [{"b64_json": "base64..."}], "created": timestamp} + """ + try: + response_data = raw_response.json() + with open("response_data.json", "w") as f: + json.dump(response_data, f) + except Exception as e: + raise self.get_error_class( + error_message=f"Error parsing Bedrock Stability response: {e}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check for errors in response + if "errors" in response_data: + raise self.get_error_class( + error_message=f"Bedrock Stability error: {response_data['errors']}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check finish_reasons + finish_reasons = response_data.get("finish_reasons", []) + if finish_reasons and finish_reasons[0]: + raise self.get_error_class( + error_message=f"Bedrock Stability error: {finish_reasons[0]}", + status_code=400, + headers=raw_response.headers, + ) + + model_response = ImageResponse() + if not model_response.data: + model_response.data = [] + + # Extract images from response + images = response_data.get("images", []) + if images: + for image_b64 in images: + if image_b64: + model_response.data.append( + ImageObject( + b64_json=image_b64, + url=None, + revised_prompt=None, + ) + ) + + if not hasattr(model_response, "_hidden_params"): + model_response._hidden_params = {} + if "additional_headers" not in model_response._hidden_params: + model_response._hidden_params["additional_headers"] = {} + + # Set cost based on model + model_info = get_model_info(model, custom_llm_provider="bedrock") + cost_per_image = model_info.get("output_cost_per_image", 0) + if cost_per_image is not None: + model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = float(cost_per_image) + + return model_response + + def use_multipart_form_data(self) -> bool: + """ + Bedrock Stability uses JSON format, not multipart/form-data. + """ + return False + + def get_complete_url( + self, + model: str, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Get the complete URL for the Bedrock Image Edit API. + + For Bedrock, this is handled by the handler which constructs the endpoint URL + based on the model ID and AWS region. This method is required by the base class + but the actual URL construction happens in BedrockImageEdit.image_edit(). + + Returns a placeholder - the real endpoint is constructed in the handler. + """ + # Bedrock URLs are constructed in the handler using boto3 + # This is a placeholder for the abstract method requirement + return "bedrock://image-edit" + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + ) -> dict: + """ + Validate environment for Bedrock Stability image edit. + + For Bedrock, AWS credentials are managed by the BaseAWSLLM class. + This method validates that headers are properly set up. + + Args: + headers: The request headers to validate/update + model: The model name being used + api_key: Optional API key (not used for Bedrock, which uses AWS credentials) + + Returns: + Updated headers dict + """ + if headers is None: + headers = {} + + # Bedrock uses AWS credentials, not API keys + # Headers are set up by the handler's get_request_headers() method + # This just ensures basic headers are present + if "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + + return headers + diff --git a/litellm/llms/bedrock/image/amazon_nova_canvas_transformation.py b/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py similarity index 100% rename from litellm/llms/bedrock/image/amazon_nova_canvas_transformation.py rename to litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py diff --git a/litellm/llms/bedrock/image/amazon_stability1_transformation.py b/litellm/llms/bedrock/image_generation/amazon_stability1_transformation.py similarity index 100% rename from litellm/llms/bedrock/image/amazon_stability1_transformation.py rename to litellm/llms/bedrock/image_generation/amazon_stability1_transformation.py diff --git a/litellm/llms/bedrock/image/amazon_stability3_transformation.py b/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py similarity index 100% rename from litellm/llms/bedrock/image/amazon_stability3_transformation.py rename to litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py diff --git a/litellm/llms/bedrock/image/amazon_titan_transformation.py b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py similarity index 100% rename from litellm/llms/bedrock/image/amazon_titan_transformation.py rename to litellm/llms/bedrock/image_generation/amazon_titan_transformation.py diff --git a/litellm/llms/bedrock/image/cost_calculator.py b/litellm/llms/bedrock/image_generation/cost_calculator.py similarity index 87% rename from litellm/llms/bedrock/image/cost_calculator.py rename to litellm/llms/bedrock/image_generation/cost_calculator.py index bc1a57b8aec..b04acc3e809 100644 --- a/litellm/llms/bedrock/image/cost_calculator.py +++ b/litellm/llms/bedrock/image_generation/cost_calculator.py @@ -1,6 +1,6 @@ from typing import Optional -from litellm.llms.bedrock.image.image_handler import BedrockImageGeneration +from litellm.llms.bedrock.image_generation.image_handler import BedrockImageGeneration from litellm.types.utils import ImageResponse diff --git a/litellm/llms/bedrock/image/image_handler.py b/litellm/llms/bedrock/image_generation/image_handler.py similarity index 97% rename from litellm/llms/bedrock/image/image_handler.py rename to litellm/llms/bedrock/image_generation/image_handler.py index 2e76596eefe..0a4cde90b27 100644 --- a/litellm/llms/bedrock/image/image_handler.py +++ b/litellm/llms/bedrock/image_generation/image_handler.py @@ -9,13 +9,13 @@ from pydantic import BaseModel import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging -from litellm.llms.bedrock.image.amazon_nova_canvas_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import ( AmazonNovaCanvasConfig, ) -from litellm.llms.bedrock.image.amazon_stability3_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_stability3_transformation import ( AmazonStability3Config, ) -from litellm.llms.bedrock.image.amazon_titan_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_titan_transformation import ( AmazonTitanImageGenerationConfig, ) from litellm.llms.custom_httpx.http_handler import ( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index a777ea3c42e..34ea598a655 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -3761,7 +3761,7 @@ class BaseLLMHTTPHandler: input=prompt, api_key="", additional_args={ - "complete_input_dict": data, + "complete_input_dict": files, "api_base": api_base, "headers": headers, }, diff --git a/litellm/llms/litellm_proxy/skills/README.md b/litellm/llms/litellm_proxy/skills/README.md new file mode 100644 index 00000000000..1dfeff1a42c --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/README.md @@ -0,0 +1,381 @@ +# LiteLLM Skills - Database-Backed Skills Storage + +This module provides database-backed skills storage as an alternative to Anthropic's cloud-based Skills API. It enables using skills with **any LLM provider** (Bedrock, OpenAI, Azure, etc.) by storing skills locally and converting them to tools + system prompt injection. + +## Architecture + +```mermaid +flowchart TB + subgraph "Skill Creation" + A[User creates skill with ZIP file] --> B{custom_llm_provider?} + B -->|anthropic| C[Forward to Anthropic API] + B -->|litellm_proxy| D[Store in LiteLLM Database] + + D --> E[Extract & store:
- display_title
- description
- instructions
- file_content ZIP] + end + + subgraph "Skill Usage in Messages API" + F[Request with container.skills] --> G[SkillsInjectionHook] + G --> H{skill_id prefix?} + + H -->|"litellm:skill_abc"| I[Fetch from LiteLLM DB] + H -->|"skill_xyz" no prefix| J[Pass to Anthropic as native skill] + + I --> K{Model provider?} + K -->|Anthropic API| L[Convert to tools] + K -->|Bedrock/OpenAI/etc| M[Convert to tools +
Inject SKILL.md into system prompt] + + J --> N[Keep in container.skills] + end + + subgraph "Skill Resolution for Non-Anthropic" + M --> O[Extract SKILL.md from ZIP] + O --> P[Add to system prompt:
# Available Skills
## Skill: My Skill
SKILL.md content...] + P --> Q[Create OpenAI-style tool:
type: function
name: skill_id
description: instructions] + Q --> R[Send to LLM Provider] + end +``` + +## Automatic Code Execution + +For skills that include executable code (Python files), LiteLLM automatically handles: + +1. **Pre-call hook** (`async_pre_call_hook`): Adds `litellm_code_execution` tool, injects SKILL.md content +2. **Post-call hook** (`async_post_call_success_deployment_hook`): Detects tool calls, executes code in Docker sandbox, continues loop +3. **Returns files**: Generated files (GIFs, images, etc.) returned directly on response + +```mermaid +sequenceDiagram + participant User + participant LiteLLM as LiteLLM SDK + participant PreHook as async_pre_call_hook + participant LLM as LLM Provider + participant PostHook as async_post_call_success_deployment_hook + participant Sandbox as Docker Sandbox + + User->>LiteLLM: litellm.acompletion(model, messages, container={skills: [...]}) + + Note over LiteLLM,PreHook: PRE-CALL HOOK + LiteLLM->>PreHook: Intercept request + PreHook->>PreHook: Fetch skill from DB (litellm:skill_id) + PreHook->>PreHook: Extract SKILL.md from ZIP + PreHook->>PreHook: Inject SKILL.md into system prompt + PreHook->>PreHook: Add litellm_code_execution tool + PreHook->>PreHook: Store skill files in metadata + PreHook-->>LiteLLM: Modified request + + LiteLLM->>LLM: Forward to provider (OpenAI/Bedrock/etc) + LLM-->>LiteLLM: Response with tool_calls + + Note over LiteLLM,PostHook: POST-CALL HOOK (Agentic Loop) + LiteLLM->>PostHook: Check response + + loop Until no more tool calls + PostHook->>PostHook: Check for litellm_code_execution tool call + alt Has code execution tool call + PostHook->>Sandbox: Execute Python code + Sandbox->>Sandbox: Copy skill files to /sandbox + Sandbox->>Sandbox: Install requirements.txt + Sandbox->>Sandbox: Run code + Sandbox-->>PostHook: Result + generated files + PostHook->>PostHook: Add tool result to messages + PostHook->>LLM: Make another LLM call + LLM-->>PostHook: New response + else No code execution + PostHook->>PostHook: Break loop + end + end + + PostHook->>PostHook: Attach files to response._litellm_generated_files + PostHook-->>LiteLLM: Modified response with files + LiteLLM-->>User: Final response with generated files +``` + +```python +import litellm +from litellm.proxy.hooks.litellm_skills import SkillsInjectionHook + +# Register the hook (done once at startup) +hook = SkillsInjectionHook() +litellm.callbacks.append(hook) + +# ONE request - LiteLLM handles everything automatically +# The container parameter triggers the SkillsInjectionHook +response = await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Create a bouncing ball GIF"}], + container={ + "skills": [{"type": "custom", "skill_id": "litellm:skill_abc123"}] + }, +) + +# Files are attached directly to response +generated_files = response._litellm_generated_files +for f in generated_files: + print(f"Generated: {f['name']} ({f['size']} bytes)") + # f['content_base64'] contains the file data +``` + +This mimics Anthropic's behavior - no manual agentic loop needed! + +### How it works + +The `SkillsInjectionHook` uses two hooks: + +1. **`async_pre_call_hook`** (proxy only): Transforms the request before LLM call + - Fetches skills from DB + - Injects SKILL.md into system prompt + - Adds `litellm_code_execution` tool + - Sets `_litellm_code_execution_enabled=True` in metadata + +2. **`async_post_call_success_deployment_hook`** (SDK + proxy): Called after LLM response + - Checks if response has `litellm_code_execution` tool call + - Executes code in Docker sandbox + - Adds result to messages, makes another LLM call + - Repeats until model gives final response + - Attaches generated files to `response._litellm_generated_files` + +## File Structure + +``` +litellm/llms/litellm_proxy/skills/ +├── __init__.py # Exports all skill components +├── handler.py # LiteLLMSkillsHandler - database CRUD operations (Prisma) +├── transformation.py # LiteLLMSkillsTransformationHandler - SDK transformation layer +├── prompt_injection.py # SkillPromptInjectionHandler - SKILL.md extraction and injection +├── sandbox_executor.py # SkillsSandboxExecutor - Docker sandbox code execution +├── code_execution.py # CodeExecutionHandler - automatic agentic loop +└── README.md # This file + +litellm/proxy/hooks/litellm_skills/ +├── __init__.py # Re-exports from SDK + SkillsInjectionHook +└── main.py # SkillsInjectionHook - CustomLogger hook for proxy +``` + +## Components + +### 1. `handler.py` - LiteLLMSkillsHandler + +Database operations for skills CRUD: + +```python +from litellm.llms.litellm_proxy.skills import LiteLLMSkillsHandler + +# Create skill +skill = await LiteLLMSkillsHandler.create_skill( + data=NewSkillRequest( + display_title="My Skill", + description="A helpful skill", + instructions="Use this skill when...", + file_content=zip_bytes, # ZIP file content + file_name="my-skill.zip", + file_type="application/zip", + ), + user_id="user_123" +) + +# List skills +skills = await LiteLLMSkillsHandler.list_skills(limit=10, offset=0) + +# Get skill +skill = await LiteLLMSkillsHandler.get_skill(skill_id="skill_abc123") + +# Delete skill +await LiteLLMSkillsHandler.delete_skill(skill_id="skill_abc123") +``` + +### 2. `transformation.py` - LiteLLMSkillsTransformationHandler + +SDK-level transformation layer that wraps handler operations: + +```python +from litellm.llms.litellm_proxy.skills import LiteLLMSkillsTransformationHandler + +handler = LiteLLMSkillsTransformationHandler() + +# Async create +skill = await handler.create_skill_handler( + display_title="My Skill", + files=[zip_file], + _is_async=True +) +``` + +## Skill ZIP Format + +Skills must be packaged as ZIP files with a `SKILL.md` file: + +``` +my-skill.zip +└── my-skill/ + └── SKILL.md +``` + +### SKILL.md Format + +```markdown +--- +name: my-skill +description: A brief description of what this skill does +--- + +# My Skill + +Detailed instructions for the LLM on how to use this skill. + +## Usage + +When the user asks about X, use this skill to... + +## Examples + +- Example 1: ... +- Example 2: ... +``` + +## SDK Usage + +### Create Skill in LiteLLM Database + +```python +import litellm + +# Create skill stored in LiteLLM DB +skill = litellm.create_skill( + display_title="Data Analysis Skill", + files=[open("data-analysis.zip", "rb")], + custom_llm_provider="litellm_proxy", # Store in LiteLLM DB +) + +print(f"Created skill: {skill.id}") # skill_abc123 +``` + +### Use Skill with Any Provider + +```python +import litellm + +# Use LiteLLM-stored skill with Bedrock +response = litellm.completion( + model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "Analyze this data..."}], + container={ + "skills": [ + {"type": "custom", "skill_id": "litellm:skill_abc123"} # litellm: prefix + ] + } +) +``` + +## How Skill Resolution Works + +### Step 1: Request with Skills + +```python +{ + "model": "bedrock/claude-3-sonnet", + "messages": [{"role": "user", "content": "Help me analyze data"}], + "container": { + "skills": [ + {"type": "custom", "skill_id": "litellm:skill_abc123"} + ] + } +} +``` + +### Step 2: SkillsInjectionHook Processing + +The hook (`litellm/proxy/hooks/litellm_skills/main.py`) intercepts the request: + +1. **Detects `litellm:` prefix** → Fetches skill from database +2. **Checks model provider** → Bedrock is not Anthropic +3. **Extracts SKILL.md** from stored ZIP file +4. **Converts skill to tool** + **Injects content into system prompt** + +### Step 3: Transformed Request + +```python +{ + "model": "bedrock/claude-3-sonnet", + "messages": [ + { + "role": "system", + "content": """ +--- + +# Available Skills + +## Skill: Data Analysis Skill + +# Data Analysis Skill + +This skill helps with data analysis tasks... + +## Usage +When the user asks about data analysis... +""" + }, + {"role": "user", "content": "Help me analyze data"} + ], + "tools": [ + { + "type": "function", + "function": { + "name": "skill_abc123", + "description": "This skill helps with data analysis tasks...", + "parameters": {"type": "object", "properties": {}, "required": []} + } + } + ] + # container is removed for non-Anthropic providers +} +``` + +## Database Schema + +Skills are stored in `LiteLLM_SkillsTable`: + +```prisma +model LiteLLM_SkillsTable { + skill_id String @id @default(uuid()) + display_title String? + description String? + instructions String? + source String @default("custom") + latest_version String? + metadata Json? @default("{}") + file_content Bytes? // ZIP file binary content + file_name String? // Original filename + file_type String? // MIME type + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? +} +``` + +## Routing Summary + +| Scenario | custom_llm_provider | skill_id Format | Behavior | +|----------|---------------------|-----------------|----------| +| Create skill on Anthropic | `anthropic` | N/A | Forward to Anthropic API | +| Create skill in LiteLLM DB | `litellm_proxy` | N/A | Store in database | +| Use Anthropic native skill | N/A | `skill_xyz` | Pass to Anthropic container.skills | +| Use LiteLLM skill on Anthropic | N/A | `litellm:skill_abc` | Convert to tools | +| Use LiteLLM skill on Bedrock/OpenAI | N/A | `litellm:skill_abc` | Convert to tools + inject SKILL.md | + +## Testing + +Run the tests: + +```bash +pytest tests/proxy_unit_tests/test_skills_db.py -v +``` + +Tests cover: +- Creating skills with file content +- Listing and retrieving skills +- Deleting skills +- Hook resolution with ZIP file extraction +- System prompt injection for non-Anthropic models + diff --git a/litellm/llms/litellm_proxy/skills/__init__.py b/litellm/llms/litellm_proxy/skills/__init__.py new file mode 100644 index 00000000000..5fb29e96bb9 --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/__init__.py @@ -0,0 +1,54 @@ +""" +LiteLLM Proxy Skills - Database-backed skills storage and execution + +This module provides: +- Database-backed skills storage (alternative to Anthropic's cloud-based skills API) +- Skill content extraction and prompt injection +- Sandboxed code execution for skills +- Automatic code execution handler + +Main components: +- handler.py: LiteLLMSkillsHandler - database CRUD operations +- transformation.py: LiteLLMSkillsTransformationHandler - SDK transformation layer +- prompt_injection.py: SkillPromptInjectionHandler - SKILL.md extraction and injection +- sandbox_executor.py: SkillsSandboxExecutor - Docker sandbox execution +- code_execution.py: CodeExecutionHandler - automatic agentic loop +""" + +from litellm.llms.litellm_proxy.skills.code_execution import ( + LITELLM_CODE_EXECUTION_TOOL, + CodeExecutionHandler, + LiteLLMInternalTools, + add_code_execution_tool, + code_execution_handler, + get_litellm_code_execution_tool, + has_code_execution_tool, +) +from litellm.llms.litellm_proxy.skills.constants import ( + DEFAULT_MAX_ITERATIONS, + DEFAULT_SANDBOX_TIMEOUT, +) +from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler +from litellm.llms.litellm_proxy.skills.prompt_injection import ( + SkillPromptInjectionHandler, +) +from litellm.llms.litellm_proxy.skills.sandbox_executor import SkillsSandboxExecutor +from litellm.llms.litellm_proxy.skills.transformation import ( + LiteLLMSkillsTransformationHandler, +) + +__all__ = [ + "LiteLLMSkillsHandler", + "LiteLLMSkillsTransformationHandler", + "SkillPromptInjectionHandler", + "SkillsSandboxExecutor", + "CodeExecutionHandler", + "LiteLLMInternalTools", + "LITELLM_CODE_EXECUTION_TOOL", + "get_litellm_code_execution_tool", + "code_execution_handler", + "has_code_execution_tool", + "add_code_execution_tool", + "DEFAULT_MAX_ITERATIONS", + "DEFAULT_SANDBOX_TIMEOUT", +] diff --git a/litellm/llms/litellm_proxy/skills/code_execution.py b/litellm/llms/litellm_proxy/skills/code_execution.py new file mode 100644 index 00000000000..d307b8b36d9 --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/code_execution.py @@ -0,0 +1,311 @@ +""" +Automatic Code Execution Handler for LiteLLM Skills + +When `litellm_code_execution` tool is present, this handler automatically: +1. Makes the LLM call +2. Executes any code the model generates +3. Continues the conversation with results +4. Returns final response with generated files inline (base64) + +This mimics Anthropic's behavior where code execution happens automatically. +Generated files are returned directly in the response - no separate storage needed. +""" + +import base64 +import json +from enum import Enum +from typing import Any, Dict, List, Optional + +from litellm._logging import verbose_logger + + +class LiteLLMInternalTools(str, Enum): + """ + Enum for internal LiteLLM tools that are injected into requests. + + These tools are handled automatically by LiteLLM hooks and are not + passed to the underlying LLM provider directly. + """ + CODE_EXECUTION = "litellm_code_execution" + + +def get_litellm_code_execution_tool() -> Dict[str, Any]: + """ + Returns the litellm_code_execution tool definition in OpenAI format. + + This tool enables automatic code execution in a sandboxed environment + when skills include executable Python code. + """ + return { + "type": "function", + "function": { + "name": LiteLLMInternalTools.CODE_EXECUTION.value, + "description": "Execute Python code in a sandboxed environment. Use this to run code that generates files, processes data, or performs computations. Generated files will be returned directly.", + "parameters": { + "type": "object", + "properties": { + "code": { + "type": "string", + "description": "Python code to execute" + } + }, + "required": ["code"] + } + } + } + + +def get_litellm_code_execution_tool_anthropic() -> Dict[str, Any]: + """ + Returns the litellm_code_execution tool definition in Anthropic/messages API format. + + This tool enables automatic code execution in a sandboxed environment + when skills include executable Python code. + """ + return { + "name": LiteLLMInternalTools.CODE_EXECUTION.value, + "description": "Execute Python code in a sandboxed environment. Use this to run code that generates files, processes data, or performs computations. Generated files will be returned directly.", + "input_schema": { + "type": "object", + "properties": { + "code": { + "type": "string", + "description": "Python code to execute" + } + }, + "required": ["code"] + } + } + + +# Singleton tool definition for backwards compatibility +LITELLM_CODE_EXECUTION_TOOL = get_litellm_code_execution_tool() + + +class CodeExecutionHandler: + """ + Handles automatic code execution for LiteLLM skills. + + When enabled, this handler intercepts LLM responses with code execution + tool calls, executes them in a sandbox, and continues the conversation + automatically until completion. + """ + + def __init__( + self, + max_iterations: Optional[int] = None, + sandbox_timeout: Optional[int] = None, + ): + from litellm.llms.litellm_proxy.skills.constants import ( + DEFAULT_MAX_ITERATIONS, + DEFAULT_SANDBOX_TIMEOUT, + ) + + self.max_iterations = max_iterations or DEFAULT_MAX_ITERATIONS + self.sandbox_timeout = sandbox_timeout or DEFAULT_SANDBOX_TIMEOUT + + async def execute_with_code_execution( + self, + model: str, + messages: List[Dict], + tools: List[Dict], + skill_files: Dict[str, bytes], + skill_id: Optional[str] = None, + **kwargs, + ) -> Dict[str, Any]: + """ + Execute an LLM call with automatic code execution handling. + + This method: + 1. Makes the initial LLM call + 2. If model calls litellm_code_execution, executes the code + 3. Continues conversation with results + 4. Repeats until model stops calling tools + 5. Returns final response with generated files inline + + Args: + model: Model to use + messages: Initial messages + tools: Tools including litellm_code_execution + skill_files: Dict of skill files for execution + skill_id: Optional skill ID for tracking + **kwargs: Additional args for litellm.acompletion + + Returns: + Dict with: + - response: Final LLM response + - files: List of generated files with content (base64) + - execution_results: List of code execution results + """ + import litellm + from litellm.llms.litellm_proxy.skills.sandbox_executor import ( + SkillsSandboxExecutor, + ) + + current_messages = list(messages) + generated_files: List[Dict[str, Any]] = [] # Files returned directly + execution_results: List[Dict] = [] + + executor = SkillsSandboxExecutor(timeout=self.sandbox_timeout) + response: Any = None # Initialize to avoid possibly unbound error + + for iteration in range(self.max_iterations): + verbose_logger.debug( + f"CodeExecutionHandler: Iteration {iteration + 1}/{self.max_iterations}" + ) + + # Make LLM call + response = await litellm.acompletion( + model=model, + messages=current_messages, + tools=tools, + **kwargs, + ) + + assistant_message = response.choices[0].message # type: ignore + stop_reason = response.choices[0].finish_reason # type: ignore + + # Build assistant message for conversation history + assistant_msg_dict: Dict[str, Any] = { + "role": "assistant", + "content": assistant_message.content, + } + if assistant_message.tool_calls: + assistant_msg_dict["tool_calls"] = [ + { + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + } + for tc in assistant_message.tool_calls + ] + current_messages.append(assistant_msg_dict) + + # Check if we're done (no tool calls or not tool_calls finish reason) + if stop_reason != "tool_calls" or not assistant_message.tool_calls: + verbose_logger.debug( + f"CodeExecutionHandler: Completed after {iteration + 1} iterations" + ) + return { + "response": response, + "files": generated_files, # Files returned directly with base64 content + "execution_results": execution_results, + "messages": current_messages, + } + + # Handle tool calls + for tool_call in assistant_message.tool_calls: + tool_name = tool_call.function.name + + if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value: + # Execute code in sandbox + try: + args = json.loads(tool_call.function.arguments) + code = args.get("code", "") + + verbose_logger.debug( + f"CodeExecutionHandler: Executing code ({len(code)} chars)" + ) + + exec_result = executor.execute( + code=code, + skill_files=skill_files, + ) + + verbose_logger.debug( + f"CodeExecutionHandler: Execution result: {exec_result}" + ) + + execution_results.append({ + "iteration": iteration, + "success": exec_result["success"], + "output": exec_result["output"], + "error": exec_result["error"], + "files": [f["name"] for f in exec_result["files"]], + }) + + # Build tool result content + tool_result = exec_result["output"] or "" + + # Collect generated files (returned directly, no storage) + if exec_result["files"]: + tool_result += "\n\nGenerated files:" + for f in exec_result["files"]: + file_content = base64.b64decode(f["content_base64"]) + # Add to generated files list (returned in response) + generated_files.append({ + "name": f["name"], + "mime_type": f["mime_type"], + "content_base64": f["content_base64"], + "size": len(file_content), + }) + tool_result += f"\n- {f['name']} ({len(file_content)} bytes)" + + verbose_logger.debug( + f"CodeExecutionHandler: Generated file {f['name']} ({len(file_content)} bytes)" + ) + + if exec_result["error"]: + tool_result += f"\n\nError:\n{exec_result['error']}" + + except Exception as e: + tool_result = f"Code execution failed: {str(e)}" + execution_results.append({ + "iteration": iteration, + "success": False, + "error": str(e), + }) + + # Add tool result to messages + current_messages.append({ + "role": "tool", + "tool_call_id": tool_call.id, + "content": tool_result, + }) + else: + # Non-code-execution tool - pass through + # In a full implementation, this would call other tool handlers + current_messages.append({ + "role": "tool", + "tool_call_id": tool_call.id, + "content": f"Tool '{tool_name}' not handled by code execution handler", + }) + + # Max iterations reached + verbose_logger.warning( + f"CodeExecutionHandler: Max iterations ({self.max_iterations}) reached" + ) + return { + "response": response, + "files": generated_files, + "execution_results": execution_results, + "messages": current_messages, + "max_iterations_reached": True, + } + + +def has_code_execution_tool(tools: Optional[List[Dict]]) -> bool: + """Check if litellm_code_execution tool is in the tools list.""" + if not tools: + return False + for tool in tools: + func = tool.get("function", {}) + if func.get("name") == LiteLLMInternalTools.CODE_EXECUTION.value: + return True + return False + + +def add_code_execution_tool(tools: Optional[List[Dict]]) -> List[Dict]: + """Add litellm_code_execution tool if not already present.""" + tools = tools or [] + if not has_code_execution_tool(tools): + tools.append(LITELLM_CODE_EXECUTION_TOOL) + return tools + + +# Global handler instance +code_execution_handler = CodeExecutionHandler() + diff --git a/litellm/llms/litellm_proxy/skills/constants.py b/litellm/llms/litellm_proxy/skills/constants.py new file mode 100644 index 00000000000..a2be6961db6 --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/constants.py @@ -0,0 +1,13 @@ +""" +Constants for LiteLLM Skills + +Centralized constants for skills processing, code execution, and sandbox configuration. +""" + +# Code execution loop settings +DEFAULT_MAX_ITERATIONS: int = 10 +"""Maximum number of iterations for the automatic code execution loop.""" + +DEFAULT_SANDBOX_TIMEOUT: int = 120 +"""Default timeout in seconds for sandbox code execution.""" + diff --git a/litellm/llms/litellm_proxy/skills/handler.py b/litellm/llms/litellm_proxy/skills/handler.py new file mode 100644 index 00000000000..f44ac4cda92 --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/handler.py @@ -0,0 +1,219 @@ +""" +Handler for LiteLLM database-backed skills operations. + +This module contains the actual database operations for skills CRUD. +Used by the transformation layer and skills injection hook. +""" + +import uuid +from typing import Any, Dict, List, Optional + +from litellm._logging import verbose_logger +from litellm.proxy._types import LiteLLM_SkillsTable, NewSkillRequest + + +def _prisma_skill_to_litellm(prisma_skill) -> LiteLLM_SkillsTable: + """ + Convert a Prisma skill record to LiteLLM_SkillsTable. + + Handles Base64 decoding of file_content field. + """ + import base64 + + data = prisma_skill.model_dump() + + # Decode Base64 file_content back to bytes + # model_dump() converts Base64 field to base64-encoded string + if data.get("file_content") is not None: + if isinstance(data["file_content"], str): + data["file_content"] = base64.b64decode(data["file_content"]) + elif isinstance(data["file_content"], bytes): + # Already bytes, no conversion needed + pass + + return LiteLLM_SkillsTable(**data) + + +class LiteLLMSkillsHandler: + """ + Handler for LiteLLM database-backed skills operations. + + This class provides static methods for CRUD operations on skills + stored in the LiteLLM proxy database (LiteLLM_SkillsTable). + """ + + @staticmethod + async def _get_prisma_client(): + """Get the prisma client from proxy server.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise ValueError( + "Prisma client is not initialized. " + "Database connection required for LiteLLM skills." + ) + return prisma_client + + @staticmethod + async def create_skill( + data: NewSkillRequest, + user_id: Optional[str] = None, + ) -> LiteLLM_SkillsTable: + """ + Create a new skill in the LiteLLM database. + + Args: + data: NewSkillRequest with skill details + user_id: Optional user ID for tracking + + Returns: + LiteLLM_SkillsTable record + """ + prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + + skill_id = f"litellm_skill_{uuid.uuid4()}" + + skill_data: Dict[str, Any] = { + "skill_id": skill_id, + "display_title": data.display_title, + "description": data.description, + "instructions": data.instructions, + "source": "custom", + "created_by": user_id, + "updated_by": user_id, + } + + # Handle metadata + if data.metadata is not None: + from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + + skill_data["metadata"] = safe_dumps(data.metadata) + + # Handle file content - wrap bytes in Base64 for Prisma + if data.file_content is not None: + from prisma.fields import Base64 + + skill_data["file_content"] = Base64.encode(data.file_content) + if data.file_name is not None: + skill_data["file_name"] = data.file_name + if data.file_type is not None: + skill_data["file_type"] = data.file_type + + verbose_logger.debug( + f"LiteLLMSkillsHandler: Creating skill {skill_id} with title={data.display_title}" + ) + + new_skill = await prisma_client.db.litellm_skillstable.create(data=skill_data) + + return _prisma_skill_to_litellm(new_skill) + + @staticmethod + async def list_skills( + limit: int = 20, + offset: int = 0, + ) -> List[LiteLLM_SkillsTable]: + """ + List skills from the LiteLLM database. + + Args: + limit: Maximum number of skills to return + offset: Number of skills to skip + + Returns: + List of LiteLLM_SkillsTable records + """ + prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + + verbose_logger.debug( + f"LiteLLMSkillsHandler: Listing skills with limit={limit}, offset={offset}" + ) + + skills = await prisma_client.db.litellm_skillstable.find_many( + take=limit, + skip=offset, + order={"created_at": "desc"}, + ) + + return [_prisma_skill_to_litellm(s) for s in skills] + + @staticmethod + async def get_skill(skill_id: str) -> LiteLLM_SkillsTable: + """ + Get a skill by ID from the LiteLLM database. + + Args: + skill_id: The skill ID to retrieve + + Returns: + LiteLLM_SkillsTable record + + Raises: + ValueError: If skill not found + """ + prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + + verbose_logger.debug(f"LiteLLMSkillsHandler: Getting skill {skill_id}") + + skill = await prisma_client.db.litellm_skillstable.find_unique( + where={"skill_id": skill_id} + ) + + if skill is None: + raise ValueError(f"Skill not found: {skill_id}") + + return _prisma_skill_to_litellm(skill) + + @staticmethod + async def delete_skill(skill_id: str) -> Dict[str, str]: + """ + Delete a skill by ID from the LiteLLM database. + + Args: + skill_id: The skill ID to delete + + Returns: + Dict with id and type of deleted skill + + Raises: + ValueError: If skill not found + """ + prisma_client = await LiteLLMSkillsHandler._get_prisma_client() + + verbose_logger.debug(f"LiteLLMSkillsHandler: Deleting skill {skill_id}") + + # Check if skill exists + skill = await prisma_client.db.litellm_skillstable.find_unique( + where={"skill_id": skill_id} + ) + + if skill is None: + raise ValueError(f"Skill not found: {skill_id}") + + # Delete the skill + await prisma_client.db.litellm_skillstable.delete(where={"skill_id": skill_id}) + + return {"id": skill_id, "type": "skill_deleted"} + + @staticmethod + async def fetch_skill_from_db(skill_id: str) -> Optional[LiteLLM_SkillsTable]: + """ + Fetch a skill from the database (used by skills injection hook). + + This is a convenience method that returns None instead of raising + an exception if the skill is not found. + + Args: + skill_id: The skill ID to fetch + + Returns: + LiteLLM_SkillsTable or None if not found + """ + try: + return await LiteLLMSkillsHandler.get_skill(skill_id) + except ValueError: + return None + except Exception as e: + verbose_logger.warning( + f"LiteLLMSkillsHandler: Error fetching skill {skill_id}: {e}" + ) + return None diff --git a/litellm/llms/litellm_proxy/skills/prompt_injection.py b/litellm/llms/litellm_proxy/skills/prompt_injection.py new file mode 100644 index 00000000000..17469274c1c --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/prompt_injection.py @@ -0,0 +1,305 @@ +""" +Prompt Injection Handler for LiteLLM Skills + +Handles extraction of skill content (SKILL.md) from stored ZIP files +and injection into the system prompt for non-Anthropic models. +""" + +import zipfile +from io import BytesIO +from typing import Any, Dict, List, Optional + +from litellm._logging import verbose_logger +from litellm.proxy._types import LiteLLM_SkillsTable + + +class SkillPromptInjectionHandler: + """ + Handles skill content extraction and system prompt injection. + + Responsibilities: + - Extract SKILL.md content from skill ZIP files + - Extract ALL files from ZIP for code execution + - Inject skill content into system message + - Create execute_code tool definition + """ + + def extract_skill_content(self, skill: LiteLLM_SkillsTable) -> Optional[str]: + """ + Extract skill content from the stored zip file. + + Looks for SKILL.md or README.md in the zip and returns its content. + This content describes the skill's capabilities and instructions. + + Args: + skill: The skill from LiteLLM database + + Returns: + The skill content as a string, or None if not available + """ + if not skill.file_content: + return skill.instructions + + try: + zip_buffer = BytesIO(skill.file_content) + with zipfile.ZipFile(zip_buffer, "r") as zf: + # Look for SKILL.md first + for name in zf.namelist(): + if name.endswith("SKILL.md"): + content = zf.read(name).decode("utf-8") + if content: + return f"## Skill: {skill.display_title or skill.skill_id}\n\n{content}" + + # Fall back to README.md + for name in zf.namelist(): + if name.endswith("README.md"): + content = zf.read(name).decode("utf-8") + if content: + return f"## Skill: {skill.display_title or skill.skill_id}\n\n{content}" + + # Fall back to any .md file + for name in zf.namelist(): + if name.endswith(".md"): + content = zf.read(name).decode("utf-8") + if content: + return f"## Skill: {skill.display_title or skill.skill_id}\n\n{content}" + except Exception as e: + verbose_logger.warning( + f"SkillPromptInjectionHandler: Error extracting content from skill {skill.skill_id}: {e}" + ) + + return skill.instructions + + def extract_all_files(self, skill: LiteLLM_SkillsTable) -> Dict[str, bytes]: + """ + Extract ALL files from skill ZIP for code execution. + + Returns a dict mapping file paths to their binary content. + The paths have the skill folder prefix removed (e.g., "slack-gif-creator/core/..." -> "core/..."). + + Args: + skill: The skill from LiteLLM database + + Returns: + Dict mapping file paths to binary content + """ + files: Dict[str, bytes] = {} + + if not skill.file_content: + return files + + try: + zip_buffer = BytesIO(skill.file_content) + with zipfile.ZipFile(zip_buffer, "r") as zf: + for name in zf.namelist(): + # Skip directories + if name.endswith("/"): + continue + + # Remove skill folder prefix (first path component) + parts = name.split("/") + if len(parts) > 1: + clean_path = "/".join(parts[1:]) + else: + clean_path = name + + if clean_path: + files[clean_path] = zf.read(name) + except Exception as e: + verbose_logger.warning( + f"SkillPromptInjectionHandler: Error extracting files from skill {skill.skill_id}: {e}" + ) + + return files + + def inject_skill_content_to_messages( + self, data: dict, skill_contents: List[str], use_anthropic_format: bool = False + ) -> dict: + """ + Inject skill content into the system prompt. + + For Anthropic messages API (use_anthropic_format=True): + - Injects into top-level 'system' parameter (not in messages array) + + For OpenAI-style APIs (use_anthropic_format=False): + - Injects into messages array with role="system" + + Args: + data: The request data dict + skill_contents: List of skill content strings to inject + use_anthropic_format: If True, use top-level 'system' param for Anthropic + + Returns: + Modified data dict with skill content in system prompt + """ + if not skill_contents: + return data + + # Build the skill injection text + skill_section = "\n\n---\n\n# Available Skills\n\n" + "\n\n---\n\n".join(skill_contents) + + if use_anthropic_format: + # Anthropic messages API: use top-level 'system' parameter + current_system = data.get("system", "") + if current_system: + data["system"] = current_system + skill_section + else: + data["system"] = skill_section.strip() + return data + + # OpenAI-style: inject into messages array + messages = data.get("messages", []) + if not messages: + return data + + # Find or create system message + system_msg_idx = None + for i, msg in enumerate(messages): + if isinstance(msg, dict) and msg.get("role") == "system": + system_msg_idx = i + break + + if system_msg_idx is not None: + # Append to existing system message + current_content = messages[system_msg_idx].get("content", "") + messages[system_msg_idx]["content"] = current_content + skill_section + else: + # Create new system message at the beginning + messages.insert(0, {"role": "system", "content": skill_section.strip()}) + + data["messages"] = messages + return data + + def create_execute_code_tool(self, skill_modules: List[str]) -> Dict[str, Any]: + """ + Create the execute_code tool definition. + + This tool allows the model to execute Python code with access + to the skill's modules (e.g., 'from core.gif_builder import GIFBuilder'). + + Args: + skill_modules: List of available module paths (e.g., ["core/gif_builder.py"]) + + Returns: + OpenAI-style tool definition + """ + # Format module list for description + module_examples = [] + for mod in skill_modules[:5]: # Limit to 5 examples + if mod.endswith(".py"): + # Convert path to import: "core/gif_builder.py" -> "from core.gif_builder import ..." + import_path = mod.replace("/", ".").replace(".py", "") + module_examples.append(f"from {import_path} import ...") + + module_hint = "" + if module_examples: + module_hint = f" Available modules: {', '.join(module_examples)}" + + return { + "type": "function", + "function": { + "name": "execute_code", + "description": f"Execute Python code in a sandboxed environment. Generated files will be returned.{module_hint}", + "parameters": { + "type": "object", + "properties": { + "code": { + "type": "string", + "description": "Python code to execute. You can import skill modules and use standard libraries." + } + }, + "required": ["code"] + } + } + } + + def convert_skill_to_tool(self, skill: LiteLLM_SkillsTable) -> Dict[str, Any]: + """ + Convert a LiteLLM skill to an OpenAI-style tool. + + The skill's instructions are used as the function description, + allowing the model to understand when and how to use the skill. + + Args: + skill: The skill from LiteLLM database + + Returns: + OpenAI-style tool definition + """ + # Create a function name from skill_id (sanitize for function naming) + func_name = skill.skill_id.replace("-", "_").replace(" ", "_") + + # Use instructions as description, fall back to description or title + description = ( + skill.instructions + or skill.description + or skill.display_title + or f"Skill: {skill.skill_id}" + ) + + # Truncate description if too long (OpenAI has limits) + max_desc_length = 1024 + if len(description) > max_desc_length: + description = description[: max_desc_length - 3] + "..." + + tool: Dict[str, Any] = { + "type": "function", + "function": { + "name": func_name, + "description": description, + "parameters": { + "type": "object", + "properties": {}, + "required": [], + }, + }, + } + + # If skill has metadata with parameter definitions, use them + if skill.metadata and isinstance(skill.metadata, dict): + params = skill.metadata.get("parameters") + if params and isinstance(params, dict): + tool["function"]["parameters"] = params + + return tool + + def convert_skill_to_anthropic_tool(self, skill: LiteLLM_SkillsTable) -> Dict[str, Any]: + """ + Convert a LiteLLM skill to an Anthropic-style tool (messages API format). + + Args: + skill: The skill from LiteLLM database + + Returns: + Anthropic-style tool definition with name, description, input_schema + """ + func_name = skill.skill_id.replace("-", "_").replace(" ", "_") + + description = ( + skill.instructions + or skill.description + or skill.display_title + or f"Skill: {skill.skill_id}" + ) + + max_desc_length = 1024 + if len(description) > max_desc_length: + description = description[: max_desc_length - 3] + "..." + + input_schema: Dict[str, Any] = { + "type": "object", + "properties": {}, + "required": [], + } + + if skill.metadata and isinstance(skill.metadata, dict): + params = skill.metadata.get("parameters") + if params and isinstance(params, dict): + input_schema = params + + return { + "name": func_name, + "description": description, + "input_schema": input_schema, + } + diff --git a/litellm/llms/litellm_proxy/skills/sandbox_executor.py b/litellm/llms/litellm_proxy/skills/sandbox_executor.py new file mode 100644 index 00000000000..7676ade5cd0 --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/sandbox_executor.py @@ -0,0 +1,286 @@ +""" +Sandbox Executor for LiteLLM Skills + +Executes skill code in a sandboxed environment using llm-sandbox. +Supports Docker, Podman, and Kubernetes backends. +""" + +import base64 +import os +from typing import Any, Dict, List, Optional + +from litellm._logging import verbose_logger + + +class SkillsSandboxExecutor: + """ + Executes skill code in llm-sandbox Docker container. + + Responsibilities: + - Create sandbox session with skill files + - Install requirements + - Execute model-generated code + - Collect generated files (GIFs, images, etc.) + """ + + def __init__( + self, + timeout: int = 60, + backend: str = "docker", + image: Optional[str] = None, + ): + """ + Initialize the sandbox executor. + + Args: + timeout: Maximum execution time in seconds + backend: Sandbox backend ("docker", "podman", "kubernetes") + image: Custom Docker image (default: uses llm-sandbox default) + """ + self.timeout = timeout + self.backend = backend + self.image = image + self._session = None + + def execute( + self, + code: str, + skill_files: Dict[str, bytes], + requirements: Optional[str] = None, + ) -> Dict[str, Any]: + """ + Execute code with skill files in sandbox. + + Args: + code: Python code to execute + skill_files: Dict mapping file paths to binary content + requirements: Optional requirements.txt content + + Returns: + { + "success": bool, + "output": str, + "error": str (if failed), + "files": [{"name": str, "content_base64": str, "mime_type": str}] + } + """ + try: + from llm_sandbox import SandboxSession + except ImportError: + verbose_logger.error( + "SkillsSandboxExecutor: llm-sandbox not installed. " + "Install with: pip install llm-sandbox" + ) + return { + "success": False, + "output": "", + "error": "llm-sandbox not installed. Install with: pip install llm-sandbox", + "files": [], + } + + try: + # Create sandbox session + session_kwargs: Dict[str, Any] = { + "lang": "python", + "verbose": False, + } + + if self.image: + session_kwargs["image"] = self.image + + with SandboxSession(**session_kwargs) as session: + # 1. Copy skill files into sandbox using copy_to_runtime + import tempfile + + # Create a temp directory to stage files + with tempfile.TemporaryDirectory() as tmpdir: + for path, content in skill_files.items(): + # Create the file in temp directory + local_path = os.path.join(tmpdir, path) + os.makedirs(os.path.dirname(local_path), exist_ok=True) + with open(local_path, "wb") as f: + f.write(content) + + # Copy to sandbox + sandbox_path = f"/sandbox/{path}" + session.copy_to_runtime(local_path, sandbox_path) + + verbose_logger.debug( + f"SkillsSandboxExecutor: Copied {len(skill_files)} files to sandbox" + ) + + # 2. Install requirements if present + req_packages = None + if requirements: + req_packages = requirements.strip().replace("\n", " ") + elif "requirements.txt" in skill_files: + req_content = skill_files["requirements.txt"].decode("utf-8") + req_packages = req_content.strip().replace("\n", " ") + + if req_packages: + # Run pip install as code + pip_code = f""" +import subprocess +subprocess.run(['pip', 'install'] + '{req_packages}'.split(), check=True) +""" + result = session.run(pip_code) + verbose_logger.debug( + "SkillsSandboxExecutor: Installed requirements" + ) + + # 3. Execute the code + # Wrap code to run from /sandbox directory + wrapped_code = f""" +import os +os.chdir('/sandbox') +import sys +sys.path.insert(0, '/sandbox') + +{code} +""" + result = session.run(wrapped_code) + + success = result.exit_code == 0 + output = result.stdout or "" + error = result.stderr or "" + + if success: + verbose_logger.debug( + "SkillsSandboxExecutor: Code execution succeeded" + ) + else: + verbose_logger.debug( + f"SkillsSandboxExecutor: Code execution failed with exit code {result.exit_code}" + ) + verbose_logger.debug( + f"SkillsSandboxExecutor: stderr: {error[:500] if error else 'No stderr'}" + ) + verbose_logger.debug( + f"SkillsSandboxExecutor: stdout: {output[:500] if output else 'No stdout'}" + ) + + # 4. Collect generated files + generated_files = self._collect_generated_files(session, skill_files) + + return { + "success": success, + "output": output, + "error": error, + "files": generated_files, + } + + except Exception as e: + verbose_logger.error( + f"SkillsSandboxExecutor: Execution failed: {e}" + ) + return { + "success": False, + "output": "", + "error": str(e), + "files": [], + } + + def _collect_generated_files( + self, + session: Any, + original_files: Dict[str, bytes], + ) -> List[Dict[str, Any]]: + """ + Collect files generated during execution. + + Looks for new files in /sandbox that weren't in the original skill files. + Focuses on common output types: GIF, PNG, JPG, PDF, CSV, etc. + + Args: + session: The sandbox session + original_files: Original skill files (to exclude) + + Returns: + List of generated files with base64 content + """ + generated_files: List[Dict[str, Any]] = [] + + try: + import tempfile + + # List files in /sandbox using Python code + list_code = """ +import os +import json +files = [] +for root, dirs, filenames in os.walk('/sandbox'): + for f in filenames: + if f.endswith(('.gif', '.png', '.jpg', '.jpeg', '.pdf', '.csv', '.json')): + files.append(os.path.join(root, f)) +print(json.dumps(files)) +""" + result = session.run(list_code) + + if result.exit_code == 0 and result.stdout: + import json + try: + filepaths = json.loads(result.stdout.strip()) + except json.JSONDecodeError: + filepaths = [] + + for filepath in filepaths: + if not filepath: + continue + + # Get relative path + rel_path = filepath.replace("/sandbox/", "") + + # Skip if it was an original file + if rel_path in original_files: + continue + + # Copy file from sandbox using copy_from_runtime + with tempfile.NamedTemporaryFile(delete=False) as tmp: + tmp_path = tmp.name + + try: + session.copy_from_runtime(filepath, tmp_path) + + with open(tmp_path, "rb") as f: + content = f.read() + + content_b64 = base64.b64encode(content).decode("utf-8") + generated_files.append({ + "name": os.path.basename(filepath), + "path": rel_path, + "content_base64": content_b64, + "mime_type": self._get_mime_type(filepath), + }) + + verbose_logger.debug( + f"SkillsSandboxExecutor: Collected generated file: {rel_path}" + ) + except Exception as e: + verbose_logger.warning( + f"SkillsSandboxExecutor: Error copying file {filepath}: {e}" + ) + finally: + if os.path.exists(tmp_path): + os.unlink(tmp_path) + + except Exception as e: + verbose_logger.warning( + f"SkillsSandboxExecutor: Error collecting generated files: {e}" + ) + + return generated_files + + def _get_mime_type(self, filename: str) -> str: + """Get MIME type for a file based on extension.""" + ext = filename.lower().split(".")[-1] + return { + "gif": "image/gif", + "png": "image/png", + "jpg": "image/jpeg", + "jpeg": "image/jpeg", + "pdf": "application/pdf", + "csv": "text/csv", + "json": "application/json", + "txt": "text/plain", + }.get(ext, "application/octet-stream") + diff --git a/litellm/llms/litellm_proxy/skills/transformation.py b/litellm/llms/litellm_proxy/skills/transformation.py new file mode 100644 index 00000000000..e7c999eacec --- /dev/null +++ b/litellm/llms/litellm_proxy/skills/transformation.py @@ -0,0 +1,336 @@ +""" +Transformation handler for LiteLLM database-backed skills. + +This module provides the SDK-level transformation layer that converts +API requests to database operations via LiteLLMSkillsHandler. + +Pattern follows litellm/llms/litellm_proxy/responses/transformation.py +""" + +from typing import TYPE_CHECKING, Any, Coroutine, Dict, List, Optional, Union + +from litellm.types.llms.anthropic_skills import ( + DeleteSkillResponse, + ListSkillsResponse, + Skill, +) +from litellm.types.utils import LlmProviders + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + +class LiteLLMSkillsTransformationHandler: + """ + Transformation handler for skills API requests to LiteLLM database operations. + + This is used when custom_llm_provider="litellm_proxy" to store/retrieve skills + from the LiteLLM proxy database instead of calling an external API. + """ + + @property + def custom_llm_provider(self) -> str: + """Return the provider name for logging.""" + return LlmProviders.LITELLM_PROXY.value + + def create_skill_handler( + self, + display_title: Optional[str] = None, + description: Optional[str] = None, + instructions: Optional[str] = None, + files: Optional[List[Any]] = None, + file_content: Optional[bytes] = None, + file_name: Optional[str] = None, + file_type: Optional[str] = None, + metadata: Optional[Dict[str, Any]] = None, + user_id: Optional[str] = None, + _is_async: bool = False, + logging_obj: Optional["LiteLLMLoggingObj"] = None, + litellm_call_id: Optional[str] = None, + **kwargs, + ) -> Union[Skill, Coroutine[Any, Any, Skill]]: + """ + Create a skill in LiteLLM database. + + Args: + display_title: Display title for the skill + description: Description of the skill + instructions: Instructions/prompt for the skill + files: Files to upload - list of tuples (filename, content, content_type) + file_content: Binary content of skill files (alternative to files) + file_name: Original filename (alternative to files) + file_type: MIME type (alternative to files) + metadata: Additional metadata + user_id: User ID for tracking + _is_async: Whether to return a coroutine + + Returns: + Skill object or coroutine that returns Skill + """ + # Pre-call logging + if logging_obj: + logging_obj.update_environment_variables( + model=None, + optional_params={"display_title": display_title}, + litellm_params={"litellm_call_id": litellm_call_id}, + custom_llm_provider=self.custom_llm_provider, + ) + + # Extract file content from files parameter if provided + # files is a list of tuples: [(filename, content, content_type), ...] + if files and not file_content: + if isinstance(files, list) and len(files) > 0: + first_file = files[0] + if isinstance(first_file, tuple) and len(first_file) >= 2: + file_name = first_file[0] + file_content = first_file[1] + file_type = first_file[2] if len(first_file) > 2 else "application/zip" + + if _is_async: + return self._async_create_skill( + display_title=display_title, + description=description, + instructions=instructions, + file_content=file_content, + file_name=file_name, + file_type=file_type, + metadata=metadata, + user_id=user_id, + ) + + import asyncio + return asyncio.get_event_loop().run_until_complete( + self._async_create_skill( + display_title=display_title, + description=description, + instructions=instructions, + file_content=file_content, + file_name=file_name, + file_type=file_type, + metadata=metadata, + user_id=user_id, + ) + ) + + async def _async_create_skill( + self, + display_title: Optional[str] = None, + description: Optional[str] = None, + instructions: Optional[str] = None, + file_content: Optional[bytes] = None, + file_name: Optional[str] = None, + file_type: Optional[str] = None, + metadata: Optional[Dict[str, Any]] = None, + user_id: Optional[str] = None, + ) -> Skill: + """Async implementation of create_skill.""" + # Lazy import to avoid SDK dependency on proxy + from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler + from litellm.proxy._types import NewSkillRequest + + skill_request = NewSkillRequest( + display_title=display_title, + description=description, + instructions=instructions, + file_content=file_content, + file_name=file_name, + file_type=file_type, + metadata=metadata, + ) + + db_skill = await LiteLLMSkillsHandler.create_skill( + data=skill_request, + user_id=user_id, + ) + + return self._db_skill_to_response(db_skill) + + def list_skills_handler( + self, + limit: int = 20, + offset: int = 0, + _is_async: bool = False, + logging_obj: Optional["LiteLLMLoggingObj"] = None, + litellm_call_id: Optional[str] = None, + **kwargs, + ) -> Union[ListSkillsResponse, Coroutine[Any, Any, ListSkillsResponse]]: + """ + List skills from LiteLLM database. + + Args: + limit: Maximum number of skills to return + offset: Number of skills to skip + _is_async: Whether to return a coroutine + logging_obj: LiteLLM logging object + litellm_call_id: Call ID for logging + + Returns: + ListSkillsResponse or coroutine that returns ListSkillsResponse + """ + # Pre-call logging + if logging_obj: + logging_obj.update_environment_variables( + model=None, + optional_params={"limit": limit, "offset": offset}, + litellm_params={"litellm_call_id": litellm_call_id}, + custom_llm_provider=self.custom_llm_provider, + ) + + if _is_async: + return self._async_list_skills(limit=limit, offset=offset) + + import asyncio + return asyncio.get_event_loop().run_until_complete( + self._async_list_skills(limit=limit, offset=offset) + ) + + async def _async_list_skills( + self, + limit: int = 20, + offset: int = 0, + ) -> ListSkillsResponse: + """Async implementation of list_skills.""" + # Lazy import to avoid SDK dependency on proxy + from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler + + db_skills = await LiteLLMSkillsHandler.list_skills( + limit=limit, + offset=offset, + ) + + skills = [self._db_skill_to_response(s) for s in db_skills] + return ListSkillsResponse( + data=skills, + has_more=len(skills) >= limit, + next_page=None, + ) + + def get_skill_handler( + self, + skill_id: str, + _is_async: bool = False, + logging_obj: Optional["LiteLLMLoggingObj"] = None, + litellm_call_id: Optional[str] = None, + **kwargs, + ) -> Union[Skill, Coroutine[Any, Any, Skill]]: + """ + Get a skill from LiteLLM database. + + Args: + skill_id: The skill ID to retrieve + _is_async: Whether to return a coroutine + logging_obj: LiteLLM logging object + litellm_call_id: Call ID for logging + + Returns: + Skill or coroutine that returns Skill + """ + # Pre-call logging + if logging_obj: + logging_obj.update_environment_variables( + model=None, + optional_params={"skill_id": skill_id}, + litellm_params={"litellm_call_id": litellm_call_id}, + custom_llm_provider=self.custom_llm_provider, + ) + + if _is_async: + return self._async_get_skill(skill_id=skill_id) + + import asyncio + return asyncio.get_event_loop().run_until_complete( + self._async_get_skill(skill_id=skill_id) + ) + + async def _async_get_skill(self, skill_id: str) -> Skill: + """Async implementation of get_skill.""" + # Lazy import to avoid SDK dependency on proxy + from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler + + db_skill = await LiteLLMSkillsHandler.get_skill(skill_id=skill_id) + return self._db_skill_to_response(db_skill) + + def delete_skill_handler( + self, + skill_id: str, + _is_async: bool = False, + logging_obj: Optional["LiteLLMLoggingObj"] = None, + litellm_call_id: Optional[str] = None, + **kwargs, + ) -> Union[DeleteSkillResponse, Coroutine[Any, Any, DeleteSkillResponse]]: + """ + Delete a skill from LiteLLM database. + + Args: + skill_id: The skill ID to delete + _is_async: Whether to return a coroutine + logging_obj: LiteLLM logging object + litellm_call_id: Call ID for logging + + Returns: + DeleteSkillResponse or coroutine that returns DeleteSkillResponse + """ + # Pre-call logging + if logging_obj: + logging_obj.update_environment_variables( + model=None, + optional_params={"skill_id": skill_id}, + litellm_params={"litellm_call_id": litellm_call_id}, + custom_llm_provider=self.custom_llm_provider, + ) + + if _is_async: + return self._async_delete_skill(skill_id=skill_id) + + import asyncio + return asyncio.get_event_loop().run_until_complete( + self._async_delete_skill(skill_id=skill_id) + ) + + async def _async_delete_skill(self, skill_id: str) -> DeleteSkillResponse: + """Async implementation of delete_skill.""" + # Lazy import to avoid SDK dependency on proxy + from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler + + result = await LiteLLMSkillsHandler.delete_skill(skill_id=skill_id) + return DeleteSkillResponse( + id=result["id"], + type=result.get("type", "skill_deleted"), + ) + + def _db_skill_to_response(self, db_skill: Any) -> Skill: + """ + Convert a database skill record to Anthropic-compatible Skill response. + + Args: + db_skill: LiteLLM_SkillsTable record + + Returns: + Skill object + """ + created_at = "" + updated_at = "" + + if hasattr(db_skill, "created_at") and db_skill.created_at: + created_at = ( + db_skill.created_at.isoformat() + if hasattr(db_skill.created_at, "isoformat") + else str(db_skill.created_at) + ) + if hasattr(db_skill, "updated_at") and db_skill.updated_at: + updated_at = ( + db_skill.updated_at.isoformat() + if hasattr(db_skill.updated_at, "isoformat") + else str(db_skill.updated_at) + ) + + return Skill( + id=db_skill.skill_id, + created_at=created_at, + updated_at=updated_at, + display_title=db_skill.display_title, + latest_version=db_skill.latest_version, + source=db_skill.source or "custom", + type="skill", + ) + diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 7ccec074703..96598c1dfe6 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -96,8 +96,8 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): validated_input.append(item.model_dump(exclude_none=True)) elif isinstance(item, dict): # Handle reasoning items specifically to filter out status=None - verbose_logger.debug(f"Handling reasoning item: {item}") if item.get("type") == "reasoning": + verbose_logger.debug(f"Handling reasoning item: {item}") # Type assertion since we know it's a dict at this point dict_item = cast(Dict[str, Any], item) filtered_item = self._handle_reasoning_item(dict_item) @@ -411,7 +411,6 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): ) raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers diff --git a/litellm/llms/stability/image_edit/__init__.py b/litellm/llms/stability/image_edit/__init__.py new file mode 100644 index 00000000000..5a9eb2e02b9 --- /dev/null +++ b/litellm/llms/stability/image_edit/__init__.py @@ -0,0 +1,37 @@ +""" +Stability AI Image Edit Module + +Factory function for getting the appropriate config class. +""" + +from litellm.llms.base_llm.image_edit.transformation import ( + BaseImageEditConfig, +) + +from .transformations import StabilityImageEditConfig + +__all__ = [ + "StabilityImageEditConfig", + "get_stability_image_edit_config", +] + + +def get_stability_image_edit_config(model: str) -> BaseImageEditConfig: + """ + Get the appropriate Stability AI config for the given model. + + Currently all models use the same config class, but this factory + allows for model-specific configs in the future. + + Args: + model: The model name (e.g., "stability/inpaint", "stability/outpaint") + + Returns: + BaseImageEditConfig instance for Stability AI + """ + # For now, all models use the same config + # In the future, we could have model-specific configs: + # - StabilityInpaintConfig for Inpaint models + # - StabilityOutpaintConfig for Outpaint models + # - etc. + return StabilityImageEditConfig() diff --git a/litellm/llms/stability/image_edit/transformations.py b/litellm/llms/stability/image_edit/transformations.py new file mode 100644 index 00000000000..173fae2d6fd --- /dev/null +++ b/litellm/llms/stability/image_edit/transformations.py @@ -0,0 +1,314 @@ +""" +Stability AI Image Edit Config + +Handles transformation between OpenAI-compatible format and Stability AI API format. + +API Reference: https://platform.stability.ai/docs/api-reference +""" + +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +import httpx +from httpx._types import RequestFiles + +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams +from litellm.types.llms.stability import ( + OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO, + STABILITY_EDIT_ENDPOINTS, +) +from litellm.types.utils import FileTypes, ImageObject, ImageResponse +from litellm.utils import get_model_info + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class StabilityImageEditConfig(BaseImageEditConfig): + """ + Configuration for Stability AI image edit. + + Supports: + - Stable Diffusion 3 (SD3, SD3.5) Image Edit + """ + + DEFAULT_BASE_URL: str = "https://api.stability.ai" + + def get_supported_openai_params( + self, model: str + ) -> List[str]: + """ + Return list of OpenAI params supported by Stability AI. + + https://platform.stability.ai/docs/api-reference + """ + return [ + "n", # Number of images (Stability always returns 1, we can loop) + "size", # Maps to aspect_ratio + "response_format", # b64_json or url (Stability only returns b64) + "mask" + ] + + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """ + Map OpenAI parameters to Stability AI parameters. + + OpenAI -> Stability mappings: + - size -> aspect_ratio + - n -> (handled separately, Stability returns 1 image per request) + """ + supported_params = self.get_supported_openai_params(model) + # Define mapping from OpenAI params to Stability params + param_mapping = { + "size": "aspect_ratio", + # "n" and "response_format" are handled separately + } + + # Create a copy to not mutate original - convert TypedDict to regular dict + mapped_params: Dict[str, Any] = dict(image_edit_optional_params) + + for k, v in image_edit_optional_params.items(): + if k in param_mapping: + # Map param if mapping exists and value is valid + if k == "size" and v in OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO: + mapped_params[param_mapping[k]] = OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO[v] # type: ignore + # Don't copy "size" itself to final dict + elif k == "n": + # Store for logic but do not add to outgoing params + mapped_params["_n"] = v + elif k == "response_format": + # Only b64 supported at Stability; store for postprocessing + mapped_params["_response_format"] = v + elif k not in supported_params: + if not drop_params: + raise ValueError( + f"Parameter {k} is not supported for model {model}. " + f"Supported parameters are {supported_params}. " + f"Set drop_params=True to drop unsupported parameters." + ) + # Otherwise, param will simply be dropped + else: + # param is supported and not mapped, keep as-is + continue + + # Remove OpenAI params that have been mapped unless they're in stability + for mapped in ["size", "n", "response_format"]: + if mapped in mapped_params: + del mapped_params[mapped] + + return mapped_params + + def _get_model_endpoint(self, model: str) -> str: + """ + Get the API endpoint for a given model. + """ + # Remove "stability/" prefix if present + model_name = model.lower() + if model_name.startswith("stability/"): + model_name = model_name[10:] # Remove "stability/" prefix + + # Check if model is in our mapping + for key, endpoint in STABILITY_EDIT_ENDPOINTS.items(): + if key in model_name: + return endpoint + + # Default to SD3 endpoint + return "/v2beta/stable-image/edit/inpaint" + + def get_complete_url( + self, + model: str, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Get the complete URL for the Stability AI API request. + """ + base_url: str = ( + api_base + or get_secret_str("STABILITY_API_BASE") + or litellm_params.get("api_base", None) + or self.DEFAULT_BASE_URL + ) + base_url = base_url.rstrip("/") + + endpoint = self._get_model_endpoint(model) + return f"{base_url}{endpoint}" + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + ) -> dict: + """ + Validate environment and set up headers for Stability AI. + """ + final_api_key: Optional[str] = api_key or get_secret_str("STABILITY_API_KEY") + + if not final_api_key: + raise ValueError( + "STABILITY_API_KEY is not set. " + "Please set it via environment variable or pass api_key parameter." + ) + + headers["Authorization"] = f"Bearer {final_api_key}" + headers["Accept"] = "application/json" + return headers + + def transform_image_edit_request( + self, + model: str, + prompt: str, + image: FileTypes, + image_edit_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[Dict, RequestFiles]: + """ + Transform OpenAI-style request to Stability AI request format. + + Note: Stability AI uses multipart/form-data, but the HTTP handler + will handle the conversion from dict to form data. + """ + # Build Stability request + # Populate multipart form-data as separate text fields (data) and files. + # Stability expects prompt/output_format/etc. as normal form fields, not file parts. + data: Dict[str, Any] = { + "prompt": prompt, + "output_format": "png", # Default to PNG + } + # Handle image parameter - could be a single file or list + image_file = image[0] if isinstance(image, list) else image # type: ignore + files: Dict[str, Any] = {"image": image_file} + + # Add optional params (already mapped in map_openai_params) + for key, value in image_edit_optional_request_params.items(): # type: ignore + # Skip internal params (prefixed with _) + if key.startswith("_") or value is None: + continue + + # File-like optional param + if key == "mask": + # Handle case where mask might be in a list + mask_value = value + if isinstance(value, list) and len(value) > 0: + mask_value = value[0] + files["mask"] = mask_value # type: ignore + continue + + # File-like optional params (init_image, style_image, etc.) + if key in ["init_image", "style_image"]: + # Handle case where value might be in a list + file_value = value + if isinstance(value, list) and len(value) > 0: + file_value = value[0] + files[key] = file_value # type: ignore + continue + + # Supported text fields + if key in [ + "negative_prompt", + "aspect_ratio", + "seed", + "mode", + "strength", + "style_preset", + "left", + "bottom", + "right", + "top", + "creativity", + "search_prompt", + "grow_mask", + "select_prompt", + "control_strength", + "composition_fidelity", + "change_strength" + ]: + data[key] = value # type: ignore + + return data, files + + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> ImageResponse: + """ + Transform Stability AI response to OpenAI-compatible ImageResponse. + + Stability returns: {"image": "base64...", "finish_reason": "SUCCESS", "seed": 123} + OpenAI expects: {"data": [{"b64_json": "base64..."}], "created": timestamp} + """ + try: + response_data = raw_response.json() + except Exception as e: + raise self.get_error_class( + error_message=f"Error parsing Stability AI response: {e}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check for errors in response + if "errors" in response_data: + raise self.get_error_class( + error_message=f"Stability AI error: {response_data['errors']}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check finish_reason + finish_reason = response_data.get("finish_reason", "") + if finish_reason == "CONTENT_FILTERED": + raise self.get_error_class( + error_message="Content was filtered by Stability AI safety systems", + status_code=400, + headers=raw_response.headers, + ) + + model_response = ImageResponse() + if not model_response.data: + model_response.data = [] + + # Extract image from response + image_b64 = response_data.get("image") + if image_b64: + model_response.data.append( + ImageObject( + b64_json=image_b64, + url=None, + revised_prompt=None, + ) + ) + + if not hasattr(model_response, "_hidden_params"): + model_response._hidden_params = {} + if "additional_headers" not in model_response._hidden_params: + model_response._hidden_params["additional_headers"] = {} + # Override: fetch model-cost from model_cost map based on the provided model name + model_info = get_model_info(model, custom_llm_provider="stability") + cost_per_image = model_info.get("output_cost_per_image", 0) + if cost_per_image is not None: + model_response._hidden_params["additional_headers"]["llm_provider-x-litellm-response-cost"] = float(cost_per_image) + return model_response + + def use_multipart_form_data(self) -> bool: + """ + Stability AI requires multipart/form-data for image generation. + """ + return True diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 6bb11430f20..03fa5b98928 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -640,14 +640,28 @@ def add_object_type(schema): if properties is not None: if "required" in schema and schema["required"] is None: schema.pop("required", None) - schema["type"] = "object" - for name, value in properties.items(): - add_object_type(value) + # Gemini doesn't accept empty properties for object types + # If properties is empty, remove it and the type field + if not properties: + schema.pop("properties", None) + schema.pop("type", None) + schema.pop("required", None) + else: + schema["type"] = "object" + for name, value in properties.items(): + add_object_type(value) items = schema.get("items", None) if items is not None: add_object_type(items) + for key in ["anyOf", "oneOf", "allOf"]: + values = schema.get(key, None) + if values is not None and isinstance(values, list): + for value in values: + if isinstance(value, dict): + add_object_type(value) + def strip_field(schema, field_name: str): schema.pop(field_name, None) diff --git a/litellm/llms/vertex_ai/ocr/common_utils.py b/litellm/llms/vertex_ai/ocr/common_utils.py new file mode 100644 index 00000000000..dc2c07420bf --- /dev/null +++ b/litellm/llms/vertex_ai/ocr/common_utils.py @@ -0,0 +1,41 @@ +""" +Common utilities for Vertex AI OCR providers. + +This module provides routing logic to determine which OCR configuration to use +based on the model name. +""" + +from typing import TYPE_CHECKING, Optional + +if TYPE_CHECKING: + from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig + + +def get_vertex_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]: + """ + Determine which Vertex AI OCR configuration to use based on the model name. + + Vertex AI supports multiple OCR services: + - Vertex AI OCR: vertex_ai/ + + Args: + model: The model name (e.g., "vertex_ai/ocr/") + + Returns: + OCR configuration instance for the specified model + + Examples: + >>> get_vertex_ai_ocr_config("vertex_ai/deepseek-ai/deepseek-ocr-maas") + + + >>> get_vertex_ai_ocr_config("vertex_ai/ocr/mistral-ocr-maas") + + """ + from litellm.llms.vertex_ai.ocr.deepseek_transformation import ( + VertexAIDeepSeekOCRConfig, + ) + from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig + if "deepseek" in model: + return VertexAIDeepSeekOCRConfig() + return VertexAIOCRConfig() + diff --git a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py new file mode 100644 index 00000000000..b16f73af3f6 --- /dev/null +++ b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py @@ -0,0 +1,394 @@ +""" +Vertex AI DeepSeek OCR transformation implementation. +""" +import json +from typing import TYPE_CHECKING, Any, Dict, Optional + +import httpx + +from litellm._logging import verbose_logger +from litellm.llms.base_llm.ocr.transformation import ( + BaseOCRConfig, + DocumentType, + OCRPage, + OCRRequestData, + OCRResponse, + OCRUsageInfo, +) +from litellm.llms.vertex_ai.vertex_llm_base import VertexBase + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class VertexAIDeepSeekOCRConfig(BaseOCRConfig): + """ + Vertex AI DeepSeek OCR transformation configuration. + + Vertex AI DeepSeek OCR uses the chat completion API format through the openapi endpoint. + This transformation converts OCR requests to chat completion format and vice versa. + """ + + def __init__(self) -> None: + super().__init__() + self.vertex_base = VertexBase() + + def validate_environment( + self, + headers: Dict, + model: str, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + litellm_params: Optional[dict] = None, + **kwargs, + ) -> Dict: + """ + Validate environment and return headers for Vertex AI OCR. + + Vertex AI uses Bearer token authentication with access token from credentials. + """ + # Extract Vertex AI parameters using safe helpers from VertexBase + # Use safe_get_* methods that don't mutate litellm_params dict + litellm_params = litellm_params or {} + + vertex_project = VertexBase.safe_get_vertex_ai_project(litellm_params=litellm_params) + vertex_credentials = VertexBase.safe_get_vertex_ai_credentials(litellm_params=litellm_params) + + # Get access token from Vertex credentials + access_token, project_id = self.vertex_base.get_access_token( + credentials=vertex_credentials, + project_id=vertex_project, + ) + + headers = { + "Authorization": f"Bearer {access_token}", + "Content-Type": "application/json", + **headers, + } + + return headers + + def get_complete_url( + self, + api_base: Optional[str], + model: str, + optional_params: dict, + litellm_params: Optional[dict] = None, + **kwargs, + ) -> str: + """ + Get complete URL for Vertex AI DeepSeek OCR endpoint. + + Vertex AI endpoint format: + https://{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/endpoints/openapi/chat/completions + + Args: + api_base: Vertex AI API base URL (optional) + model: Model name (e.g., "deepseek-ai/deepseek-ocr-maas") + optional_params: Optional parameters + litellm_params: LiteLLM parameters containing vertex_project, vertex_location + + Returns: Complete URL for Vertex AI OCR endpoint + """ + # Extract Vertex AI parameters using safe helpers from VertexBase + # Use safe_get_* methods that don't mutate litellm_params dict + litellm_params = litellm_params or {} + + vertex_project = VertexBase.safe_get_vertex_ai_project(litellm_params=litellm_params) + vertex_location = VertexBase.safe_get_vertex_ai_location(litellm_params=litellm_params) + + if vertex_project is None: + raise ValueError( + "Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter" + ) + + if vertex_location is None: + vertex_location = "us-central1" + + # Get API base URL + if api_base is None: + api_base = "https://aiplatform.googleapis.com" + + # Ensure no trailing slash + api_base = api_base.rstrip("/") + + # Vertex AI DeepSeek OCR endpoint format + # Format: https://{region}-aiplatform.googleapis.com/v1/projects/{project}/locations/{region}/endpoints/openapi/chat/completions + return f"{api_base}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi/chat/completions" + + def transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Transform OCR request to chat completion format for Vertex AI DeepSeek OCR. + + Converts OCR document format to chat completion messages format: + - Input: {"type": "image_url", "image_url": "gs://..."} + - Output: {"model": "deepseek-ai/deepseek-ocr-maas", "messages": [{"role": "user", "content": [{"type": "image_url", "image_url": "gs://..."}]}]} + + Args: + model: Model name (e.g., "deepseek-ai/deepseek-ocr-maas") + document: Document dict from user (Mistral OCR format) + optional_params: Already mapped optional parameters + headers: Request headers + **kwargs: Additional arguments + + Returns: + OCRRequestData with JSON data in chat completion format + """ + verbose_logger.debug("Vertex AI DeepSeek OCR transform_ocr_request (sync) called") + + if not isinstance(document, dict): + raise ValueError(f"Expected document dict, got {type(document)}") + + # Extract document type and URL + doc_type = document.get("type") + image_url = None + document_url = None + + if doc_type == "image_url": + image_url = document.get("image_url", "") + elif doc_type == "document_url": + document_url = document.get("document_url", "") + else: + raise ValueError(f"Unsupported document type: {doc_type}. Expected 'image_url' or 'document_url'") + + # Build chat completion message content + content_item = {} + if image_url: + content_item = { + "type": "image_url", + "image_url": image_url + } + elif document_url: + # For document URLs, we use image_url type as well (Vertex AI supports both) + content_item = { + "type": "image_url", + "image_url": document_url + } + + # Build chat completion request + data = { + "model": "deepseek-ai/" + model, + "messages": [ + { + "role": "user", + "content": [content_item] + } + ] + } + + # Add optional parameters (stream, temperature, etc.) + # Filter out OCR-specific params that don't apply to chat completion + chat_completion_params = {} + for key, value in optional_params.items(): + # Include common chat completion params + if key in ["stream", "temperature", "max_tokens", "top_p", "n", "stop"]: + chat_completion_params[key] = value + + data.update(chat_completion_params) + + verbose_logger.debug("Vertex AI DeepSeek OCR: Transformed request to chat completion format") + + return OCRRequestData(data=data, files=None) + + async def async_transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Transform OCR request to chat completion format for Vertex AI DeepSeek OCR (async). + + Same as sync version - no async-specific logic needed. + + Args: + model: Model name + document: Document dict from user + optional_params: Already mapped optional parameters + headers: Request headers + **kwargs: Additional arguments + + Returns: + OCRRequestData with JSON data in chat completion format + """ + return self.transform_ocr_request( + model=model, + document=document, + optional_params=optional_params, + headers=headers, + **kwargs, + ) + + def transform_ocr_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs, + ) -> OCRResponse: + """ + Transform chat completion response to OCR format. + + Vertex AI DeepSeek OCR returns chat completion format: + { + "id": "...", + "object": "chat.completion", + "choices": [{ + "message": { + "role": "assistant", + "content": "" + } + }], + "usage": {...} + } + + We need to extract the content and convert it to OCRResponse format. + + Args: + model: Model name + raw_response: Raw HTTP response from Vertex AI + logging_obj: Logging object + **kwargs: Additional arguments + + Returns: + OCRResponse in standard format + """ + verbose_logger.debug("Vertex AI DeepSeek OCR transform_ocr_response called") + verbose_logger.debug(f"Raw response: {raw_response.text}") + + try: + response_json = raw_response.json() + + # Extract content from chat completion response + choices = response_json.get("choices", []) + if not choices: + raise ValueError("No choices in chat completion response") + + message = choices[0].get("message", {}) + content = message.get("content", "") + + if not content: + raise ValueError("No content in chat completion response") + + # Try to parse content as JSON (OCR result might be JSON string) + ocr_data = None + try: + # If content is a JSON string, parse it + if isinstance(content, str) and content.strip().startswith("{"): + ocr_data = json.loads(content) + elif isinstance(content, dict): + ocr_data = content + else: + # If content is markdown text, create a single page with the markdown + ocr_data = { + "pages": [ + { + "index": 0, + "markdown": content + } + ], + "model": model, + "usage_info": response_json.get("usage", {}) + } + except json.JSONDecodeError: + # If JSON parsing fails, treat content as markdown + ocr_data = { + "pages": [ + { + "index": 0, + "markdown": content + } + ], + "model": model, + "usage_info": response_json.get("usage", {}) + } + + # Ensure we have the expected structure + if "pages" not in ocr_data: + # If OCR data doesn't have pages, wrap the content in a page + ocr_data = { + "pages": [ + { + "index": 0, + "markdown": content if isinstance(content, str) else json.dumps(content) + } + ], + "model": ocr_data.get("model", model), + "usage_info": ocr_data.get("usage_info", response_json.get("usage", {})) + } + + # Convert usage info if present + usage_info = None + if "usage_info" in ocr_data: + usage_dict = ocr_data["usage_info"] + if isinstance(usage_dict, dict): + usage_info = OCRUsageInfo(**usage_dict) + + # Build OCRResponse + pages = [] + for page_data in ocr_data.get("pages", []): + # Ensure page has required fields + if isinstance(page_data, dict): + page = OCRPage( + index=page_data.get("index", 0), + markdown=page_data.get("markdown", ""), + images=page_data.get("images"), + dimensions=page_data.get("dimensions") + ) + pages.append(page) + + if not pages: + # Create a default page if none exist + pages = [OCRPage(index=0, markdown=content if isinstance(content, str) else "")] + + return OCRResponse( + pages=pages, + model=ocr_data.get("model", model), + document_annotation=ocr_data.get("document_annotation"), + usage_info=usage_info, + object="ocr", + ) + + except Exception as e: + verbose_logger.error(f"Error parsing Vertex AI DeepSeek OCR response: {e}") + raise e + + async def async_transform_ocr_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs, + ) -> OCRResponse: + """ + Async transform chat completion response to OCR format. + + Same as sync version - no async-specific logic needed. + + Args: + model: Model name + raw_response: Raw HTTP response + logging_obj: Logging object + **kwargs: Additional arguments + + Returns: + OCRResponse in standard format + """ + return self.transform_ocr_response( + model=model, + raw_response=raw_response, + logging_obj=logging_obj, + **kwargs, + ) + diff --git a/litellm/main.py b/litellm/main.py index b46208cc432..0715dd8e61b 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -165,7 +165,8 @@ from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion from .llms.azure_ai.embed import AzureAIEmbedding from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM from .llms.bedrock.embed.embedding import BedrockEmbedding -from .llms.bedrock.image.image_handler import BedrockImageGeneration +from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration +from .llms.bedrock.image_edit.handler import BedrockImageEdit from .llms.bytez.chat.transformation import BytezChatConfig from .llms.clarifai.chat.transformation import ClarifaiConfig from .llms.codestral.completion.handler import CodestralTextCompletion @@ -271,6 +272,7 @@ codestral_text_completions = CodestralTextCompletion() bedrock_converse_chat_completion = BedrockConverseLLM() bedrock_embedding = BedrockEmbedding() bedrock_image_generation = BedrockImageGeneration() +bedrock_image_edit = BedrockImageEdit() vertex_chat_completion = VertexLLM() vertex_embedding = VertexEmbedding() vertex_multimodal_embedding = VertexMultimodalEmbedding() diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 0f5f61e708d..d0bbbe6d5df 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -24483,6 +24483,90 @@ "output_cost_per_image": 0.08, "supported_endpoints": ["/v1/images/generations"] }, + "stability/inpaint": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/outpaint": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.004, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/erase": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/search-and-replace": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/search-and-recolor": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/remove-background": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/replace-background-and-relight": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.008, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/sketch": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/structure": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/style": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/style-transfer": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.008, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/fast": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.002, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/conservative": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.04, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/creative": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.06, + "supported_endpoints": ["/v1/images/edits"] + }, "stability/stable-image-core": { "litellm_provider": "stability", "mode": "image_generation", @@ -24531,6 +24615,84 @@ "mode": "image_generation", "output_cost_per_image": 0.14 }, + "stability.stable-conservative-upscale-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.40 + }, + "stability.stable-creative-upscale-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.60 + }, + "stability.stable-fast-upscale-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.03 + }, + "stability.stable-outpaint-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.06 + }, + "stability.stable-image-control-sketch-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-control-structure-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-erase-object-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-inpaint-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-remove-background-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-search-recolor-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-search-replace-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-style-guide-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-style-transfer-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.08 + }, "standard/1024-x-1024/dall-e-3": { "input_cost_per_pixel": 3.81469e-08, "litellm_provider": "openai", @@ -27777,6 +27939,14 @@ ], "source": "https://cloud.google.com/generative-ai-app-builder/pricing" }, + "vertex_ai/deepseek-ai/deepseek-ocr-maas": { + "litellm_provider": "vertex_ai", + "mode": "ocr", + "input_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, + "ocr_cost_per_page": 3e-04, + "source": "https://cloud.google.com/vertex-ai/pricing" + }, "vertex_ai/openai/gpt-oss-120b-maas": { "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-openai_models", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 744de60c5a3..3172f866177 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -16,7 +16,11 @@ from typing_extensions import Required, TypedDict from litellm._uuid import uuid from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import AllMessageValues, OpenAIFileObject +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIFileObject, + ResponsesAPIResponse, +) from litellm.types.mcp import ( MCPAuth, MCPAuthType, @@ -1140,6 +1144,60 @@ class MakeMCPServersPublicRequest(LiteLLMPydanticObjectBase): mcp_server_ids: List[str] +######## Skills API Types ######## + + +class NewSkillRequest(LiteLLMPydanticObjectBase): + """Request to create a new skill in LiteLLM database""" + + display_title: Optional[str] = None + description: Optional[str] = None + instructions: Optional[str] = None + file_content: Optional[bytes] = None # Binary content of skill files (zip) + file_name: Optional[str] = None # Original filename + file_type: Optional[str] = None # MIME type (e.g., "application/zip") + metadata: Optional[Dict[str, Any]] = None + + +class UpdateSkillRequest(LiteLLMPydanticObjectBase): + """Request to update an existing skill""" + + skill_id: str + display_title: Optional[str] = None + description: Optional[str] = None + instructions: Optional[str] = None + file_content: Optional[bytes] = None # Binary content of skill files (zip) + file_name: Optional[str] = None # Original filename + file_type: Optional[str] = None # MIME type + metadata: Optional[Dict[str, Any]] = None + + +class LiteLLM_SkillsTable(LiteLLMPydanticObjectBase): + """Represents a LiteLLM_SkillsTable record""" + + skill_id: str + display_title: Optional[str] = None + description: Optional[str] = None + instructions: Optional[str] = None + source: str = "custom" + latest_version: Optional[str] = None + file_content: Optional[bytes] = None # Binary content of skill files (zip) + file_name: Optional[str] = None # Original filename + file_type: Optional[str] = None # MIME type + metadata: Optional[Dict[str, Any]] = None + created_at: Optional[datetime] = None + created_by: Optional[str] = None + updated_at: Optional[datetime] = None + updated_by: Optional[str] = None + + +class ListSkillsRequest(LiteLLMPydanticObjectBase): + """Request to list skills from LiteLLM database""" + + limit: Optional[int] = 20 + offset: Optional[int] = 0 + + class NewUserRequestTeam(LiteLLMPydanticObjectBase): team_id: str max_budget_in_team: Optional[float] = None @@ -3740,8 +3798,8 @@ class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase): unified_object_id: str model_object_id: str - file_purpose: Literal["batch", "fine-tune"] - file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob] + file_purpose: Literal["batch", "fine-tune", "response"] + file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, ResponsesAPIResponse] class EnterpriseLicenseData(TypedDict, total=False): diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 8637bc88c57..f798d218f1d 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -885,14 +885,16 @@ class ProxyBaseLLMRequestProcessing: @staticmethod def _get_pre_call_type( - route_type: Literal["acompletion", "aembedding", "aresponses"], - ) -> Literal["completion", "embeddings", "responses"]: + route_type: Literal["acompletion", "aembedding", "aresponses", "allm_passthrough_route"], + ) -> Literal["completion", "embeddings", "responses", "allm_passthrough_route"]: if route_type == "acompletion": return "completion" elif route_type == "aembedding": return "embeddings" elif route_type == "aresponses": return "responses" + elif route_type == "allm_passthrough_route": + return "allm_passthrough_route" ######################################################### # Proxy Level Streaming Data Generator diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 0dac30f72b2..a1bbf36ac0c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -128,7 +128,7 @@ class UnifiedLLMGuardrails(CustomLogger): endpoint_guardrail_translation_mappings = ( load_guardrail_translation_mappings() ) - if CallTypes(call_type) not in endpoint_guardrail_translation_mappings: + if call_type is not None and CallTypes(call_type) not in endpoint_guardrail_translation_mappings: return data endpoint_translation = endpoint_guardrail_translation_mappings[ @@ -180,10 +180,10 @@ class UnifiedLLMGuardrails(CustomLogger): call_type: Optional[CallTypesLiteral] = None if user_api_key_dict.request_route is not None: call_types = get_call_types_for_route(user_api_key_dict.request_route) - if call_types is not None and len(call_types) > 0: - call_type = call_types[0] + if call_types is not None and len(call_types) > 0: # type: ignore + call_type = call_types[0] # type: ignore if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=response) + call_type = _infer_call_type(call_type=None, completion_response=response) # type: ignore if call_type is None: return response @@ -308,10 +308,10 @@ class UnifiedLLMGuardrails(CustomLogger): if call_type is None and user_api_key_dict.request_route is not None: call_types = get_call_types_for_route(user_api_key_dict.request_route) if call_types is not None: - call_type = call_types[0] + call_type = call_types[0].value if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=item) + call_type = _infer_call_type(call_type=None, completion_response=item) # type: ignore # If call type not supported, just pass through all chunks if ( diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index ccb1d0c7bd7..1d1e559d4be 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -3,6 +3,7 @@ from typing import Literal, Union from . import * from .cache_control_check import _PROXY_CacheControlCheck +from .litellm_skills import SkillsInjectionHook from .max_budget_limiter import _PROXY_MaxBudgetLimiter from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 @@ -21,6 +22,7 @@ PROXY_HOOKS = { "parallel_request_limiter": _PROXY_MaxParallelRequestsHandler_v3, "cache_control_check": _PROXY_CacheControlCheck, "responses_id_security": ResponsesIDSecurity, + "litellm_skills": SkillsInjectionHook, } ## FEATURE FLAG HOOKS ## diff --git a/litellm/proxy/hooks/litellm_skills/__init__.py b/litellm/proxy/hooks/litellm_skills/__init__.py new file mode 100644 index 00000000000..057cf3d8b38 --- /dev/null +++ b/litellm/proxy/hooks/litellm_skills/__init__.py @@ -0,0 +1,39 @@ +""" +LiteLLM Skills Hook - Proxy integration for skills + +This module provides the CustomLogger hook for skills processing. +The actual skill logic is in litellm/llms/litellm_proxy/skills/. + +Usage: + from litellm.proxy.hooks.litellm_skills import SkillsInjectionHook + + # Register hook in proxy + litellm.callbacks.append(SkillsInjectionHook()) +""" + +# Re-export from the SDK location for convenience +from litellm.llms.litellm_proxy.skills import ( + LITELLM_CODE_EXECUTION_TOOL, + CodeExecutionHandler, + LiteLLMInternalTools, + SkillPromptInjectionHandler, + SkillsSandboxExecutor, + code_execution_handler, + get_litellm_code_execution_tool, +) +from litellm.proxy.hooks.litellm_skills.main import ( + SkillsInjectionHook, + skills_injection_hook, +) + +__all__ = [ + "SkillsInjectionHook", + "skills_injection_hook", + "CodeExecutionHandler", + "LiteLLMInternalTools", + "LITELLM_CODE_EXECUTION_TOOL", + "get_litellm_code_execution_tool", + "code_execution_handler", + "SkillPromptInjectionHandler", + "SkillsSandboxExecutor", +] diff --git a/litellm/proxy/hooks/litellm_skills/main.py b/litellm/proxy/hooks/litellm_skills/main.py new file mode 100644 index 00000000000..26d4cbe1de7 --- /dev/null +++ b/litellm/proxy/hooks/litellm_skills/main.py @@ -0,0 +1,869 @@ +""" +Skills Injection Hook for LiteLLM Proxy + +Main hook that orchestrates skill processing: +- Fetches skills from LiteLLM DB +- Injects SKILL.md content into system prompt +- Adds litellm_code_execution tool for automatic code execution +- Handles agentic loop internally when litellm_code_execution is called + +For non-Anthropic models (e.g., Bedrock, OpenAI, etc.): +- Skills are converted to OpenAI-style tools +- Skill file content (SKILL.md) is extracted and injected into the system prompt +- litellm_code_execution tool is added - when model calls it, LiteLLM handles + execution automatically and returns final response with file_ids + +Usage: + # Simple - LiteLLM handles everything automatically via proxy + # The container parameter triggers the SkillsInjectionHook + response = await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Create a bouncing ball GIF"}], + container={"skills": [{"skill_id": "litellm:skill_abc123"}]}, + ) + # Response includes file_ids for generated files +""" + +import base64 +import json +from typing import Any, Dict, List, Optional, Union + +from litellm._logging import verbose_proxy_logger +from litellm.caching.caching import DualCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.litellm_proxy.skills.prompt_injection import ( + SkillPromptInjectionHandler, +) +from litellm.proxy._types import LiteLLM_SkillsTable, UserAPIKeyAuth +from litellm.types.utils import CallTypes, CallTypesLiteral + + +class SkillsInjectionHook(CustomLogger): + """ + Pre/Post-call hook that processes skills from container.skills parameter. + + Pre-call (async_pre_call_hook): + - Skills with 'litellm:' prefix are fetched from LiteLLM DB + - For Anthropic models: native skills pass through, LiteLLM skills converted to tools + - For non-Anthropic models: LiteLLM skills are converted to tools + execute_code tool + + Post-call (async_post_call_success_deployment_hook): + - If response has litellm_code_execution tool call, automatically execute code + - Continue conversation loop until model gives final response + - Return response with generated files inline + + This hook is called automatically by litellm during completion calls. + """ + + def __init__(self, **kwargs): + from litellm.llms.litellm_proxy.skills.constants import ( + DEFAULT_MAX_ITERATIONS, + DEFAULT_SANDBOX_TIMEOUT, + ) + + self.optional_params = kwargs + self.prompt_handler = SkillPromptInjectionHandler() + self.max_iterations = kwargs.get("max_iterations", DEFAULT_MAX_ITERATIONS) + self.sandbox_timeout = kwargs.get("sandbox_timeout", DEFAULT_SANDBOX_TIMEOUT) + super().__init__(**kwargs) + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: CallTypesLiteral, + ) -> Optional[Union[Exception, str, dict]]: + """ + Process skills from container.skills before the LLM call. + + 1. Check if container.skills exists in request + 2. Separate skills by prefix (litellm: vs native) + 3. Fetch LiteLLM skills from database + 4. For Anthropic: keep native skills in container + 5. For non-Anthropic: convert LiteLLM skills to tools, inject content, add execute_code + """ + # Only process completion-type calls + if call_type not in ["completion", "acompletion", "anthropic_messages"]: + return data + + container = data.get("container") + if not container or not isinstance(container, dict): + return data + + skills = container.get("skills") + if not skills or not isinstance(skills, list): + return data + + verbose_proxy_logger.debug(f"SkillsInjectionHook: Processing {len(skills)} skills") + + litellm_skills: List[LiteLLM_SkillsTable] = [] + anthropic_skills: List[Dict[str, Any]] = [] + + # Separate skills by prefix + for skill in skills: + if not isinstance(skill, dict): + continue + + skill_id = skill.get("skill_id", "") + if skill_id.startswith("litellm_"): + # Fetch from LiteLLM DB + db_skill = await self._fetch_skill_from_db(skill_id) + if db_skill: + litellm_skills.append(db_skill) + else: + verbose_proxy_logger.warning( + f"SkillsInjectionHook: Skill '{skill_id}' not found in LiteLLM DB" + ) + else: + # Native Anthropic skill - pass through + anthropic_skills.append(skill) + + # Check if using messages API spec (anthropic_messages call type) + # Messages API always uses Anthropic-style tool format + use_anthropic_format = call_type == "anthropic_messages" + + if len(litellm_skills) > 0: + data = self._process_for_messages_api( + data=data, + litellm_skills=litellm_skills, + use_anthropic_format=use_anthropic_format, + ) + + return data + + + def _process_for_messages_api( + self, + data: dict, + litellm_skills: List[LiteLLM_SkillsTable], + use_anthropic_format: bool = True, + ) -> dict: + """ + Process skills for messages API (Anthropic format tools). + + - Converts skills to Anthropic-style tools (name, description, input_schema) + - Extracts and injects SKILL.md content into system prompt + - Adds litellm_code_execution tool for code execution + - Stores skill files in metadata for sandbox execution + """ + from litellm.llms.litellm_proxy.skills.code_execution import ( + get_litellm_code_execution_tool_anthropic, + ) + + tools = data.get("tools", []) + skill_contents: List[str] = [] + all_skill_files: Dict[str, Dict[str, bytes]] = {} + all_module_paths: List[str] = [] + + for skill in litellm_skills: + # Convert skill to Anthropic-style tool + tools.append(self.prompt_handler.convert_skill_to_anthropic_tool(skill)) + + # Extract skill content from file if available + content = self.prompt_handler.extract_skill_content(skill) + if content: + skill_contents.append(content) + + # Extract all files for code execution + skill_files = self.prompt_handler.extract_all_files(skill) + if skill_files: + all_skill_files[skill.skill_id] = skill_files + for path in skill_files.keys(): + if path.endswith(".py"): + all_module_paths.append(path) + + if tools: + data["tools"] = tools + + # Inject skill content into system prompt + # For Anthropic messages API, use top-level 'system' param instead of messages array + if skill_contents: + data = self.prompt_handler.inject_skill_content_to_messages( + data, skill_contents, use_anthropic_format=use_anthropic_format + ) + + # Add litellm_code_execution tool if we have skill files + if all_skill_files: + code_exec_tool = get_litellm_code_execution_tool_anthropic() + data["tools"] = data.get("tools", []) + [code_exec_tool] + + # Store skill files in litellm_metadata for automatic code execution + data["litellm_metadata"] = data.get("litellm_metadata", {}) + data["litellm_metadata"]["_skill_files"] = all_skill_files + data["litellm_metadata"]["_litellm_code_execution_enabled"] = True + + # Remove container (not supported by underlying providers) + data.pop("container", None) + + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Messages API - converted {len(litellm_skills)} skills to Anthropic tools, " + f"injected {len(skill_contents)} skill contents, " + f"added litellm_code_execution tool with {len(all_module_paths)} modules" + ) + + return data + + def _process_non_anthropic_model( + self, + data: dict, + litellm_skills: List[LiteLLM_SkillsTable], + ) -> dict: + """ + Process skills for non-Anthropic models (OpenAI format tools). + + - Converts skills to OpenAI-style tools + - Extracts and injects SKILL.md content + - Adds execute_code tool for code execution + - Stores skill files in metadata for sandbox execution + """ + tools = data.get("tools", []) + skill_contents: List[str] = [] + all_skill_files: Dict[str, Dict[str, bytes]] = {} + all_module_paths: List[str] = [] + + for skill in litellm_skills: + # Convert skill to OpenAI-style tool + tools.append(self.prompt_handler.convert_skill_to_tool(skill)) + + # Extract skill content from file if available + content = self.prompt_handler.extract_skill_content(skill) + if content: + skill_contents.append(content) + + # Extract all files for code execution + skill_files = self.prompt_handler.extract_all_files(skill) + if skill_files: + all_skill_files[skill.skill_id] = skill_files + # Collect Python module paths + for path in skill_files.keys(): + if path.endswith(".py"): + all_module_paths.append(path) + + if tools: + data["tools"] = tools + + # Inject skill content into system prompt + if skill_contents: + data = self.prompt_handler.inject_skill_content_to_messages(data, skill_contents) + + # Add litellm_code_execution tool if we have skill files + if all_skill_files: + from litellm.llms.litellm_proxy.skills.code_execution import ( + get_litellm_code_execution_tool, + ) + data["tools"] = data.get("tools", []) + [get_litellm_code_execution_tool()] + + # Store skill files in litellm_metadata for automatic code execution + # Using litellm_metadata instead of metadata to avoid conflicts with user metadata + data["litellm_metadata"] = data.get("litellm_metadata", {}) + data["litellm_metadata"]["_skill_files"] = all_skill_files + data["litellm_metadata"]["_litellm_code_execution_enabled"] = True + + # Remove container for non-Anthropic (they don't support it) + data.pop("container", None) + + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Non-Anthropic model - converted {len(litellm_skills)} skills to tools, " + f"injected {len(skill_contents)} skill contents, " + f"added execute_code tool with {len(all_module_paths)} modules" + ) + + return data + + async def _fetch_skill_from_db(self, skill_id: str) -> Optional[LiteLLM_SkillsTable]: + """ + Fetch a skill from the LiteLLM database. + + Args: + skill_id: The skill ID (without 'litellm:' prefix) + + Returns: + LiteLLM_SkillsTable or None if not found + """ + try: + from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler + + return await LiteLLMSkillsHandler.fetch_skill_from_db(skill_id) + except Exception as e: + verbose_proxy_logger.warning( + f"SkillsInjectionHook: Error fetching skill {skill_id}: {e}" + ) + return None + + def _is_anthropic_model(self, model: str) -> bool: + """ + Check if the model is an Anthropic model using get_llm_provider. + + Args: + model: The model name/identifier + + Returns: + True if Anthropic model, False otherwise + """ + try: + from litellm.litellm_core_utils.get_llm_provider_logic import ( + get_llm_provider, + ) + + _, custom_llm_provider, _, _ = get_llm_provider(model=model) + return custom_llm_provider == "anthropic" + except Exception: + # Fallback to simple check if get_llm_provider fails + return "claude" in model.lower() or model.lower().startswith("anthropic/") + + async def async_post_call_success_deployment_hook( + self, + request_data: dict, + response: Any, + call_type: Optional[CallTypes], + ) -> Optional[Any]: + """ + Post-call hook to handle automatic code execution. + + Handles both OpenAI format (response.choices) and Anthropic/messages API + format (response["content"]). + + If the response contains a tool call (litellm_code_execution or skill tool): + 1. Execute the code in sandbox + 2. Add result to messages + 3. Make another LLM call + 4. Repeat until model gives final response + 5. Return modified response with generated files + """ + from litellm.llms.litellm_proxy.skills.code_execution import ( + LiteLLMInternalTools, + ) + + # Check if code execution is enabled for this request + litellm_metadata = request_data.get("litellm_metadata", {}) + metadata = request_data.get("metadata", {}) + + code_exec_enabled = ( + litellm_metadata.get("_litellm_code_execution_enabled") or + metadata.get("_litellm_code_execution_enabled") + ) + if not code_exec_enabled: + return None + + # Get skill files + skill_files_by_id = ( + litellm_metadata.get("_skill_files") or + metadata.get("_skill_files", {}) + ) + all_skill_files: Dict[str, bytes] = {} + for files_dict in skill_files_by_id.values(): + all_skill_files.update(files_dict) + + if not all_skill_files: + verbose_proxy_logger.warning( + "SkillsInjectionHook: No skill files found, cannot execute code" + ) + return None + + # Check for tool calls - handle both Anthropic and OpenAI formats + tool_calls = self._extract_tool_calls(response) + if not tool_calls: + return None + + # Check if any tool call needs execution (litellm_code_execution or skill tool) + has_executable_tool = False + for tc in tool_calls: + tool_name = tc.get("name", "") + # Execute if it's litellm_code_execution OR a skill tool (skill_xxx) + if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value or tool_name.startswith("skill_"): + has_executable_tool = True + break + + if not has_executable_tool: + return None + + verbose_proxy_logger.debug( + "SkillsInjectionHook: Detected tool call, starting execution loop" + ) + + # Start the agentic loop + return await self._execute_code_loop_messages_api( + data=request_data, + response=response, + skill_files=all_skill_files, + ) + + def _extract_tool_calls(self, response: Any) -> List[Dict[str, Any]]: + """Extract tool calls from response, handling both formats.""" + tool_calls = [] + + # Get content - handle both dict and object responses + content = None + if isinstance(response, dict): + content = response.get("content", []) + elif hasattr(response, "content"): + content = response.content + + # Anthropic/messages API format: response has "content" list with tool_use blocks + if content: + for block in content: + if isinstance(block, dict) and block.get("type") == "tool_use": + tool_calls.append({ + "id": block.get("id"), + "name": block.get("name"), + "input": block.get("input", {}), + }) + elif hasattr(block, "type") and getattr(block, "type", None) == "tool_use": + tool_calls.append({ + "id": getattr(block, "id", None), + "name": getattr(block, "name", None), + "input": getattr(block, "input", {}), + }) + + # OpenAI format: response has choices[0].message.tool_calls + if not tool_calls and hasattr(response, "choices") and response.choices: # type: ignore[union-attr] + msg = response.choices[0].message # type: ignore[union-attr] + if hasattr(msg, "tool_calls") and msg.tool_calls: + for tc in msg.tool_calls: + tool_calls.append({ + "id": tc.id, + "name": tc.function.name, + "input": json.loads(tc.function.arguments) if tc.function.arguments else {}, + }) + + return tool_calls + + async def _execute_code_loop_messages_api( + self, + data: dict, + response: Any, + skill_files: Dict[str, bytes], + ) -> Any: + """ + Execute the code execution loop for messages API (Anthropic format). + + Returns the final response with generated files inline. + """ + import litellm + from litellm.llms.litellm_proxy.skills.code_execution import ( + LiteLLMInternalTools, + ) + from litellm.llms.litellm_proxy.skills.sandbox_executor import ( + SkillsSandboxExecutor, + ) + + # Ensure response is not None + if response is None: + verbose_proxy_logger.error( + "SkillsInjectionHook: Response is None, cannot execute code loop" + ) + return None + + model = data.get("model", "") + messages = list(data.get("messages", [])) + tools = data.get("tools", []) + max_tokens = data.get("max_tokens", 4096) + + executor = SkillsSandboxExecutor(timeout=self.sandbox_timeout) + generated_files: List[Dict[str, Any]] = [] + current_response = response + + for iteration in range(self.max_iterations): + # Extract tool calls from current response + tool_calls = self._extract_tool_calls(current_response) + stop_reason = current_response.get("stop_reason") if isinstance(current_response, dict) else getattr(current_response, "stop_reason", None) + + # Get content for assistant message - convert to plain dicts + raw_content = current_response.get("content", []) if isinstance(current_response, dict) else getattr(current_response, "content", []) + content_blocks = [] + for block in raw_content or []: + if isinstance(block, dict): + content_blocks.append(block) + elif hasattr(block, "model_dump"): + content_blocks.append(block.model_dump()) + elif hasattr(block, "__dict__"): + content_blocks.append(dict(block.__dict__)) + else: + content_blocks.append({"type": "text", "text": str(block)}) + + # Build assistant message for conversation history (Anthropic format) + assistant_msg = {"role": "assistant", "content": content_blocks} + messages.append(assistant_msg) + + # Check if we're done (no tool calls) + if stop_reason != "tool_use" or not tool_calls: + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Loop completed after {iteration + 1} iterations, " + f"{len(generated_files)} files generated" + ) + return self._attach_files_to_response(current_response, generated_files) + + # Process tool calls + tool_results = [] + for tc in tool_calls: + tool_name = tc.get("name", "") + tool_id = tc.get("id", "") + tool_input = tc.get("input", {}) + + # Execute if it's litellm_code_execution OR a skill tool + if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value: + code = tool_input.get("code", "") + result = await self._execute_code(code, skill_files, executor, generated_files) + elif tool_name.startswith("skill_"): + # Skill tool - execute the skill's code + result = await self._execute_skill_tool(tool_name, tool_input, skill_files, executor, generated_files) + else: + result = f"Tool '{tool_name}' not handled" + + tool_results.append({ + "type": "tool_result", + "tool_use_id": tool_id, + "content": result, + }) + + # Add tool results to messages (Anthropic format) + messages.append({"role": "user", "content": tool_results}) + + # Make next LLM call + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Making LLM call iteration {iteration + 2}" + ) + try: + current_response = await litellm.anthropic.acreate( + model=model, + messages=messages, + tools=tools, + max_tokens=max_tokens, + ) + if current_response is None: + verbose_proxy_logger.error( + "SkillsInjectionHook: LLM call returned None" + ) + return self._attach_files_to_response(response, generated_files) + except Exception as e: + verbose_proxy_logger.error( + f"SkillsInjectionHook: LLM call failed: {e}" + ) + return self._attach_files_to_response(response, generated_files) + + verbose_proxy_logger.warning( + f"SkillsInjectionHook: Max iterations ({self.max_iterations}) reached" + ) + return self._attach_files_to_response(current_response, generated_files) + + async def _execute_code( + self, + code: str, + skill_files: Dict[str, bytes], + executor: Any, + generated_files: List[Dict[str, Any]], + ) -> str: + """Execute code in sandbox and return result string.""" + try: + verbose_proxy_logger.debug(f"SkillsInjectionHook: Executing code ({len(code)} chars)") + + exec_result = executor.execute(code=code, skill_files=skill_files) + + result = exec_result.get("output", "") or "" + + # Collect generated files + if exec_result.get("files"): + for f in exec_result["files"]: + generated_files.append({ + "name": f["name"], + "mime_type": f["mime_type"], + "content_base64": f["content_base64"], + "size": len(base64.b64decode(f["content_base64"])), + }) + result += f"\n\nGenerated file: {f['name']}" + + if exec_result.get("error"): + result += f"\n\nError: {exec_result['error']}" + + return result or "Code executed successfully" + except Exception as e: + return f"Code execution failed: {str(e)}" + + async def _execute_skill_tool( + self, + tool_name: str, + tool_input: Dict[str, Any], + skill_files: Dict[str, bytes], + executor: Any, + generated_files: List[Dict[str, Any]], + ) -> str: + """Execute a skill tool by generating and running code based on skill content.""" + # Generate code based on available skill modules + # Look for Python modules in the skill + python_modules = [p for p in skill_files.keys() if p.endswith(".py") and not p.endswith("__init__.py")] + + # Try to find the main builder/creator module + main_module = None + for mod in python_modules: + if "builder" in mod.lower() or "creator" in mod.lower() or "generator" in mod.lower(): + main_module = mod + break + + if not main_module and python_modules: + # Use first non-init module + main_module = python_modules[0] + + if main_module: + # Convert path to import: "core/gif_builder.py" -> "core.gif_builder" + import_path = main_module.replace("/", ".").replace(".py", "") + + # Generate code that imports and uses the module + code = f""" +# Auto-generated code to execute skill +import sys +sys.path.insert(0, '/sandbox') + +from {import_path} import * + +# Try to find and use a Builder/Creator class +import inspect +module = __import__('{import_path}', fromlist=['']) + +for name, obj in inspect.getmembers(module): + if inspect.isclass(obj) and name != 'object': + try: + instance = obj() + # Try common methods + if hasattr(instance, 'create'): + result = instance.create() + elif hasattr(instance, 'build'): + result = instance.build() + elif hasattr(instance, 'generate'): + result = instance.generate() + elif hasattr(instance, 'save'): + instance.save('output.gif') + print(f'Used {{name}} class') + break + except Exception as e: + print(f'Error with {{name}}: {{e}}') + continue + +# List generated files +import os +for f in os.listdir('.'): + if f.endswith(('.gif', '.png', '.jpg')): + print(f'Generated: {{f}}') +""" + else: + # Fallback generic code + code = """ +print('No executable skill module found') +""" + + return await self._execute_code(code, skill_files, executor, generated_files) + + async def _execute_code_loop( + self, + data: dict, + response: Any, + skill_files: Dict[str, bytes], + ) -> Any: + """ + Execute the code execution loop until model gives final response. + + Returns the final response with generated files inline. + """ + import litellm + from litellm.llms.litellm_proxy.skills.code_execution import ( + LiteLLMInternalTools, + ) + from litellm.llms.litellm_proxy.skills.sandbox_executor import ( + SkillsSandboxExecutor, + ) + + model = data.get("model", "") + messages = list(data.get("messages", [])) + tools = data.get("tools", []) + + # Keys to exclude when passing through to acompletion + # These are either handled explicitly or are internal LiteLLM fields + _EXCLUDED_ACOMPLETION_KEYS = frozenset({ + "messages", + "model", + "tools", + "metadata", + "litellm_metadata", + "container", + }) + + kwargs = { + k: v for k, v in data.items() + if k not in _EXCLUDED_ACOMPLETION_KEYS + } + + executor = SkillsSandboxExecutor(timeout=self.sandbox_timeout) + generated_files: List[Dict[str, Any]] = [] + current_response: Any = response + + for iteration in range(self.max_iterations): + # OpenAI format response has choices[0].message + assistant_message = current_response.choices[0].message # type: ignore[union-attr] + stop_reason = current_response.choices[0].finish_reason # type: ignore[union-attr] + + # Build assistant message for conversation history + assistant_msg_dict: Dict[str, Any] = { + "role": "assistant", + "content": assistant_message.content, + } + if assistant_message.tool_calls: + assistant_msg_dict["tool_calls"] = [ + { + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + } + for tc in assistant_message.tool_calls + ] + messages.append(assistant_msg_dict) + + # Check if we're done (no tool calls) + if stop_reason != "tool_calls" or not assistant_message.tool_calls: + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Code execution loop completed after " + f"{iteration + 1} iterations, {len(generated_files)} files generated" + ) + # Attach generated files to response + return self._attach_files_to_response(current_response, generated_files) + + # Process tool calls + for tool_call in assistant_message.tool_calls: + tool_name = tool_call.function.name + + if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value: + tool_result = await self._execute_code_tool( + tool_call=tool_call, + skill_files=skill_files, + executor=executor, + generated_files=generated_files, + ) + else: + # Non-code-execution tool - cannot handle + tool_result = f"Tool '{tool_name}' not handled automatically" + + messages.append({ + "role": "tool", + "tool_call_id": tool_call.id, + "content": tool_result, + }) + + # Make next LLM call using the messages API + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Making LLM call iteration {iteration + 2}" + ) + current_response = await litellm.anthropic.acreate( + model=model, + messages=messages, + tools=tools, + max_tokens=kwargs.get("max_tokens", 4096), + ) + + # Max iterations reached + verbose_proxy_logger.warning( + f"SkillsInjectionHook: Max iterations ({self.max_iterations}) reached" + ) + return self._attach_files_to_response(current_response, generated_files) + + async def _execute_code_tool( + self, + tool_call: Any, + skill_files: Dict[str, bytes], + executor: Any, + generated_files: List[Dict[str, Any]], + ) -> str: + """Execute a litellm_code_execution tool call and return result string.""" + try: + args = json.loads(tool_call.function.arguments) + code = args.get("code", "") + + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Executing code ({len(code)} chars)" + ) + + exec_result = executor.execute( + code=code, + skill_files=skill_files, + ) + + # Build tool result content + tool_result = exec_result.get("output", "") or "" + + # Collect generated files + if exec_result.get("files"): + tool_result += "\n\nGenerated files:" + for f in exec_result["files"]: + file_content = base64.b64decode(f["content_base64"]) + generated_files.append({ + "name": f["name"], + "mime_type": f["mime_type"], + "content_base64": f["content_base64"], + "size": len(file_content), + }) + tool_result += f"\n- {f['name']} ({len(file_content)} bytes)" + + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Generated file {f['name']} " + f"({len(file_content)} bytes)" + ) + + if exec_result.get("error"): + tool_result += f"\n\nError:\n{exec_result['error']}" + + return tool_result + + except Exception as e: + verbose_proxy_logger.error( + f"SkillsInjectionHook: Code execution failed: {e}" + ) + return f"Code execution failed: {str(e)}" + + def _attach_files_to_response( + self, + response: Any, + generated_files: List[Dict[str, Any]], + ) -> Any: + """ + Attach generated files to the response object. + + Files are added to response._litellm_generated_files for easy access. + For dict responses, files are added as a key. + """ + if not generated_files: + return response + + # Handle dict response (Anthropic/messages API format) + if isinstance(response, dict): + response["_litellm_generated_files"] = generated_files + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Attached {len(generated_files)} files to dict response" + ) + return response + + # Handle object response (OpenAI format) + try: + response._litellm_generated_files = generated_files + except AttributeError: + pass + + # Also add to model_extra if available (for serialization) + if hasattr(response, "model_extra"): + if response.model_extra is None: + response.model_extra = {} + response.model_extra["_litellm_generated_files"] = generated_files + + verbose_proxy_logger.debug( + f"SkillsInjectionHook: Attached {len(generated_files)} files to response" + ) + + return response + + +# Global instance for registration +skills_injection_hook = SkillsInjectionHook() + +import litellm + +litellm.logging_callback_manager.add_litellm_callback(skills_injection_hook) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 647c433ca1a..50c0a86793c 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1912,14 +1912,14 @@ async def info_key_fn( Example Curl: ``` - curl -X GET "http://0.0.0.0:4000/key/info?key=sk-02Wr4IAlN3NvPXvL5JVvDA" \ + curl -X GET "http://0.0.0.0:4000/key/info?key=sk-test-example-key-123" \ -H "Authorization: Bearer sk-1234" ``` Example Curl - if no key is passed, it will use the Key Passed in Authorization Header ``` curl -X GET "http://0.0.0.0:4000/key/info" \ --H "Authorization: Bearer sk-02Wr4IAlN3NvPXvL5JVvDA" +-H "Authorization: Bearer sk-test-example-key-123" ``` """ from litellm.proxy.proxy_server import prisma_client @@ -2312,31 +2312,70 @@ async def _team_key_deletion_check( return False -async def can_delete_verification_token( +async def can_modify_verification_token( key_info: LiteLLM_VerificationToken, user_api_key_cache: DualCache, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, ) -> bool: """ - - check if user is proxy admin - - check if user is team admin and key is a team key - - check if key is personal key + Check if user has permission to modify (delete/regenerate) a verification token. + + Rules: + - Proxy admin can modify any key + - For team keys: only team admin or key owner can modify + - For personal keys: only key owner can modify + + Args: + key_info: The verification token to check + user_api_key_cache: Cache for user API keys + user_api_key_dict: The user making the request + prisma_client: Prisma client for database access + + Returns: + True if user can modify the key, False otherwise """ is_team_key = _is_team_key(data=key_info) + + # 1. Proxy admin can modify any key if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return True - elif is_team_key and key_info.team_id is not None: - return await _team_key_deletion_check( - user_api_key_dict=user_api_key_dict, - key_info=key_info, + + # 2. For team keys: only team admin or key owner can modify + if is_team_key and key_info.team_id is not None: + # Get team object to check if user is team admin + team_table = await get_team_object( + team_id=key_info.team_id, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, + check_db_only=True, ) - elif key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id: - return True - else: + + if team_table is None: + return False + + # Check if user is team admin + if _is_user_team_admin( + user_api_key_dict=user_api_key_dict, + team_obj=team_table, + ): + return True + + # Check if the key belongs to the user (they own it) + if key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id: + return True + + # Not team admin and doesn't own the key return False + + # 3. For personal keys: only key owner can modify + if key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id: + return True + + # Default: deny + return False + + async def delete_verification_tokens( @@ -2385,7 +2424,7 @@ async def delete_verification_tokens( else: authorized_keys: List[LiteLLM_VerificationToken] = [] for key in _keys_being_deleted: - if await can_delete_verification_token( + if await can_modify_verification_token( key_info=key, user_api_key_cache=user_api_key_cache, user_api_key_dict=user_api_key_dict, @@ -2818,6 +2857,18 @@ async def regenerate_key_fn( user_api_key_cache=user_api_key_cache, ) + # check if user has ownership permission to regenerate key + if not await can_modify_verification_token( + key_info=_key_in_db, + user_api_key_cache=user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": "You are not authorized to regenerate this key"}, + ) + verbose_proxy_logger.debug("key_in_db: %s", _key_in_db) new_token = get_new_token(data=data) @@ -2856,14 +2907,8 @@ async def regenerate_key_fn( ### 3. remove existing key entry from cache ###################################################################### - if key: - await _delete_cache_key_object( - hashed_token=hash_token(key), - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - if hashed_api_key: + if hashed_api_key or key: await _delete_cache_key_object( hashed_token=hash_token(key), user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index d51336ef0b3..2ff1183579f 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -2,13 +2,13 @@ import base64 import mimetypes import re from dataclasses import dataclass, field -from typing import List, Literal, Optional, Union +from typing import TYPE_CHECKING, List, Literal, Optional, Union -from fastapi import Request - -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.types.utils import SpecialEnums +if TYPE_CHECKING: + from fastapi import Request + def _is_base64_encoded_unified_file_id(b64_uid: str) -> Union[str, Literal[False]]: # Ensure b64_uid is a string and not a mock object @@ -554,7 +554,7 @@ class FileCreationParams: async def extract_file_creation_params( - request: Request, + request: "Request", request_body: Optional[dict] = None, target_model_names_form: Optional[str] = None, target_storage_form: Optional[str] = None, @@ -571,6 +571,8 @@ async def extract_file_creation_params( Returns: FileCreationParams: Structured parameters extracted from the request """ + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + if request_body is None: request_body = await _read_request_body(request=request) or {} @@ -621,7 +623,7 @@ def _extract_target_model_names_simple(target_model_names_form: Optional[str] = return [] -def _extract_model_param(request: Request, request_body: dict) -> Optional[str]: +def _extract_model_param(request: "Request", request_body: dict) -> Optional[str]: """ Extract model parameter from request. diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index a773e934ef1..2191968e86c 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -1,10 +1,5 @@ model_list: - - model_name: gemini/* + - model_name: anthropic/* litellm_params: - model: gemini/* + model: anthropic/* -litellm_settings: - callbacks: ["dynamic_rate_limiter_v3"] - priority_reservation: - "prod": 0.9 # 90% reserved for production - "dev": 0.1 # 10% reserved for development diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0e591c5e10f..267e0d77422 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4434,7 +4434,7 @@ class ProxyStartupEvent: ) @classmethod - async def initialize_scheduled_background_jobs( + async def initialize_scheduled_background_jobs( # noqa: PLR0915 cls, general_settings: dict, prisma_client: PrismaClient, @@ -4629,6 +4629,37 @@ class ProxyStartupEvent: ) pass + ### CHECK RESPONSES COST ### + if llm_router is not None: + try: + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + check_responses_cost_job = CheckResponsesCost( + proxy_logging_obj=proxy_logging_obj, + prisma_client=prisma_client, + llm_router=llm_router, + ) + scheduler.add_job( + check_responses_cost_job.check_responses_cost, + "interval", + seconds=proxy_batch_polling_interval + + random.randint(0, 30), # Add small random offset + # REMOVED jitter parameter - major cause of memory leak + id="check_responses_cost_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) + verbose_proxy_logger.info("Responses cost check job scheduled successfully") + + except Exception as e: + verbose_proxy_logger.error(f"Failed to setup responses cost checking: {e}") + verbose_proxy_logger.debug( + "Checking responses cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." + ) + pass + # MEMORY LEAK FIX: Start scheduler with paused=False to avoid backlog processing # Do NOT reset job times to "now" as this can trigger the memory leak # The misfire_grace_time and coalesce settings will handle any missed runs properly diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 252b3a7d384..623e8408862 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1,6 +1,6 @@ import asyncio import time -from typing import Any, AsyncIterator, cast +from typing import Any, AsyncIterator, Optional, cast from uuid import uuid4 from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -155,7 +155,7 @@ async def responses_api( # Normal response flow processor = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + response = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -173,6 +173,48 @@ async def responses_api( user_api_base=user_api_base, version=version, ) + + # Store in managed objects table if background mode is enabled + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + from litellm_enterprise.proxy.hooks.managed_files import ( # type: ignore + _PROXY_LiteLLMManagedFiles, + ) + managed_files_obj = cast( + Optional[_PROXY_LiteLLMManagedFiles], + proxy_logging_obj.get_proxy_hook("managed_files"), + ) + + if managed_files_obj and llm_router: + try: + # Get the actual deployment model_id from hidden params + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if not model_id: + verbose_proxy_logger.warning( + f"No model_id found in response hidden params for response {response.id}, skipping managed object storage" + ) + raise Exception("No model_id found in response hidden params") + # Store in managed objects table + await managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + + verbose_proxy_logger.info( + f"Stored background response {response.id} in managed objects table with unified_id={response.id}" + ) + except Exception as e: + verbose_proxy_logger.error( + f"Failed to store background response in managed objects table: {str(e)}" + ) + + return response except ModifyResponseException as e: # Guardrail passthrough: return violation message in Responses API format (200) _data = e.request_data diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index d5fb82808c1..2dd0c5e7556 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -824,4 +824,22 @@ model LiteLLM_UISettings { ui_settings Json created_at DateTime @default(now()) updated_at DateTime @updatedAt +} + +// Skills table for storing LiteLLM-managed skills +model LiteLLM_SkillsTable { + skill_id String @id @default(uuid()) + display_title String? + description String? + instructions String? // The skill instructions/prompt (from SKILL.md) + source String @default("custom") // "custom" or "anthropic" + latest_version String? + file_content Bytes? // Binary content of the skill files (zip) + file_name String? // Original filename + file_type String? // MIME type (e.g., "application/zip") + metadata Json? @default("{}") + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? } \ No newline at end of file diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 774b971de3a..f8ece80707f 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1938,7 +1938,7 @@ async def view_spend_logs( # noqa: PLR0915 Example Request for specific api_key ``` - curl -X GET "http://0.0.0.0:8000/spend/logs?api_key=sk-Fn8Ej39NkBQmUagFEoUWPQ" \ + curl -X GET "http://0.0.0.0:8000/spend/logs?api_key=sk-test-example-key-123" \ -H "Authorization: Bearer sk-1234" ``` diff --git a/litellm/rag/__init__.py b/litellm/rag/__init__.py index f87e72f0c17..54f4d3ccaa0 100644 --- a/litellm/rag/__init__.py +++ b/litellm/rag/__init__.py @@ -5,9 +5,9 @@ Provides an all-in-one API for document ingestion: Upload -> (OCR) -> Chunk -> Embed -> Vector Store """ -from litellm.rag.main import aingest, ingest +from litellm.rag.main import aingest, aquery, ingest, query -__all__ = ["ingest", "aingest"] +__all__ = ["ingest", "aingest", "query", "aquery"] # Expose at litellm.rag level for convenience diff --git a/litellm/rag/main.py b/litellm/rag/main.py index e7a9d3a241f..b8461a8daa6 100644 --- a/litellm/rag/main.py +++ b/litellm/rag/main.py @@ -7,12 +7,22 @@ Upload -> (OCR) -> Chunk -> Embed -> Vector Store from __future__ import annotations -__all__ = ["ingest", "aingest"] +__all__ = ["ingest", "aingest", "query", "aquery"] import asyncio import contextvars from functools import partial -from typing import TYPE_CHECKING, Any, Coroutine, Dict, Optional, Tuple, Type, Union +from typing import ( + TYPE_CHECKING, + Any, + Coroutine, + Dict, + List, + Optional, + Tuple, + Type, + Union, +) import httpx @@ -21,7 +31,14 @@ from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion -from litellm.types.rag import RAGIngestOptions, RAGIngestResponse +from litellm.rag.rag_query import RAGQuery +from litellm.types.rag import ( + RAGIngestOptions, + RAGIngestResponse, + RAGQueryRequest, + RAGQueryResponse, +) +from litellm.types.utils import ModelResponse from litellm.utils import client if TYPE_CHECKING: @@ -172,6 +189,163 @@ async def aingest( ) +async def _execute_query_pipeline( + model: str, + messages: List[Any], + retrieval_config: Dict[str, Any], + rerank: Optional[Dict[str, Any]] = None, + stream: bool = False, + **kwargs, +) -> ModelResponse: + """ + Execute the RAG query pipeline. + """ + # 1. Extract query from last user message + query_text = RAGQuery.extract_query_from_messages(messages) + if not query_text: + raise ValueError("No query found in messages for RAG query") + + # 2. Search vector store + search_response = await litellm.vector_stores.asearch( + vector_store_id=retrieval_config["vector_store_id"], + query=query_text, + max_num_results=retrieval_config.get("top_k", 10), + custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"), + **kwargs, + ) + + rerank_response = None + context_chunks = search_response.get("data", []) + + # 3. Optional rerank + if rerank and rerank.get("enabled"): + documents = RAGQuery.extract_documents_from_search(search_response) + if documents: + rerank_response = await litellm.arerank( + model=rerank["model"], + query=query_text, + documents=documents, + top_n=rerank.get("top_n", 5), + ) + context_chunks = RAGQuery.get_top_chunks_from_rerank( + search_response, rerank_response + ) + + # 4. Build context message and call completion + context_message = RAGQuery.build_context_message(context_chunks) + modified_messages = messages[:-1] + [context_message] + [messages[-1]] + + response = await litellm.acompletion( + model=model, + messages=modified_messages, + stream=stream, + **kwargs, + ) + + # 5. Attach search results to response + if not stream and isinstance(response, ModelResponse): + response = RAGQuery.add_search_results_to_response( + response=response, + search_results=search_response, + rerank_results=rerank_response, + ) + + return response # type: ignore[return-value] + + +@client +async def aquery( + model: str, + messages: List[Any], + retrieval_config: Dict[str, Any], + rerank: Optional[Dict[str, Any]] = None, + stream: bool = False, + **kwargs, +) -> ModelResponse: + """ + Async: Query a RAG pipeline. + """ + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["aquery"] = True + + func = partial( + query, + model=model, + messages=messages, + retrieval_config=retrieval_config, + rerank=rerank, + stream=stream, + **kwargs, + ) + + ctx = contextvars.copy_context() + func_with_context = partial(ctx.run, func) + init_response = await loop.run_in_executor(None, func_with_context) + + if asyncio.iscoroutine(init_response): + response = await init_response + else: + response = init_response + + return response + except Exception as e: + raise litellm.exception_type( + model=model, + custom_llm_provider=retrieval_config.get("custom_llm_provider"), + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def query( + model: str, + messages: List[Any], + retrieval_config: Dict[str, Any], + rerank: Optional[Dict[str, Any]] = None, + stream: bool = False, + **kwargs, +) -> Union[ModelResponse, Coroutine[Any, Any, ModelResponse]]: + """ + Query a RAG pipeline. + """ + local_vars = locals() + try: + _is_async = kwargs.pop("aquery", False) is True + + if _is_async: + return _execute_query_pipeline( + model=model, + messages=messages, + retrieval_config=retrieval_config, + rerank=rerank, + stream=stream, + **kwargs, + ) + else: + return asyncio.get_event_loop().run_until_complete( + _execute_query_pipeline( + model=model, + messages=messages, + retrieval_config=retrieval_config, + rerank=rerank, + stream=stream, + **kwargs, + ) + ) + except Exception as e: + raise litellm.exception_type( + model=model, + custom_llm_provider=retrieval_config.get("custom_llm_provider"), + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + @client def ingest( ingest_options: Dict[str, Any], diff --git a/litellm/rag/rag_query.py b/litellm/rag/rag_query.py new file mode 100644 index 00000000000..53cc6d0089c --- /dev/null +++ b/litellm/rag/rag_query.py @@ -0,0 +1,120 @@ +from typing import Any, Dict, List, Optional, Union, cast + +import litellm +from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage +from litellm.types.utils import ModelResponse +from litellm.types.vector_stores import ( + VectorStoreResultContent, + VectorStoreSearchResponse, + VectorStoreSearchResult, +) + + +class RAGQuery: + CONTENT_PREFIX_STRING = "Context:\n\n" + + @staticmethod + def extract_query_from_messages(messages: List[AllMessageValues]) -> Optional[str]: + """ + Extract the query from the last user message. + """ + if not messages or len(messages) == 0: + return None + + last_message = messages[-1] + if not isinstance(last_message, dict) or "content" not in last_message: + return None + + content = last_message["content"] + + if isinstance(content, str): + return content + elif isinstance(content, list) and len(content) > 0: + # Handle list of content items, extract text from first text item + for item in content: + if ( + isinstance(item, dict) + and item.get("type") == "text" + and "text" in item + ): + return item["text"] + + return None + + @staticmethod + def build_context_message(context_chunks: List[Any]) -> ChatCompletionUserMessage: + """ + Process search results and build a context message. + """ + context_content = RAGQuery.CONTENT_PREFIX_STRING + + for chunk in context_chunks: + if isinstance(chunk, dict): + result_content: Optional[List[VectorStoreResultContent]] = chunk.get( + "content" + ) + if result_content: + for content_item in result_content: + content_text: Optional[str] = content_item.get("text") + if content_text: + context_content += content_text + "\n\n" + elif "text" in chunk: # Fallback for simple dict with text + context_content += chunk["text"] + "\n\n" + elif isinstance(chunk, str): + context_content += chunk + "\n\n" + + return { + "role": "user", + "content": context_content, + } + + @staticmethod + def add_search_results_to_response( + response: ModelResponse, + search_results: VectorStoreSearchResponse, + rerank_results: Optional[Any] = None, + ) -> ModelResponse: + """ + Add search results to the response choices. + """ + if hasattr(response, "choices") and response.choices: + for choice in response.choices: + message = getattr(choice, "message", None) + if message is not None: + # Get existing provider_specific_fields or create new dict + provider_fields = ( + getattr(message, "provider_specific_fields", None) or {} + ) + + # Add search results + provider_fields["search_results"] = search_results + if rerank_results: + provider_fields["rerank_results"] = rerank_results + + # Set the provider_specific_fields + setattr(message, "provider_specific_fields", provider_fields) + return response + + @staticmethod + def extract_documents_from_search( + search_response: Any, + ) -> List[Union[str, Dict[str, Any]]]: + """Extract text documents from vector store search response.""" + documents: List[Union[str, Dict[str, Any]]] = [] + for result in search_response.get("data", []): + content_list = result.get("content", []) + for content in content_list: + if content.get("type") == "text" and content.get("text"): + documents.append(content["text"]) + return documents + + @staticmethod + def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> List[Any]: + """Get the original search results corresponding to the top reranked results.""" + top_chunks = [] + original_results = search_response.get("data", []) + for result in rerank_response.get("results", []): + index = result.get("index") + if index is not None and index < len(original_results): + top_chunks.append(original_results[index]) + return top_chunks diff --git a/litellm/skills/main.py b/litellm/skills/main.py index 2baeb60518e..f6abd9043d4 100644 --- a/litellm/skills/main.py +++ b/litellm/skills/main.py @@ -23,12 +23,27 @@ from litellm.types.llms.anthropic_skills import ( Skill, ) from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager, client # Initialize HTTP handler base_llm_http_handler = BaseLLMHTTPHandler() DEFAULT_ANTHROPIC_API_BASE = "https://api.anthropic.com/v1" +# Initialize LiteLLM skills handler (lazy - only used when custom_llm_provider="litellm") +_litellm_skills_handler = None + + +def _get_litellm_skills_handler(): + """Lazy initialization of LiteLLM skills handler to avoid import overhead.""" + global _litellm_skills_handler + if _litellm_skills_handler is None: + from litellm.llms.litellm_proxy.skills.transformation import ( + LiteLLMSkillsTransformationHandler, + ) + _litellm_skills_handler = LiteLLMSkillsTransformationHandler() + return _litellm_skills_handler + @client async def acreate_skill( @@ -133,18 +148,6 @@ def create_skill( if custom_llm_provider is None: custom_llm_provider = "anthropic" - # Get provider config - skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( - ProviderConfigManager.get_provider_skills_api_config( - provider=litellm.LlmProviders(custom_llm_provider), - ) - ) - - if skills_api_provider_config is None: - raise ValueError( - f"CREATE skill is not supported for {custom_llm_provider}" - ) - # Build create request create_request: CreateSkillRequest = {} if display_title is not None: @@ -156,6 +159,30 @@ def create_skill( if extra_body: create_request.update(extra_body) # type: ignore + # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" + if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: + return _get_litellm_skills_handler().create_skill_handler( + display_title=display_title, + files=files, + metadata=extra_body.get("metadata") if extra_body else None, + user_id=kwargs.get("user_id"), + _is_async=_is_async, + logging_obj=litellm_logging_obj, + litellm_call_id=litellm_call_id, + ) + + # Get provider config for external providers (Anthropic, etc.) + skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( + ProviderConfigManager.get_provider_skills_api_config( + provider=litellm.LlmProviders(custom_llm_provider), + ) + ) + + if skills_api_provider_config is None: + raise ValueError( + f"CREATE skill is not supported for {custom_llm_provider}" + ) + # Validate environment and get headers headers = extra_headers or {} headers = skills_api_provider_config.validate_environment( @@ -316,7 +343,17 @@ def list_skills( if custom_llm_provider is None: custom_llm_provider = "anthropic" - # Get provider config + # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" + if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: + return _get_litellm_skills_handler().list_skills_handler( + limit=limit or 20, + offset=0, + _is_async=_is_async, + logging_obj=litellm_logging_obj, + litellm_call_id=litellm_call_id, + ) + + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), @@ -481,7 +518,16 @@ def get_skill( if custom_llm_provider is None: custom_llm_provider = "anthropic" - # Get provider config + # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" + if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: + return _get_litellm_skills_handler().get_skill_handler( + skill_id=skill_id, + _is_async=_is_async, + logging_obj=litellm_logging_obj, + litellm_call_id=litellm_call_id, + ) + + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), @@ -638,7 +684,16 @@ def delete_skill( if custom_llm_provider is None: custom_llm_provider = "anthropic" - # Get provider config + # Route to LiteLLM DB if custom_llm_provider="litellm_proxy" + if custom_llm_provider == LlmProviders.LITELLM_PROXY.value: + return _get_litellm_skills_handler().delete_skill_handler( + skill_id=skill_id, + _is_async=_is_async, + logging_obj=litellm_logging_obj, + litellm_call_id=litellm_call_id, + ) + + # Get provider config for external providers (Anthropic, etc.) skills_api_provider_config: Optional[BaseSkillsAPIConfig] = ( ProviderConfigManager.get_provider_skills_api_config( provider=litellm.LlmProviders(custom_llm_provider), diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index 23dd661e9ad..371f008c04b 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -358,6 +358,7 @@ class AnthropicMessagesRequestOptionalParams(TypedDict, total=False): top_p: Optional[float] mcp_servers: Optional[List[AnthropicMcpServerTool]] context_management: Optional[Dict[str, Any]] + container: Optional[Dict[str, Any]] # Container config with skills for code execution class AnthropicMessagesRequest(AnthropicMessagesRequestOptionalParams, total=False): diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index dbbab6c1fdc..ceeae958a80 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -903,6 +903,7 @@ class ChatCompletionRequest(TypedDict, total=False): functions: List user: str metadata: dict # litellm specific param + reasoning_effort: str # OpenAI o1/o3 reasoning parameter class ChatCompletionDeltaChunk(TypedDict, total=False): @@ -1028,6 +1029,19 @@ OpenAIImageGenerationOptionalParams = Literal[ "user", ] +OpenAIImageEditOptionalParams = Literal[ + "background", + "n", + "mask" + "output_compression", + "output_format", + "quality", + "partial_images", + "response_format", + "size", + "style", + "user", +] class ComputerToolParam(TypedDict, total=False): display_height: Required[float] diff --git a/litellm/types/llms/stability.py b/litellm/types/llms/stability.py index 33199ff769d..7dd92e380c7 100644 --- a/litellm/types/llms/stability.py +++ b/litellm/types/llms/stability.py @@ -29,6 +29,13 @@ class StabilityImageGenerationRequest(TypedDict, total=False): strength: Optional[float] # How much to transform the image (0-1) style_preset: Optional[str] # Style preset name +class StabilityImageEditRequest(StabilityImageGenerationRequest): + """ + Request parameters for Stability AI image edit endpoint. + + Endpoint: /v2beta/stable-image/edit/inpaint + """ + mask: Optional[str] # Base64-encoded mask (white = edit, black = keep) class StabilityImageGenerationResponse(TypedDict, total=False): """ @@ -197,16 +204,12 @@ STABILITY_EDIT_ENDPOINTS = { "search-and-replace": "/v2beta/stable-image/edit/search-and-replace", "search-and-recolor": "/v2beta/stable-image/edit/search-and-recolor", "remove-background": "/v2beta/stable-image/edit/remove-background", -} - -STABILITY_UPSCALE_ENDPOINTS = { + "replace-background-and-relight": "/v2beta/stable-image/edit/replace-background-and-relight", "fast": "/v2beta/stable-image/upscale/fast", "conservative": "/v2beta/stable-image/upscale/conservative", "creative": "/v2beta/stable-image/upscale/creative", -} - -STABILITY_CONTROL_ENDPOINTS = { "sketch": "/v2beta/stable-image/control/sketch", "structure": "/v2beta/stable-image/control/structure", "style": "/v2beta/stable-image/control/style", + "style-transfer": "/v2beta/stable-image/control/style-transfer", } diff --git a/litellm/types/rag.py b/litellm/types/rag.py index dd724ca217a..fe237a13431 100644 --- a/litellm/types/rag.py +++ b/litellm/types/rag.py @@ -7,6 +7,8 @@ from typing import Any, Dict, List, Literal, Optional, Union from pydantic import BaseModel, ConfigDict from typing_extensions import TypedDict +from litellm.types.utils import ModelResponse + class RAGChunkingStrategy(TypedDict, total=False): """ @@ -187,3 +189,39 @@ class RAGIngestRequest(BaseModel): model_config = ConfigDict(extra="allow") # Allow additional fields + +class RAGRetrievalConfig(TypedDict, total=False): + """Configuration for vector store retrieval.""" + + vector_store_id: str + custom_llm_provider: str + top_k: int # max results from vector store + filters: Optional[Dict[str, Any]] # optional - vector store filters + + +class RAGRerankConfig(TypedDict, total=False): + """Configuration for reranking results.""" + + enabled: bool + model: str + top_n: int # final number of chunks after reranking + return_documents: Optional[bool] + + +class RAGQueryRequest(BaseModel): + """Request body for RAG query API.""" + + model: str + messages: List[Any] + retrieval_config: RAGRetrievalConfig + rerank: Optional[RAGRerankConfig] = None + stream: Optional[bool] = False + + model_config = ConfigDict(extra="allow") + + +class RAGQueryResponse(ModelResponse): + """Response from RAG query API.""" + + pass + diff --git a/litellm/utils.py b/litellm/utils.py index 5fb78e141c1..ce6b2aa9c6a 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7966,6 +7966,18 @@ class ProviderConfigManager: ) return get_vertex_ai_image_edit_config(model) + elif LlmProviders.STABILITY == provider: + from litellm.llms.stability.image_edit import ( + get_stability_image_edit_config, + ) + + return get_stability_image_edit_config(model) + elif LlmProviders.BEDROCK == provider: + from litellm.llms.bedrock.image_edit.stability_transformation import ( + BedrockStabilityImageEditConfig, + ) + + return BedrockStabilityImageEditConfig() return None @staticmethod @@ -7984,9 +7996,13 @@ class ProviderConfigManager: return get_azure_ai_ocr_config(model=model) + if provider == litellm.LlmProviders.VERTEX_AI: + from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config + + return get_vertex_ai_ocr_config(model=model) + PROVIDER_TO_CONFIG_MAP = { litellm.LlmProviders.MISTRAL: MistralOCRConfig, - litellm.LlmProviders.VERTEX_AI: VertexAIOCRConfig, } config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None) if config_class is None: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 0f5f61e708d..8acab0d72d6 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -24483,6 +24483,90 @@ "output_cost_per_image": 0.08, "supported_endpoints": ["/v1/images/generations"] }, + "stability/inpaint": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/outpaint": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.004, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/erase": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/search-and-replace": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/search-and-recolor": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/remove-background": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/replace-background-and-relight": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.008, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/sketch": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/structure": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/style": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.005, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/style-transfer": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.008, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/fast": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.002, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/conservative": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.04, + "supported_endpoints": ["/v1/images/edits"] + }, + "stability/creative": { + "litellm_provider": "stability", + "mode": "image_edit", + "output_cost_per_image": 0.06, + "supported_endpoints": ["/v1/images/edits"] + }, "stability/stable-image-core": { "litellm_provider": "stability", "mode": "image_generation", @@ -24510,6 +24594,84 @@ "mode": "image_generation", "output_cost_per_image": 0.04 }, + "stability.stable-conservative-upscale-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.40 + }, + "stability.stable-creative-upscale-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.60 + }, + "stability.stable-fast-upscale-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.03 + }, + "stability.stable-outpaint-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.06 + }, + "stability.stable-image-control-sketch-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-control-structure-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-erase-object-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-inpaint-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-remove-background-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-search-recolor-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-search-replace-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-image-style-guide-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.07 + }, + "stability.stable-style-transfer-v1:0": { + "litellm_provider": "bedrock", + "max_input_tokens": 77, + "mode": "image_edit", + "output_cost_per_image": 0.08 + }, "stability.stable-image-core-v1:1": { "litellm_provider": "bedrock", "max_input_tokens": 77, @@ -27777,6 +27939,14 @@ ], "source": "https://cloud.google.com/generative-ai-app-builder/pricing" }, + "vertex_ai/deepseek-ai/deepseek-ocr-maas": { + "litellm_provider": "vertex_ai", + "mode": "ocr", + "input_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, + "ocr_cost_per_page": 3e-04, + "source": "https://cloud.google.com/vertex-ai/pricing" + }, "vertex_ai/openai/gpt-oss-120b-maas": { "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-openai_models", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 01d237d4eaf..9f3d6f1bf93 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -84,6 +84,23 @@ "a2a": true } }, + "amazon_nova": { + "display_name": "Amazon Nova (`amazon_nova`)", + "url": "https://docs.litellm.ai/docs/providers/amazon_nova", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": true, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": true + } + }, "anthropic": { "display_name": "Anthropic (`anthropic`)", "url": "https://docs.litellm.ai/docs/providers/anthropic", diff --git a/requirements.txt b/requirements.txt index d95cd47f4e5..d08351c602d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -48,6 +48,7 @@ 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.15 # 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 tiktoken==0.8.0 # for calculating usage diff --git a/schema.prisma b/schema.prisma index d5fb82808c1..2dd0c5e7556 100644 --- a/schema.prisma +++ b/schema.prisma @@ -824,4 +824,22 @@ model LiteLLM_UISettings { ui_settings Json created_at DateTime @default(now()) updated_at DateTime @updatedAt +} + +// Skills table for storing LiteLLM-managed skills +model LiteLLM_SkillsTable { + skill_id String @id @default(uuid()) + display_title String? + description String? + instructions String? // The skill instructions/prompt (from SKILL.md) + source String @default("custom") // "custom" or "anthropic" + latest_version String? + file_content Bytes? // Binary content of the skill files (zip) + file_name String? // Original filename + file_type String? // MIME type (e.g., "application/zip") + metadata Json? @default("{}") + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? } \ No newline at end of file diff --git a/tests/agent_tests/local_vertex_agent.py b/tests/agent_tests/local_vertex_agent.py index 3cc9f868612..cfc202936b3 100644 --- a/tests/agent_tests/local_vertex_agent.py +++ b/tests/agent_tests/local_vertex_agent.py @@ -21,7 +21,7 @@ from google.auth.transport.requests import Request import httpx # Configuration - update these for your agent -PROJECT_ID = "gen-lang-client-0682925754" # Your GCP project ID +PROJECT_ID = "test-gcp-project-id-123" # Your GCP project ID (test value) LOCATION = "us-central1" # Your agent's location # For Reasoning Engines, use just the numeric ID at the end diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index 90bfd6e6479..328589ac2f8 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -136,4 +136,5 @@ polars: >=1.31.0 # Unknown license, the license.md allows free of charge use semantic_router: >=0.1.10 # Unknown license pondpond: >=1.4.1 # Apache 2.0 License fastuuid: >=0.13.0 # BSD-3-Clause license +llm-sandbox: >=0.3.31 # MIT License - https://github.com/vndee/llm-sandbox diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py index 73331547772..96da3271829 100644 --- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py +++ b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py @@ -10,7 +10,7 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path -from litellm.llms.bedrock.image.amazon_nova_canvas_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import ( AmazonNovaCanvasConfig, ) @@ -22,15 +22,15 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import pytest -from litellm.llms.bedrock.image.cost_calculator import cost_calculator +from litellm.llms.bedrock.image_generation.cost_calculator import cost_calculator from litellm.types.utils import ImageResponse, ImageObject import os import litellm -from litellm.llms.bedrock.image.amazon_stability3_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_stability3_transformation import ( AmazonStability3Config, ) -from litellm.llms.bedrock.image.amazon_stability1_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_stability1_transformation import ( AmazonStabilityConfig, ) from litellm.types.llms.bedrock import ( @@ -38,7 +38,7 @@ from litellm.types.llms.bedrock import ( AmazonStability3TextToImageResponse, ) from unittest.mock import MagicMock, patch -from litellm.llms.bedrock.image.image_handler import ( +from litellm.llms.bedrock.image_generation.image_handler import ( BedrockImageGeneration, BedrockImagePreparedRequest, ) @@ -530,7 +530,7 @@ def test_backward_compatibility_regular_nova_model(): def test_amazon_titan_image_gen(): from litellm import image_generation - model_id = "bedrock/amazon.titan-image-generator-v1" + model_id = "bedrock/stability.stable-image-core-v1:1" response = litellm.image_generation( model=model_id, diff --git a/tests/litellm_utils_tests/test_cyberark.py b/tests/litellm_utils_tests/test_cyberark.py index b94e5949534..b7cb25791a9 100644 --- a/tests/litellm_utils_tests/test_cyberark.py +++ b/tests/litellm_utils_tests/test_cyberark.py @@ -13,7 +13,7 @@ from unittest.mock import AsyncMock, MagicMock, patch from litellm._uuid import uuid # Set up environment variables for testing -os.environ["CYBERARK_API_KEY"] = "2syke5r262b6je2f4et1x3jptmry3frfx83t65e6417zad632e5qq8a" +os.environ["CYBERARK_API_KEY"] = "test-cyberark-api-key-909" os.environ["CYBERARK_API_BASE"] = "http://0.0.0.0:8080" os.environ["CYBERARK_ACCOUNT"] = "default" os.environ["CYBERARK_USERNAME"] = "admin" diff --git a/tests/llm_translation/test_azure_openai.py b/tests/llm_translation/test_azure_openai.py index 216da5db8d4..970e68e478b 100644 --- a/tests/llm_translation/test_azure_openai.py +++ b/tests/llm_translation/test_azure_openai.py @@ -193,7 +193,7 @@ def test_process_azure_endpoint_url(api_base, model, expected_endpoint): "azure_deployment": model, "max_retries": 2, "timeout": 600, - "api_key": "f28ab7b695af4154bc53498e5bdccb07", + "api_key": "sk-test-mock-key-505", }, "model": model, } diff --git a/tests/llm_translation/test_bedrock_agentcore.py b/tests/llm_translation/test_bedrock_agentcore.py index 029bdf4e37b..3afb01482ac 100644 --- a/tests/llm_translation/test_bedrock_agentcore.py +++ b/tests/llm_translation/test_bedrock_agentcore.py @@ -218,7 +218,7 @@ def test_bedrock_agentcore_with_api_key_bearer_token(): from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() - test_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.SflKxwRJSMeKKF2QT4fwpMeJf36POk6yJV_adQssw5c" + test_jwt_token = "test-jwt-token-header.payload.signature" with patch.object(client, "post", return_value=MagicMock()) as mock_post: try: diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index bd08d4444f6..78c9f94239b 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -295,7 +295,7 @@ def bedrock_session_token_creds(): aws_role_name = ( "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci" ) - aws_web_identity_token = "oidc/circleci_v2/" + aws_web_identity_token = "test-oidc-token-123" creds = bllm.get_credentials( aws_region_name=aws_region_name, diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator.zip b/tests/llm_translation/test_skills_data/slack-gif-creator.zip new file mode 100644 index 00000000000..15c60e3667d Binary files /dev/null and b/tests/llm_translation/test_skills_data/slack-gif-creator.zip differ diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/LICENSE.txt b/tests/llm_translation/test_skills_data/slack-gif-creator/LICENSE.txt new file mode 100644 index 00000000000..7a4a3ea2424 --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/LICENSE.txt @@ -0,0 +1,202 @@ + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. \ No newline at end of file diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/SKILL.md b/tests/llm_translation/test_skills_data/slack-gif-creator/SKILL.md new file mode 100644 index 00000000000..16660d8ceb7 --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/SKILL.md @@ -0,0 +1,254 @@ +--- +name: slack-gif-creator +description: Knowledge and utilities for creating animated GIFs optimized for Slack. Provides constraints, validation tools, and animation concepts. Use when users request animated GIFs for Slack like "make me a GIF of X doing Y for Slack." +license: Complete terms in LICENSE.txt +--- + +# Slack GIF Creator + +A toolkit providing utilities and knowledge for creating animated GIFs optimized for Slack. + +## Slack Requirements + +**Dimensions:** +- Emoji GIFs: 128x128 (recommended) +- Message GIFs: 480x480 + +**Parameters:** +- FPS: 10-30 (lower is smaller file size) +- Colors: 48-128 (fewer = smaller file size) +- Duration: Keep under 3 seconds for emoji GIFs + +## Core Workflow + +```python +from core.gif_builder import GIFBuilder +from PIL import Image, ImageDraw + +# 1. Create builder +builder = GIFBuilder(width=128, height=128, fps=10) + +# 2. Generate frames +for i in range(12): + frame = Image.new('RGB', (128, 128), (240, 248, 255)) + draw = ImageDraw.Draw(frame) + + # Draw your animation using PIL primitives + # (circles, polygons, lines, etc.) + + builder.add_frame(frame) + +# 3. Save with optimization +builder.save('output.gif', num_colors=48, optimize_for_emoji=True) +``` + +## Drawing Graphics + +### Working with User-Uploaded Images +If a user uploads an image, consider whether they want to: +- **Use it directly** (e.g., "animate this", "split this into frames") +- **Use it as inspiration** (e.g., "make something like this") + +Load and work with images using PIL: +```python +from PIL import Image + +uploaded = Image.open('file.png') +# Use directly, or just as reference for colors/style +``` + +### Drawing from Scratch +When drawing graphics from scratch, use PIL ImageDraw primitives: + +```python +from PIL import ImageDraw + +draw = ImageDraw.Draw(frame) + +# Circles/ovals +draw.ellipse([x1, y1, x2, y2], fill=(r, g, b), outline=(r, g, b), width=3) + +# Stars, triangles, any polygon +points = [(x1, y1), (x2, y2), (x3, y3), ...] +draw.polygon(points, fill=(r, g, b), outline=(r, g, b), width=3) + +# Lines +draw.line([(x1, y1), (x2, y2)], fill=(r, g, b), width=5) + +# Rectangles +draw.rectangle([x1, y1, x2, y2], fill=(r, g, b), outline=(r, g, b), width=3) +``` + +**Don't use:** Emoji fonts (unreliable across platforms) or assume pre-packaged graphics exist in this skill. + +### Making Graphics Look Good + +Graphics should look polished and creative, not basic. Here's how: + +**Use thicker lines** - Always set `width=2` or higher for outlines and lines. Thin lines (width=1) look choppy and amateurish. + +**Add visual depth**: +- Use gradients for backgrounds (`create_gradient_background`) +- Layer multiple shapes for complexity (e.g., a star with a smaller star inside) + +**Make shapes more interesting**: +- Don't just draw a plain circle - add highlights, rings, or patterns +- Stars can have glows (draw larger, semi-transparent versions behind) +- Combine multiple shapes (stars + sparkles, circles + rings) + +**Pay attention to colors**: +- Use vibrant, complementary colors +- Add contrast (dark outlines on light shapes, light outlines on dark shapes) +- Consider the overall composition + +**For complex shapes** (hearts, snowflakes, etc.): +- Use combinations of polygons and ellipses +- Calculate points carefully for symmetry +- Add details (a heart can have a highlight curve, snowflakes have intricate branches) + +Be creative and detailed! A good Slack GIF should look polished, not like placeholder graphics. + +## Available Utilities + +### GIFBuilder (`core.gif_builder`) +Assembles frames and optimizes for Slack: +```python +builder = GIFBuilder(width=128, height=128, fps=10) +builder.add_frame(frame) # Add PIL Image +builder.add_frames(frames) # Add list of frames +builder.save('out.gif', num_colors=48, optimize_for_emoji=True, remove_duplicates=True) +``` + +### Validators (`core.validators`) +Check if GIF meets Slack requirements: +```python +from core.validators import validate_gif, is_slack_ready + +# Detailed validation +passes, info = validate_gif('my.gif', is_emoji=True, verbose=True) + +# Quick check +if is_slack_ready('my.gif'): + print("Ready!") +``` + +### Easing Functions (`core.easing`) +Smooth motion instead of linear: +```python +from core.easing import interpolate + +# Progress from 0.0 to 1.0 +t = i / (num_frames - 1) + +# Apply easing +y = interpolate(start=0, end=400, t=t, easing='ease_out') + +# Available: linear, ease_in, ease_out, ease_in_out, +# bounce_out, elastic_out, back_out +``` + +### Frame Helpers (`core.frame_composer`) +Convenience functions for common needs: +```python +from core.frame_composer import ( + create_blank_frame, # Solid color background + create_gradient_background, # Vertical gradient + draw_circle, # Helper for circles + draw_text, # Simple text rendering + draw_star # 5-pointed star +) +``` + +## Animation Concepts + +### Shake/Vibrate +Offset object position with oscillation: +- Use `math.sin()` or `math.cos()` with frame index +- Add small random variations for natural feel +- Apply to x and/or y position + +### Pulse/Heartbeat +Scale object size rhythmically: +- Use `math.sin(t * frequency * 2 * math.pi)` for smooth pulse +- For heartbeat: two quick pulses then pause (adjust sine wave) +- Scale between 0.8 and 1.2 of base size + +### Bounce +Object falls and bounces: +- Use `interpolate()` with `easing='bounce_out'` for landing +- Use `easing='ease_in'` for falling (accelerating) +- Apply gravity by increasing y velocity each frame + +### Spin/Rotate +Rotate object around center: +- PIL: `image.rotate(angle, resample=Image.BICUBIC)` +- For wobble: use sine wave for angle instead of linear + +### Fade In/Out +Gradually appear or disappear: +- Create RGBA image, adjust alpha channel +- Or use `Image.blend(image1, image2, alpha)` +- Fade in: alpha from 0 to 1 +- Fade out: alpha from 1 to 0 + +### Slide +Move object from off-screen to position: +- Start position: outside frame bounds +- End position: target location +- Use `interpolate()` with `easing='ease_out'` for smooth stop +- For overshoot: use `easing='back_out'` + +### Zoom +Scale and position for zoom effect: +- Zoom in: scale from 0.1 to 2.0, crop center +- Zoom out: scale from 2.0 to 1.0 +- Can add motion blur for drama (PIL filter) + +### Explode/Particle Burst +Create particles radiating outward: +- Generate particles with random angles and velocities +- Update each particle: `x += vx`, `y += vy` +- Add gravity: `vy += gravity_constant` +- Fade out particles over time (reduce alpha) + +## Optimization Strategies + +Only when asked to make the file size smaller, implement a few of the following methods: + +1. **Fewer frames** - Lower FPS (10 instead of 20) or shorter duration +2. **Fewer colors** - `num_colors=48` instead of 128 +3. **Smaller dimensions** - 128x128 instead of 480x480 +4. **Remove duplicates** - `remove_duplicates=True` in save() +5. **Emoji mode** - `optimize_for_emoji=True` auto-optimizes + +```python +# Maximum optimization for emoji +builder.save( + 'emoji.gif', + num_colors=48, + optimize_for_emoji=True, + remove_duplicates=True +) +``` + +## Philosophy + +This skill provides: +- **Knowledge**: Slack's requirements and animation concepts +- **Utilities**: GIFBuilder, validators, easing functions +- **Flexibility**: Create the animation logic using PIL primitives + +It does NOT provide: +- Rigid animation templates or pre-made functions +- Emoji font rendering (unreliable across platforms) +- A library of pre-packaged graphics built into the skill + +**Note on user uploads**: This skill doesn't include pre-built graphics, but if a user uploads an image, use PIL to load and work with it - interpret based on their request whether they want it used directly or just as inspiration. + +Be creative! Combine concepts (bouncing + rotating, pulsing + sliding, etc.) and use PIL's full capabilities. + +## Dependencies + +```bash +pip install pillow imageio numpy +``` diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/core/__init__.py b/tests/llm_translation/test_skills_data/slack-gif-creator/core/__init__.py new file mode 100644 index 00000000000..8b137891791 --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/core/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/core/easing.py b/tests/llm_translation/test_skills_data/slack-gif-creator/core/easing.py new file mode 100644 index 00000000000..772fa830235 --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/core/easing.py @@ -0,0 +1,234 @@ +#!/usr/bin/env python3 +""" +Easing Functions - Timing functions for smooth animations. + +Provides various easing functions for natural motion and timing. +All functions take a value t (0.0 to 1.0) and return eased value (0.0 to 1.0). +""" + +import math + + +def linear(t: float) -> float: + """Linear interpolation (no easing).""" + return t + + +def ease_in_quad(t: float) -> float: + """Quadratic ease-in (slow start, accelerating).""" + return t * t + + +def ease_out_quad(t: float) -> float: + """Quadratic ease-out (fast start, decelerating).""" + return t * (2 - t) + + +def ease_in_out_quad(t: float) -> float: + """Quadratic ease-in-out (slow start and end).""" + if t < 0.5: + return 2 * t * t + return -1 + (4 - 2 * t) * t + + +def ease_in_cubic(t: float) -> float: + """Cubic ease-in (slow start).""" + return t * t * t + + +def ease_out_cubic(t: float) -> float: + """Cubic ease-out (fast start).""" + return (t - 1) * (t - 1) * (t - 1) + 1 + + +def ease_in_out_cubic(t: float) -> float: + """Cubic ease-in-out.""" + if t < 0.5: + return 4 * t * t * t + return (t - 1) * (2 * t - 2) * (2 * t - 2) + 1 + + +def ease_in_bounce(t: float) -> float: + """Bounce ease-in (bouncy start).""" + return 1 - ease_out_bounce(1 - t) + + +def ease_out_bounce(t: float) -> float: + """Bounce ease-out (bouncy end).""" + if t < 1 / 2.75: + return 7.5625 * t * t + elif t < 2 / 2.75: + t -= 1.5 / 2.75 + return 7.5625 * t * t + 0.75 + elif t < 2.5 / 2.75: + t -= 2.25 / 2.75 + return 7.5625 * t * t + 0.9375 + else: + t -= 2.625 / 2.75 + return 7.5625 * t * t + 0.984375 + + +def ease_in_out_bounce(t: float) -> float: + """Bounce ease-in-out.""" + if t < 0.5: + return ease_in_bounce(t * 2) * 0.5 + return ease_out_bounce(t * 2 - 1) * 0.5 + 0.5 + + +def ease_in_elastic(t: float) -> float: + """Elastic ease-in (spring effect).""" + if t == 0 or t == 1: + return t + return -math.pow(2, 10 * (t - 1)) * math.sin((t - 1.1) * 5 * math.pi) + + +def ease_out_elastic(t: float) -> float: + """Elastic ease-out (spring effect).""" + if t == 0 or t == 1: + return t + return math.pow(2, -10 * t) * math.sin((t - 0.1) * 5 * math.pi) + 1 + + +def ease_in_out_elastic(t: float) -> float: + """Elastic ease-in-out.""" + if t == 0 or t == 1: + return t + t = t * 2 - 1 + if t < 0: + return -0.5 * math.pow(2, 10 * t) * math.sin((t - 0.1) * 5 * math.pi) + return math.pow(2, -10 * t) * math.sin((t - 0.1) * 5 * math.pi) * 0.5 + 1 + + +# Convenience mapping +EASING_FUNCTIONS = { + "linear": linear, + "ease_in": ease_in_quad, + "ease_out": ease_out_quad, + "ease_in_out": ease_in_out_quad, + "bounce_in": ease_in_bounce, + "bounce_out": ease_out_bounce, + "bounce": ease_in_out_bounce, + "elastic_in": ease_in_elastic, + "elastic_out": ease_out_elastic, + "elastic": ease_in_out_elastic, +} + + +def get_easing(name: str = "linear"): + """Get easing function by name.""" + return EASING_FUNCTIONS.get(name, linear) + + +def interpolate(start: float, end: float, t: float, easing: str = "linear") -> float: + """ + Interpolate between two values with easing. + + Args: + start: Start value + end: End value + t: Progress from 0.0 to 1.0 + easing: Name of easing function + + Returns: + Interpolated value + """ + ease_func = get_easing(easing) + eased_t = ease_func(t) + return start + (end - start) * eased_t + + +def ease_back_in(t: float) -> float: + """Back ease-in (slight overshoot backward before forward motion).""" + c1 = 1.70158 + c3 = c1 + 1 + return c3 * t * t * t - c1 * t * t + + +def ease_back_out(t: float) -> float: + """Back ease-out (overshoot forward then settle back).""" + c1 = 1.70158 + c3 = c1 + 1 + return 1 + c3 * pow(t - 1, 3) + c1 * pow(t - 1, 2) + + +def ease_back_in_out(t: float) -> float: + """Back ease-in-out (overshoot at both ends).""" + c1 = 1.70158 + c2 = c1 * 1.525 + if t < 0.5: + return (pow(2 * t, 2) * ((c2 + 1) * 2 * t - c2)) / 2 + return (pow(2 * t - 2, 2) * ((c2 + 1) * (t * 2 - 2) + c2) + 2) / 2 + + +def apply_squash_stretch( + base_scale: tuple[float, float], intensity: float, direction: str = "vertical" +) -> tuple[float, float]: + """ + Calculate squash and stretch scales for more dynamic animation. + + Args: + base_scale: (width_scale, height_scale) base scales + intensity: Squash/stretch intensity (0.0-1.0) + direction: 'vertical', 'horizontal', or 'both' + + Returns: + (width_scale, height_scale) with squash/stretch applied + """ + width_scale, height_scale = base_scale + + if direction == "vertical": + # Compress vertically, expand horizontally (preserve volume) + height_scale *= 1 - intensity * 0.5 + width_scale *= 1 + intensity * 0.5 + elif direction == "horizontal": + # Compress horizontally, expand vertically + width_scale *= 1 - intensity * 0.5 + height_scale *= 1 + intensity * 0.5 + elif direction == "both": + # General squash (both dimensions) + width_scale *= 1 - intensity * 0.3 + height_scale *= 1 - intensity * 0.3 + + return (width_scale, height_scale) + + +def calculate_arc_motion( + start: tuple[float, float], end: tuple[float, float], height: float, t: float +) -> tuple[float, float]: + """ + Calculate position along a parabolic arc (natural motion path). + + Args: + start: (x, y) starting position + end: (x, y) ending position + height: Arc height at midpoint (positive = upward) + t: Progress (0.0-1.0) + + Returns: + (x, y) position along arc + """ + x1, y1 = start + x2, y2 = end + + # Linear interpolation for x + x = x1 + (x2 - x1) * t + + # Parabolic interpolation for y + # y = start + progress * (end - start) + arc_offset + # Arc offset peaks at t=0.5 + arc_offset = 4 * height * t * (1 - t) + y = y1 + (y2 - y1) * t - arc_offset + + return (x, y) + + +# Add new easing functions to the convenience mapping +EASING_FUNCTIONS.update( + { + "back_in": ease_back_in, + "back_out": ease_back_out, + "back_in_out": ease_back_in_out, + "anticipate": ease_back_in, # Alias + "overshoot": ease_back_out, # Alias + } +) diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/core/frame_composer.py b/tests/llm_translation/test_skills_data/slack-gif-creator/core/frame_composer.py new file mode 100644 index 00000000000..1afe434811b --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/core/frame_composer.py @@ -0,0 +1,176 @@ +#!/usr/bin/env python3 +""" +Frame Composer - Utilities for composing visual elements into frames. + +Provides functions for drawing shapes, text, emojis, and compositing elements +together to create animation frames. +""" + +from typing import Optional + +import numpy as np +from PIL import Image, ImageDraw, ImageFont + + +def create_blank_frame( + width: int, height: int, color: tuple[int, int, int] = (255, 255, 255) +) -> Image.Image: + """ + Create a blank frame with solid color background. + + Args: + width: Frame width + height: Frame height + color: RGB color tuple (default: white) + + Returns: + PIL Image + """ + return Image.new("RGB", (width, height), color) + + +def draw_circle( + frame: Image.Image, + center: tuple[int, int], + radius: int, + fill_color: Optional[tuple[int, int, int]] = None, + outline_color: Optional[tuple[int, int, int]] = None, + outline_width: int = 1, +) -> Image.Image: + """ + Draw a circle on a frame. + + Args: + frame: PIL Image to draw on + center: (x, y) center position + radius: Circle radius + fill_color: RGB fill color (None for no fill) + outline_color: RGB outline color (None for no outline) + outline_width: Outline width in pixels + + Returns: + Modified frame + """ + draw = ImageDraw.Draw(frame) + x, y = center + bbox = [x - radius, y - radius, x + radius, y + radius] + draw.ellipse(bbox, fill=fill_color, outline=outline_color, width=outline_width) + return frame + + +def draw_text( + frame: Image.Image, + text: str, + position: tuple[int, int], + color: tuple[int, int, int] = (0, 0, 0), + centered: bool = False, +) -> Image.Image: + """ + Draw text on a frame. + + Args: + frame: PIL Image to draw on + text: Text to draw + position: (x, y) position (top-left unless centered=True) + color: RGB text color + centered: If True, center text at position + + Returns: + Modified frame + """ + draw = ImageDraw.Draw(frame) + + # Uses Pillow's default font. + # If the font should be changed for the emoji, add additional logic here. + font = ImageFont.load_default() + + if centered: + bbox = draw.textbbox((0, 0), text, font=font) + text_width = bbox[2] - bbox[0] + text_height = bbox[3] - bbox[1] + x = position[0] - text_width // 2 + y = position[1] - text_height // 2 + position = (x, y) + + draw.text(position, text, fill=color, font=font) + return frame + + +def create_gradient_background( + width: int, + height: int, + top_color: tuple[int, int, int], + bottom_color: tuple[int, int, int], +) -> Image.Image: + """ + Create a vertical gradient background. + + Args: + width: Frame width + height: Frame height + top_color: RGB color at top + bottom_color: RGB color at bottom + + Returns: + PIL Image with gradient + """ + frame = Image.new("RGB", (width, height)) + draw = ImageDraw.Draw(frame) + + # Calculate color step for each row + r1, g1, b1 = top_color + r2, g2, b2 = bottom_color + + for y in range(height): + # Interpolate color + ratio = y / height + r = int(r1 * (1 - ratio) + r2 * ratio) + g = int(g1 * (1 - ratio) + g2 * ratio) + b = int(b1 * (1 - ratio) + b2 * ratio) + + # Draw horizontal line + draw.line([(0, y), (width, y)], fill=(r, g, b)) + + return frame + + +def draw_star( + frame: Image.Image, + center: tuple[int, int], + size: int, + fill_color: tuple[int, int, int], + outline_color: Optional[tuple[int, int, int]] = None, + outline_width: int = 1, +) -> Image.Image: + """ + Draw a 5-pointed star. + + Args: + frame: PIL Image to draw on + center: (x, y) center position + size: Star size (outer radius) + fill_color: RGB fill color + outline_color: RGB outline color (None for no outline) + outline_width: Outline width + + Returns: + Modified frame + """ + import math + + draw = ImageDraw.Draw(frame) + x, y = center + + # Calculate star points + points = [] + for i in range(10): + angle = (i * 36 - 90) * math.pi / 180 # 36 degrees per point, start at top + radius = size if i % 2 == 0 else size * 0.4 # Alternate between outer and inner + px = x + radius * math.cos(angle) + py = y + radius * math.sin(angle) + points.append((px, py)) + + # Draw star + draw.polygon(points, fill=fill_color, outline=outline_color, width=outline_width) + + return frame diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/core/gif_builder.py b/tests/llm_translation/test_skills_data/slack-gif-creator/core/gif_builder.py new file mode 100644 index 00000000000..5759f144fe3 --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/core/gif_builder.py @@ -0,0 +1,269 @@ +#!/usr/bin/env python3 +""" +GIF Builder - Core module for assembling frames into GIFs optimized for Slack. + +This module provides the main interface for creating GIFs from programmatically +generated frames, with automatic optimization for Slack's requirements. +""" + +from pathlib import Path +from typing import Optional + +import imageio.v3 as imageio +import numpy as np +from PIL import Image + + +class GIFBuilder: + """Builder for creating optimized GIFs from frames.""" + + def __init__(self, width: int = 480, height: int = 480, fps: int = 15): + """ + Initialize GIF builder. + + Args: + width: Frame width in pixels + height: Frame height in pixels + fps: Frames per second + """ + self.width = width + self.height = height + self.fps = fps + self.frames: list[np.ndarray] = [] + + def add_frame(self, frame: np.ndarray | Image.Image): + """ + Add a frame to the GIF. + + Args: + frame: Frame as numpy array or PIL Image (will be converted to RGB) + """ + if isinstance(frame, Image.Image): + frame = np.array(frame.convert("RGB")) + + # Ensure frame is correct size + if frame.shape[:2] != (self.height, self.width): + pil_frame = Image.fromarray(frame) + pil_frame = pil_frame.resize( + (self.width, self.height), Image.Resampling.LANCZOS + ) + frame = np.array(pil_frame) + + self.frames.append(frame) + + def add_frames(self, frames: list[np.ndarray | Image.Image]): + """Add multiple frames at once.""" + for frame in frames: + self.add_frame(frame) + + def optimize_colors( + self, num_colors: int = 128, use_global_palette: bool = True + ) -> list[np.ndarray]: + """ + Reduce colors in all frames using quantization. + + Args: + num_colors: Target number of colors (8-256) + use_global_palette: Use a single palette for all frames (better compression) + + Returns: + List of color-optimized frames + """ + optimized = [] + + if use_global_palette and len(self.frames) > 1: + # Create a global palette from all frames + # Sample frames to build palette + sample_size = min(5, len(self.frames)) + sample_indices = [ + int(i * len(self.frames) / sample_size) for i in range(sample_size) + ] + sample_frames = [self.frames[i] for i in sample_indices] + + # Combine sample frames into a single image for palette generation + # Flatten each frame to get all pixels, then stack them + all_pixels = np.vstack( + [f.reshape(-1, 3) for f in sample_frames] + ) # (total_pixels, 3) + + # Create a properly-shaped RGB image from the pixel data + # We'll make a roughly square image from all the pixels + total_pixels = len(all_pixels) + width = min(512, int(np.sqrt(total_pixels))) # Reasonable width, max 512 + height = (total_pixels + width - 1) // width # Ceiling division + + # Pad if necessary to fill the rectangle + pixels_needed = width * height + if pixels_needed > total_pixels: + padding = np.zeros((pixels_needed - total_pixels, 3), dtype=np.uint8) + all_pixels = np.vstack([all_pixels, padding]) + + # Reshape to proper RGB image format (H, W, 3) + img_array = ( + all_pixels[:pixels_needed].reshape(height, width, 3).astype(np.uint8) + ) + combined_img = Image.fromarray(img_array, mode="RGB") + + # Generate global palette + global_palette = combined_img.quantize(colors=num_colors, method=2) + + # Apply global palette to all frames + for frame in self.frames: + pil_frame = Image.fromarray(frame) + quantized = pil_frame.quantize(palette=global_palette, dither=1) + optimized.append(np.array(quantized.convert("RGB"))) + else: + # Use per-frame quantization + for frame in self.frames: + pil_frame = Image.fromarray(frame) + quantized = pil_frame.quantize(colors=num_colors, method=2, dither=1) + optimized.append(np.array(quantized.convert("RGB"))) + + return optimized + + def deduplicate_frames(self, threshold: float = 0.9995) -> int: + """ + Remove duplicate or near-duplicate consecutive frames. + + Args: + threshold: Similarity threshold (0.0-1.0). Higher = more strict (0.9995 = nearly identical). + Use 0.9995+ to preserve subtle animations, 0.98 for aggressive removal. + + Returns: + Number of frames removed + """ + if len(self.frames) < 2: + return 0 + + deduplicated = [self.frames[0]] + removed_count = 0 + + for i in range(1, len(self.frames)): + # Compare with previous frame + prev_frame = np.array(deduplicated[-1], dtype=np.float32) + curr_frame = np.array(self.frames[i], dtype=np.float32) + + # Calculate similarity (normalized) + diff = np.abs(prev_frame - curr_frame) + similarity = 1.0 - (np.mean(diff) / 255.0) + + # Keep frame if sufficiently different + # High threshold (0.9995+) means only remove nearly identical frames + if similarity < threshold: + deduplicated.append(self.frames[i]) + else: + removed_count += 1 + + self.frames = deduplicated + return removed_count + + def save( + self, + output_path: str | Path, + num_colors: int = 128, + optimize_for_emoji: bool = False, + remove_duplicates: bool = False, + ) -> dict: + """ + Save frames as optimized GIF for Slack. + + Args: + output_path: Where to save the GIF + num_colors: Number of colors to use (fewer = smaller file) + optimize_for_emoji: If True, optimize for emoji size (128x128, fewer colors) + remove_duplicates: If True, remove duplicate consecutive frames (opt-in) + + Returns: + Dictionary with file info (path, size, dimensions, frame_count) + """ + if not self.frames: + raise ValueError("No frames to save. Add frames with add_frame() first.") + + output_path = Path(output_path) + + # Remove duplicate frames to reduce file size + if remove_duplicates: + removed = self.deduplicate_frames(threshold=0.9995) + if removed > 0: + print( + f" Removed {removed} nearly identical frames (preserved subtle animations)" + ) + + # Optimize for emoji if requested + if optimize_for_emoji: + if self.width > 128 or self.height > 128: + print( + f" Resizing from {self.width}x{self.height} to 128x128 for emoji" + ) + self.width = 128 + self.height = 128 + # Resize all frames + resized_frames = [] + for frame in self.frames: + pil_frame = Image.fromarray(frame) + pil_frame = pil_frame.resize((128, 128), Image.Resampling.LANCZOS) + resized_frames.append(np.array(pil_frame)) + self.frames = resized_frames + num_colors = min(num_colors, 48) # More aggressive color limit for emoji + + # More aggressive FPS reduction for emoji + if len(self.frames) > 12: + print( + f" Reducing frames from {len(self.frames)} to ~12 for emoji size" + ) + # Keep every nth frame to get close to 12 frames + keep_every = max(1, len(self.frames) // 12) + self.frames = [ + self.frames[i] for i in range(0, len(self.frames), keep_every) + ] + + # Optimize colors with global palette + optimized_frames = self.optimize_colors(num_colors, use_global_palette=True) + + # Calculate frame duration in milliseconds + frame_duration = 1000 / self.fps + + # Save GIF + imageio.imwrite( + output_path, + optimized_frames, + duration=frame_duration, + loop=0, # Infinite loop + ) + + # Get file info + file_size_kb = output_path.stat().st_size / 1024 + file_size_mb = file_size_kb / 1024 + + info = { + "path": str(output_path), + "size_kb": file_size_kb, + "size_mb": file_size_mb, + "dimensions": f"{self.width}x{self.height}", + "frame_count": len(optimized_frames), + "fps": self.fps, + "duration_seconds": len(optimized_frames) / self.fps, + "colors": num_colors, + } + + # Print info + print(f"\n✓ GIF created successfully!") + print(f" Path: {output_path}") + print(f" Size: {file_size_kb:.1f} KB ({file_size_mb:.2f} MB)") + print(f" Dimensions: {self.width}x{self.height}") + print(f" Frames: {len(optimized_frames)} @ {self.fps} fps") + print(f" Duration: {info['duration_seconds']:.1f}s") + print(f" Colors: {num_colors}") + + # Size info + if optimize_for_emoji: + print(f" Optimized for emoji (128x128, reduced colors)") + if file_size_mb > 1.0: + print(f"\n Note: Large file size ({file_size_kb:.1f} KB)") + print(" Consider: fewer frames, smaller dimensions, or fewer colors") + + return info + + def clear(self): + """Clear all frames (useful for creating multiple GIFs).""" + self.frames = [] diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/core/validators.py b/tests/llm_translation/test_skills_data/slack-gif-creator/core/validators.py new file mode 100644 index 00000000000..a6f5bdf28dd --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/core/validators.py @@ -0,0 +1,136 @@ +#!/usr/bin/env python3 +""" +Validators - Check if GIFs meet Slack's requirements. + +These validators help ensure your GIFs meet Slack's size and dimension constraints. +""" + +from pathlib import Path + + +def validate_gif( + gif_path: str | Path, is_emoji: bool = True, verbose: bool = True +) -> tuple[bool, dict]: + """ + Validate GIF for Slack (dimensions, size, frame count). + + Args: + gif_path: Path to GIF file + is_emoji: True for emoji (128x128 recommended), False for message GIF + verbose: Print validation details + + Returns: + Tuple of (passes: bool, results: dict with all details) + """ + from PIL import Image + + gif_path = Path(gif_path) + + if not gif_path.exists(): + return False, {"error": f"File not found: {gif_path}"} + + # Get file size + size_bytes = gif_path.stat().st_size + size_kb = size_bytes / 1024 + size_mb = size_kb / 1024 + + # Get dimensions and frame info + try: + with Image.open(gif_path) as img: + width, height = img.size + + # Count frames + frame_count = 0 + try: + while True: + img.seek(frame_count) + frame_count += 1 + except EOFError: + pass + + # Get duration + try: + duration_ms = img.info.get("duration", 100) + total_duration = (duration_ms * frame_count) / 1000 + fps = frame_count / total_duration if total_duration > 0 else 0 + except: + total_duration = None + fps = None + + except Exception as e: + return False, {"error": f"Failed to read GIF: {e}"} + + # Validate dimensions + if is_emoji: + optimal = width == height == 128 + acceptable = width == height and 64 <= width <= 128 + dim_pass = acceptable + else: + aspect_ratio = ( + max(width, height) / min(width, height) + if min(width, height) > 0 + else float("inf") + ) + dim_pass = aspect_ratio <= 2.0 and 320 <= min(width, height) <= 640 + + results = { + "file": str(gif_path), + "passes": dim_pass, + "width": width, + "height": height, + "size_kb": size_kb, + "size_mb": size_mb, + "frame_count": frame_count, + "duration_seconds": total_duration, + "fps": fps, + "is_emoji": is_emoji, + "optimal": optimal if is_emoji else None, + } + + # Print if verbose + if verbose: + print(f"\nValidating {gif_path.name}:") + print( + f" Dimensions: {width}x{height}" + + ( + f" ({'optimal' if optimal else 'acceptable'})" + if is_emoji and acceptable + else "" + ) + ) + print( + f" Size: {size_kb:.1f} KB" + + (f" ({size_mb:.2f} MB)" if size_mb >= 1.0 else "") + ) + print( + f" Frames: {frame_count}" + + (f" @ {fps:.1f} fps ({total_duration:.1f}s)" if fps else "") + ) + + if not dim_pass: + print( + f" Note: {'Emoji should be 128x128' if is_emoji else 'Unusual dimensions for Slack'}" + ) + + if size_mb > 5.0: + print(f" Note: Large file size - consider fewer frames/colors") + + return dim_pass, results + + +def is_slack_ready( + gif_path: str | Path, is_emoji: bool = True, verbose: bool = True +) -> bool: + """ + Quick check if GIF is ready for Slack. + + Args: + gif_path: Path to GIF file + is_emoji: True for emoji GIF, False for message GIF + verbose: Print feedback + + Returns: + True if dimensions are acceptable + """ + passes, _ = validate_gif(gif_path, is_emoji, verbose) + return passes diff --git a/tests/llm_translation/test_skills_data/slack-gif-creator/requirements.txt b/tests/llm_translation/test_skills_data/slack-gif-creator/requirements.txt new file mode 100644 index 00000000000..8bc4493e916 --- /dev/null +++ b/tests/llm_translation/test_skills_data/slack-gif-creator/requirements.txt @@ -0,0 +1,4 @@ +pillow>=10.0.0 +imageio>=2.31.0 +imageio-ffmpeg>=0.4.9 +numpy>=1.24.0 \ No newline at end of file diff --git a/tests/llm_translation/test_skills_e2e.py b/tests/llm_translation/test_skills_e2e.py new file mode 100644 index 00000000000..9329919ae21 --- /dev/null +++ b/tests/llm_translation/test_skills_e2e.py @@ -0,0 +1,187 @@ +""" +End-to-end test for LiteLLM Skills with Messages API. + +Tests the slack-gif-creator skill with GPT-4o via messages API +to verify skills work correctly and can generate a GIF. +""" + +import os +import sys +import zipfile +from io import BytesIO +from pathlib import Path + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +import litellm.proxy.proxy_server +from litellm.caching.caching import DualCache +from litellm.proxy._types import NewSkillRequest, UserAPIKeyAuth +from litellm.proxy.utils import PrismaClient, ProxyLogging + +proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + + +def create_skill_zip_from_folder(skill_name: str) -> bytes: + """Create a ZIP file from a skill folder in test_skills_data.""" + test_dir = Path(__file__).parent / "test_skills_data" + skill_dir = test_dir / skill_name + + zip_buffer = BytesIO() + with zipfile.ZipFile(zip_buffer, "w", zipfile.ZIP_DEFLATED) as zf: + for file_path in skill_dir.rglob("*"): + if file_path.is_file(): + arcname = f"{skill_name}/{file_path.relative_to(skill_dir)}" + zf.write(file_path, arcname=arcname) + + return zip_buffer.getvalue() + + +@pytest.fixture +def prisma_client(): + """Set up prisma client for tests.""" + from litellm.proxy.proxy_cli import append_query_params + + params = {"connection_limit": 100, "pool_timeout": 60} + database_url = os.getenv("DATABASE_URL") + if not database_url: + pytest.skip("DATABASE_URL not set") + + modified_url = append_query_params(database_url, params) + os.environ["DATABASE_URL"] = modified_url + + prisma_client = PrismaClient( + database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj + ) + + return prisma_client + + +@pytest.mark.asyncio +async def test_slack_gif_skill_creates_gif(prisma_client): + """ + Test slack-gif-creator skill generates a GIF using GPT-4o via messages API. + + Flow: + 1. Store skill in LiteLLM DB + 2. Hook resolves skill, adds litellm_code_execution tool, injects SKILL.md + 3. Make GPT-4o call via messages API + 4. Hook handles code execution loop + 5. Verify GIF is generated + """ + litellm._turn_on_debug() + if not os.getenv("OPENAI_API_KEY"): + pytest.skip("OPENAI_API_KEY not set") + + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + await litellm.proxy.proxy_server.prisma_client.connect() + + from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler + from litellm.proxy.hooks.litellm_skills import SkillsInjectionHook + from litellm.types.utils import CallTypes + + # 1. Store skill in DB + skill_name = "slack-gif-creator" + zip_content = create_skill_zip_from_folder(skill_name) + + skill_request = NewSkillRequest( + display_title="Slack GIF Creator", + description="Create animated GIFs optimized for Slack", + instructions="Use this skill to create animated GIFs for Slack emoji", + file_content=zip_content, + file_name=f"{skill_name}.zip", + file_type="application/zip", + ) + created_skill = await LiteLLMSkillsHandler.create_skill( + data=skill_request, + user_id="test_user", + ) + + print(f"\nCreated skill: {created_skill.skill_id}") + + hook = SkillsInjectionHook() + + try: + # 2. Build request with container.skills (messages API spec) + request_data = { + "model": "claude-sonnet-4-5", + "max_tokens": 4096, + "messages": [ + { + "role": "user", + "content": "Create a simple bouncing red ball GIF for Slack emoji." + } + ], + "container": { + "skills": [ + {"type": "custom", "skill_id": f"litellm:{created_skill.skill_id}"} + ] + }, + } + + # 3. Pre-call hook resolves skill + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + cache = DualCache() + + transformed = await hook.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=request_data, + call_type="anthropic_messages", + ) + assert isinstance(transformed, dict) + + # Hook returns Anthropic-format tools for messages API + tool_names = [t.get('name') for t in transformed.get('tools', [])] + print(f"\nTools after hook: {tool_names}") + assert "litellm_code_execution" in tool_names, "Should have litellm_code_execution tool" + + # 4. Make GPT-4o call via messages API (tools already in Anthropic format) + print("\n--- Making GPT-4o call via messages API ---") + response = await litellm.anthropic.acreate( + model=transformed["model"], + max_tokens=transformed.get("max_tokens", 4096), + messages=transformed["messages"], + tools=transformed.get("tools"), + ) + + print(f"Initial response: {response}") + + # 5. Post-call hook handles code execution loop + final_response = await hook.async_post_call_success_deployment_hook( + request_data=transformed, + response=response, + call_type=CallTypes.anthropic_messages, + ) + + if final_response: + response = final_response + print("Code execution completed!") + + # 6. Check for generated files (handle both dict and object response) + if isinstance(response, dict): + generated_files = response.get("_litellm_generated_files", []) + else: + generated_files = getattr(response, "_litellm_generated_files", []) + print(f"\nGenerated files: {len(generated_files)}") + + if generated_files: + import base64 + for f in generated_files: + print(f" - {f['name']} ({f['size']} bytes)") + if f['name'].endswith('.gif'): + content = base64.b64decode(f['content_base64']) + assert content[:6] in [b'GIF89a', b'GIF87a'], "Should be valid GIF" + print(" Valid GIF!") + print("\nSUCCESS - GIF generated!") + else: + # Print response for debugging + if hasattr(response, "choices"): + print(f"\nResponse: {response.choices[0].message}") + else: + print(f"\nResponse: {response}") + + finally: + await LiteLLMSkillsHandler.delete_skill(skill_id=created_skill.skill_id) diff --git a/tests/local_testing/test_alangfuse.py b/tests/local_testing/test_alangfuse.py index a20370135f9..306c7749f18 100644 --- a/tests/local_testing/test_alangfuse.py +++ b/tests/local_testing/test_alangfuse.py @@ -1019,7 +1019,7 @@ generation_params = { ], }, }, - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "sk-test-mock-api-key-123", "litellm_api_version": "0.0.0", "user_api_key_user_id": "default_user_id", "user_api_key_spend": 0.0, @@ -1142,7 +1142,7 @@ def test_langfuse_prompt_type(prompt): ], }, }, - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "sk-test-mock-api-key-123", "litellm_api_version": "0.0.0", "user_api_key_user_id": "default_user_id", "user_api_key_spend": 0.0, diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index d06568c8796..d72dcbb974b 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -4160,10 +4160,10 @@ def test_openai_hallucinated_tool_call_util(function_name, expect_modification): def test_langfuse_completion(monkeypatch): monkeypatch.setenv( - "LANGFUSE_PUBLIC_KEY", "pk-lf-b3db7e8e-c2f6-4fc7-825c-a541a8fbe003" + "LANGFUSE_PUBLIC_KEY", "test-langfuse-public-key-123" ) monkeypatch.setenv( - "LANGFUSE_SECRET_KEY", "sk-lf-b11ef3a8-361c-4445-9652-12318b8596e4" + "LANGFUSE_SECRET_KEY", "test-langfuse-secret-key-456" ) monkeypatch.setenv("LANGFUSE_HOST", "https://us.cloud.langfuse.com") litellm.set_verbose = True diff --git a/tests/local_testing/test_completion_cost.py b/tests/local_testing/test_completion_cost.py index 40efcc23868..2f78f27361e 100644 --- a/tests/local_testing/test_completion_cost.py +++ b/tests/local_testing/test_completion_cost.py @@ -401,7 +401,7 @@ def test_dalle_3_azure_cost_tracking(): { "b64_json": None, "revised_prompt": "A close-up image of an adorable baby sea otter. Its fur is thick and fluffy to provide buoyancy and insulation against the cold water. Its eyes are round, curious and full of life. It's lying on its back, floating effortlessly on the calm sea surface under the warm sun. Surrounding the otter are patches of colorful kelp drifting along the gentle waves, giving the scene a touch of vibrancy. The sea otter has its small paws folded on its chest, and it seems to be taking a break from its play.", - "url": "https://dalleprodsec.blob.core.windows.net/private/images/3e5d00f3-700e-4b75-869d-2de73c3c975d/generated_00.png?se=2024-03-13T17%3A49%3A51Z&sig=R9RJD5oOSe0Vp9Eg7ze%2FZ8QR7ldRyGH6XhMxiau16Jc%3D&ske=2024-03-19T11%3A08%3A03Z&skoid=e52d5ed7-0657-4f62-bc12-7e5dbb260a96&sks=b&skt=2024-03-12T11%3A08%3A03Z&sktid=33e01921-4d64-4f8c-a055-5bdaffd5e33d&skv=2020-10-02&sp=r&spr=https&sr=b&sv=2020-10-02", + "url": "test-azure-blob-url-with-sas-token", } ], ) diff --git a/tests/local_testing/test_exceptions.py b/tests/local_testing/test_exceptions.py index a27a64dd6e3..987c213d5ca 100644 --- a/tests/local_testing/test_exceptions.py +++ b/tests/local_testing/test_exceptions.py @@ -176,7 +176,7 @@ def invalid_auth(model): # set the model key to an invalid key, depending on th elif "togethercomputer" in model: temporary_key = os.environ["TOGETHERAI_API_KEY"] os.environ["TOGETHERAI_API_KEY"] = ( - "84060c79880fc49df126d3e87b53f8a463ff6e1c6d27fe64207cde25cdfcd1f24a" + "sk-test-togetherai-key-808" ) elif model in litellm.openrouter_models: temporary_key = os.environ["OPENROUTER_API_KEY"] diff --git a/tests/local_testing/test_gcs_bucket.py b/tests/local_testing/test_gcs_bucket.py index fbca0e0060d..2f7d5cd0dec 100644 --- a/tests/local_testing/test_gcs_bucket.py +++ b/tests/local_testing/test_gcs_bucket.py @@ -83,7 +83,7 @@ async def test_aaabasic_gcs_logger(): mock_response="Hi!", metadata={ "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456", "user_api_key_alias": None, "user_api_end_user_max_budget": None, "litellm_api_version": "0.0.0", @@ -155,7 +155,7 @@ async def test_aaabasic_gcs_logger(): assert ( gcs_payload["metadata"]["user_api_key_hash"] - == "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b" + == "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456" ) assert gcs_payload["metadata"]["user_api_key_user_id"] == "116544810872468347480" @@ -191,7 +191,7 @@ async def test_basic_gcs_logger_failure(): metadata={ "gcs_log_id": gcs_log_id, "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456", "user_api_key_alias": None, "user_api_end_user_max_budget": None, "litellm_api_version": "0.0.0", @@ -259,7 +259,7 @@ async def test_basic_gcs_logger_failure(): assert ( gcs_payload["metadata"]["user_api_key_hash"] - == "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b" + == "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456" ) assert gcs_payload["metadata"]["user_api_key_user_id"] == "116544810872468347480" @@ -599,7 +599,7 @@ async def test_basic_gcs_logger_with_folder_in_bucket_name(): mock_response="Hi!", metadata={ "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456", "user_api_key_alias": None, "user_api_end_user_max_budget": None, "litellm_api_version": "0.0.0", @@ -671,7 +671,7 @@ async def test_basic_gcs_logger_with_folder_in_bucket_name(): assert ( gcs_payload["metadata"]["user_api_key_hash"] - == "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b" + == "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456" ) assert gcs_payload["metadata"]["user_api_key_user_id"] == "116544810872468347480" diff --git a/tests/local_testing/test_pass_through_endpoints.py b/tests/local_testing/test_pass_through_endpoints.py index 29cf9682a7c..1c6a7f2c5d8 100644 --- a/tests/local_testing/test_pass_through_endpoints.py +++ b/tests/local_testing/test_pass_through_endpoints.py @@ -446,7 +446,7 @@ async def test_aaapass_through_endpoint_pass_through_keys_langfuse( response = client.post( "/api/public/ingestion", json=_json_data, - headers={"Authorization": "Basic c2stbXktdGVzdC1rZXk6YW55dGhpbmc="}, + headers={"Authorization": "Basic test-base64-auth-token-123"}, ) print("JSON response: ", _json_data) diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py index ac7f5cd6aa1..8a691e7618d 100644 --- a/tests/logging_callback_tests/test_alerting.py +++ b/tests/logging_callback_tests/test_alerting.py @@ -488,7 +488,7 @@ async def test_send_token_budget_crossed_alerts(alerting_type): with patch.object(slack_alerting, "send_alert", new=AsyncMock()) as mock_send_alert: user_info = { - "token": "50e55ca5bfbd0759697538e8d23c0cd5031f52d9e19e176d7233b20c7c4d3403", + "token": "sk-test-mock-token-606", "spend": 86, "max_budget": 100, "user_id": "ishaan@berri.ai", @@ -528,7 +528,7 @@ async def test_webhook_alerting(alerting_type): slack_alerting, "send_webhook_alert", new=AsyncMock() ) as mock_send_alert: user_info = { - "token": "50e55ca5bfbd0759697538e8d23c0cd5031f52d9e19e176d7233b20c7c4d3403", + "token": "sk-test-mock-token-606", "spend": 1, "max_budget": 0, "user_id": "ishaan@berri.ai", @@ -559,7 +559,7 @@ async def test_webhook_alerting(alerting_type): # slack_alerting, "send_webhook_alert", new=AsyncMock() # ) as mock_send_alert: # user_info = { -# "token": "50e55ca5bfbd0759697538e8d23c0cd5031f52d9e19e176d7233b20c7c4d3403", +# "token": "sk-test-mock-token-606", # "spend": 1, # "max_budget": 0, # "user_id": "ishaan@berri.ai", diff --git a/tests/logging_callback_tests/test_spend_logs.py b/tests/logging_callback_tests/test_spend_logs.py index 10c067b7bc9..4f6d4438285 100644 --- a/tests/logging_callback_tests/test_spend_logs.py +++ b/tests/logging_callback_tests/test_spend_logs.py @@ -54,7 +54,7 @@ def test_spend_logs_payload(model_id: Optional[str]): }, "litellm_params": { "acompletion": True, - "api_key": "23c217a5b59f41b6b7a198017f4792f2", + "api_key": "sk-test-mock-key-707", "force_timeout": 600, "logger_fn": None, "verbose": False, @@ -65,7 +65,7 @@ def test_spend_logs_payload(model_id: Optional[str]): "completion_call_id": None, "metadata": { "tags": ["model-anthropic-claude-v2.1", "app-ishaan-prod"], - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "sk-test-mock-api-key-123", "user_api_key_alias": "custom-key-alias", "user_api_end_user_max_budget": None, "litellm_api_version": "0.0.0", @@ -243,7 +243,7 @@ def test_spend_logs_payload_whisper(): "litellm_params": { "api_base": "", "metadata": { - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "sk-test-mock-api-key-123", "user_api_key_alias": None, "user_api_key_end_user_id": "test-user", "user_api_end_user_max_budget": None, diff --git a/tests/logging_callback_tests/test_view_request_resp_logs.py b/tests/logging_callback_tests/test_view_request_resp_logs.py index 34e8d01303a..ea778a44e67 100644 --- a/tests/logging_callback_tests/test_view_request_resp_logs.py +++ b/tests/logging_callback_tests/test_view_request_resp_logs.py @@ -42,7 +42,7 @@ mock_response_data = { "response_time": 0.1622769832611084, "model": "my-fake-model", "metadata": { - "user_api_key_hash": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key_hash": "sk-test-mock-api-key-123", "user_api_key_alias": None, "user_api_key_team_id": None, "user_api_key_org_id": None, diff --git a/tests/ocr_tests/base_ocr_unit_tests.py b/tests/ocr_tests/base_ocr_unit_tests.py index aaa135a4d6b..88d6caf1435 100644 --- a/tests/ocr_tests/base_ocr_unit_tests.py +++ b/tests/ocr_tests/base_ocr_unit_tests.py @@ -41,6 +41,13 @@ class BaseOCRTest(ABC): pytest.skip(f"Rate limit exceeded - {error_msg}") except litellm.InternalServerError: pytest.skip("Model is overloaded") + except litellm.BadRequestError as e: + # Handle URL rejection errors from Vertex AI + error_msg = str(e) + if "URL_REJECTED" in error_msg or "Cannot fetch content from the provided URL" in error_msg: + pytest.skip(f"URL rejected by provider - {error_msg}") + else: + raise @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio diff --git a/tests/ocr_tests/test_ocr_vertex_ai.py b/tests/ocr_tests/test_ocr_vertex_ai.py index 3118871bca8..9b9c10452c5 100644 --- a/tests/ocr_tests/test_ocr_vertex_ai.py +++ b/tests/ocr_tests/test_ocr_vertex_ai.py @@ -1,5 +1,5 @@ """ -Test OCR functionality with Vertex AI Mistral OCR API. +Test OCR functionality with Vertex AI OCR APIs (Mistral and DeepSeek). Note: Vertex AI OCR automatically converts URLs to base64 data URIs since the Vertex AI endpoint doesn't have internet access. @@ -7,6 +7,7 @@ the Vertex AI endpoint doesn't have internet access. import os import json import tempfile +import pytest from base_ocr_unit_tests import BaseOCRTest @@ -50,7 +51,8 @@ def load_vertex_ai_credentials(): # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) -class TestVertexAIOCR(BaseOCRTest): + +class TestVertexAIMistralOCR(BaseOCRTest): """ Test class for Vertex AI Mistral OCR functionality. Inherits from BaseOCRTest and provides Vertex AI-specific configuration. @@ -61,7 +63,7 @@ class TestVertexAIOCR(BaseOCRTest): def get_base_ocr_call_args(self) -> dict: """ - Return the base OCR call args for Vertex AI. + Return the base OCR call args for Vertex AI Mistral OCR. """ load_vertex_ai_credentials() return { @@ -69,3 +71,58 @@ class TestVertexAIOCR(BaseOCRTest): "vertex_location": "us-central1", } + +class TestVertexAIDeepSeekOCR(BaseOCRTest): + """ + Test class for Vertex AI DeepSeek OCR functionality. + Inherits from BaseOCRTest and provides Vertex AI-specific configuration. + + Note: DeepSeek OCR uses the chat completion API format through the openapi endpoint. + Note: DeepSeek OCR does not support PDF URLs - only image URLs and base64 data. + """ + + def get_base_ocr_call_args(self) -> dict: + """ + Return the base OCR call args for Vertex AI DeepSeek OCR. + """ + load_vertex_ai_credentials() + return { + "model": "vertex_ai/deepseek-ocr-maas", + "vertex_location": "us-central1", + } + + # Skip PDF URL tests for DeepSeek OCR as it doesn't support PDF URLs + @pytest.mark.skip(reason="DeepSeek OCR does not support PDF URLs") + async def test_basic_ocr_with_url(self, sync_mode): + """Skip this test for DeepSeek OCR - PDF URLs not supported""" + pass + + @pytest.mark.skip(reason="DeepSeek OCR does not support PDF URLs") + def test_ocr_response_structure(self): + """Skip this test for DeepSeek OCR - PDF URLs not supported""" + pass + + +def test_vertex_ai_ocr_routing(): + """ + Test that Vertex AI OCR routing correctly selects the right config based on model name. + """ + from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config + from litellm.llms.vertex_ai.ocr.deepseek_transformation import VertexAIDeepSeekOCRConfig + from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig + + # Test DeepSeek OCR routing + deepseek_config = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas") + assert isinstance(deepseek_config, VertexAIDeepSeekOCRConfig), \ + "DeepSeek model should route to VertexAIDeepSeekOCRConfig" + + # Test Mistral OCR routing (should use default VertexAIOCRConfig) + mistral_config = get_vertex_ai_ocr_config("vertex_ai/mistral-ocr-2505") + assert isinstance(mistral_config, VertexAIOCRConfig), \ + "Mistral model should route to VertexAIOCRConfig" + + # Test other DeepSeek variants + deepseek_variant = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas") + assert isinstance(deepseek_variant, VertexAIDeepSeekOCRConfig), \ + "DeepSeek variant should route to VertexAIDeepSeekOCRConfig" + diff --git a/tests/old_proxy_tests/tests/test_anthropic_sdk.py b/tests/old_proxy_tests/tests/test_anthropic_sdk.py index 073fafb079b..289fc845549 100644 --- a/tests/old_proxy_tests/tests/test_anthropic_sdk.py +++ b/tests/old_proxy_tests/tests/test_anthropic_sdk.py @@ -6,7 +6,7 @@ client = Anthropic( # This is the default and can be omitted base_url="http://localhost:4000", # this is a litellm proxy key :) - not a real anthropic key - api_key="sk-s4xN1IiLTCytwtZFJaYQrA", + api_key="sk-test-proxy-key-123", ) message = client.messages.create( diff --git a/tests/otel_tests/test_prometheus.py b/tests/otel_tests/test_prometheus.py index 883562e8820..c15d2d9f050 100644 --- a/tests/otel_tests/test_prometheus.py +++ b/tests/otel_tests/test_prometheus.py @@ -8,6 +8,7 @@ import asyncio from litellm._uuid import uuid import os import sys +import hashlib from openai import AsyncOpenAI from typing import Dict, Any @@ -93,7 +94,7 @@ async def test_proxy_failure_metrics(): async with aiohttp.ClientSession() as session: # Make a bad chat completion call status, response_text = await make_bad_chat_completion_request( - session, "sk-1234" + session, "sk-test-1234" ) # Check if the request failed as expected @@ -105,8 +106,12 @@ async def test_proxy_failure_metrics(): print("/metrics", metrics) + # Compute expected hash for test key + test_key = "sk-test-1234" + expected_hash = hashlib.sha256(test_key.encode()).hexdigest() + # Check if the failure metric is present and correct - use pattern matching for robustness - expected_metric_pattern = 'litellm_proxy_failed_requests_metric_total{api_key_alias="None",end_user="None",exception_class="Openai.RateLimitError",exception_status="429",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",route="/chat/completions",team="None",team_alias="None",user="default_user_id",user_email="None"}' + expected_metric_pattern = f'litellm_proxy_failed_requests_metric_total{{api_key_alias="None",end_user="None",exception_class="Openai.RateLimitError",exception_status="429",hashed_api_key="{expected_hash}",requested_model="fake-azure-endpoint",route="/chat/completions",team="None",team_alias="None",user="default_user_id",user_email="None"}}' # Check if the pattern is in metrics (this metric doesn't include user_email field) assert any( @@ -114,7 +119,7 @@ async def test_proxy_failure_metrics(): ), f"Expected failure metric pattern not found in /metrics. Pattern: {expected_metric_pattern}" # Check total requests metric which includes user_email - total_requests_pattern = 'litellm_proxy_total_requests_metric_total{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",route="/chat/completions",status_code="429",team="None",team_alias="None",user="default_user_id",user_email="None"}' + total_requests_pattern = f'litellm_proxy_total_requests_metric_total{{api_key_alias="None",end_user="None",hashed_api_key="{expected_hash}",requested_model="fake-azure-endpoint",route="/chat/completions",status_code="429",team="None",team_alias="None",user="default_user_id",user_email="None"}}' assert any( total_requests_pattern in line for line in metrics.split("\n") @@ -133,7 +138,7 @@ async def test_proxy_success_metrics(): async with aiohttp.ClientSession() as session: # Make a good chat completion call status, response_text = await make_good_chat_completion_request( - session, "sk-1234" + session, "sk-test-1234" ) # Check if the request succeeded as expected @@ -147,14 +152,18 @@ async def test_proxy_success_metrics(): assert END_USER_ID not in metrics + # Compute expected hash for test key + test_key = "sk-test-1234" + expected_hash = hashlib.sha256(test_key.encode()).hexdigest() + # Check if the success metric is present and correct assert ( - 'litellm_request_total_latency_metric_bucket{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",le="0.005",model="fake",requested_model="fake-openai-endpoint",team="None",team_alias="None",user="default_user_id"}' + f'litellm_request_total_latency_metric_bucket{{api_key_alias="None",end_user="None",hashed_api_key="{expected_hash}",le="0.005",model="fake",requested_model="fake-openai-endpoint",team="None",team_alias="None",user="default_user_id"}}' in metrics ) assert ( - 'litellm_llm_api_latency_metric_bucket{api_key_alias="None",end_user="None",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",le="0.005",model="fake",requested_model="fake-openai-endpoint",team="None",team_alias="None",user="default_user_id"}' + f'litellm_llm_api_latency_metric_bucket{{api_key_alias="None",end_user="None",hashed_api_key="{expected_hash}",le="0.005",model="fake",requested_model="fake-openai-endpoint",team="None",team_alias="None",user="default_user_id"}}' in metrics ) @@ -215,7 +224,7 @@ async def test_proxy_fallback_metrics(): async with aiohttp.ClientSession() as session: # Make a good chat completion call - await make_chat_completion_request_with_fallback(session, "sk-1234") + await make_chat_completion_request_with_fallback(session, "sk-test-1234") # Get metrics async with session.get("http://0.0.0.0:4000/metrics") as response: @@ -223,15 +232,19 @@ async def test_proxy_fallback_metrics(): print("/metrics", metrics) + # Compute expected hash for test key + test_key = "sk-test-1234" + expected_hash = hashlib.sha256(test_key.encode()).hexdigest() + # Check if successful fallback metric is incremented assert ( - 'litellm_deployment_successful_fallbacks_total{api_key_alias="None",exception_class="Openai.RateLimitError",exception_status="429",fallback_model="fake-openai-endpoint",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",team="None",team_alias="None"} 1.0' + f'litellm_deployment_successful_fallbacks_total{{api_key_alias="None",exception_class="Openai.RateLimitError",exception_status="429",fallback_model="fake-openai-endpoint",hashed_api_key="{expected_hash}",requested_model="fake-azure-endpoint",team="None",team_alias="None"}} 1.0' in metrics ) # Check if failed fallback metric is incremented assert ( - 'litellm_deployment_failed_fallbacks_total{api_key_alias="None",exception_class="Openai.RateLimitError",exception_status="429",fallback_model="unknown-model",hashed_api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b",requested_model="fake-azure-endpoint",team="None",team_alias="None"} 1.0' + f'litellm_deployment_failed_fallbacks_total{{api_key_alias="None",exception_class="Openai.RateLimitError",exception_status="429",fallback_model="unknown-model",hashed_api_key="{expected_hash}",requested_model="fake-azure-endpoint",team="None",team_alias="None"}} 1.0' in metrics ) @@ -242,7 +255,7 @@ async def create_test_team( """Create a new team and return the team_id""" url = "http://0.0.0.0:4000/team/new" headers = { - "Authorization": "Bearer sk-1234", + "Authorization": "Bearer sk-test-1234", "Content-Type": "application/json", } @@ -260,7 +273,7 @@ async def create_test_user( """Create a new user and return the user info""" url = "http://0.0.0.0:4000/user/new" headers = { - "Authorization": "Bearer sk-1234", + "Authorization": "Bearer sk-test-1234", "Content-Type": "application/json", } @@ -307,7 +320,7 @@ async def create_test_key(session: aiohttp.ClientSession, team_id: str) -> str: """Generate a new key for the team and return it""" url = "http://0.0.0.0:4000/key/generate" headers = { - "Authorization": "Bearer sk-1234", + "Authorization": "Bearer sk-test-1234", "Content-Type": "application/json", } data = { @@ -326,7 +339,7 @@ async def get_team_info(session: aiohttp.ClientSession, team_id: str) -> Dict[st """Fetch team info and return the response""" url = f"http://0.0.0.0:4000/team/info?team_id={team_id}" headers = { - "Authorization": "Bearer sk-1234", + "Authorization": "Bearer sk-test-1234", } async with session.get(url, headers=headers) as response: @@ -415,7 +428,7 @@ async def create_test_key_with_budget( """Generate a new key with budget constraints and return it""" url = "http://0.0.0.0:4000/key/generate" headers = { - "Authorization": "Bearer sk-1234", + "Authorization": "Bearer sk-test-1234", "Content-Type": "application/json", } print("budget_data", budget_data) diff --git a/tests/otel_tests/test_team_member_permissions.py b/tests/otel_tests/test_team_member_permissions.py index d8187e2bc15..062f96de475 100644 --- a/tests/otel_tests/test_team_member_permissions.py +++ b/tests/otel_tests/test_team_member_permissions.py @@ -20,11 +20,12 @@ Valid Permissions: - User tries editing a key with team_id = team_id -> expect to pass. Valid Permissions - - User tries deleting a key with team_id = team_id -> expect to pass. Valid Permissions - + - Note: Delete/regenerate require key ownership or team admin status, not just team member permissions + - User tries deleting a key with team_id = team_id -> expect to fail (403) unless user owns the key or is team admin + - User tries regenerating a key with team_id = team_id -> expect to fail (403) unless user owns the key or is team admin + Invalid Permissions: - User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions - - User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions - User tries calling /key/info with team_id, expect to get valid response @@ -303,10 +304,11 @@ async def test_default_member_permissions(): key=user_key, key_id=team_key, ) - assert "status" in delete_result and delete_result["status"] == 401, "User should not be able to delete keys for team" + assert "status" in delete_result and delete_result["status"] == 403, "User should not be able to delete keys for team" error_data = json.loads(delete_result["error"]) print("error response =", json.dumps(error_data, indent=4)) - assert error_data["error"]["type"] == ProxyErrorTypes.team_member_permission_error.value, "Error should be a team member permission error" + # Delete endpoint now returns 403 with authorization error, not team_member_permission_error + assert "error" in error_data, "Error should contain error field" # User tries regenerating a key with team_id print("Regular team member trying to regenerate a key with team_id. Expecting error.") @@ -318,7 +320,8 @@ async def test_default_member_permissions(): assert "status" in regenerate_result and regenerate_result["status"] == 401, "User should not be able to regenerate keys for team" error_data = json.loads(regenerate_result["error"]) print("error response =", json.dumps(error_data, indent=4)) - assert error_data["error"]["type"] == ProxyErrorTypes.team_member_permission_error.value, "Error should be a team member permission error" + # Regenerate endpoint now returns 403 with authorization error, not team_member_permission_error + assert "error" in error_data, "Error should contain error field" # Test valid permissions # User tries calling /key/info with team_id @@ -378,13 +381,15 @@ async def test_edit_delete_permissions(): ) assert "status" not in update_result, "User should be able to update keys for team" - # User tries deleting a key with team_id - test this last + # User tries deleting a key with team_id + # Note: Even with /key/delete permission, users can only delete keys they own or if they're team admin + # The delete endpoint checks ownership/team admin status, not just team member permissions delete_result = await delete_key( session=session, key=user_key, key_id=key_id ) - assert "status" not in delete_result, "User should be able to delete keys for team" + assert "status" in delete_result and delete_result["status"] == 403, "User should not be able to delete keys they don't own (even with /key/delete permission, ownership is required)" # Test invalid permissions # User tries creating a key with team_id @@ -396,13 +401,14 @@ async def test_edit_delete_permissions(): assert "status" in create_result and create_result["status"] != 200, "User should not be able to create keys for team" # User tries regenerating a key with team_id + # Note: Even with /key/regenerate permission, users can only regenerate keys they own or if they're team admin regenerate_result = await regenerate_key( session=session, key=user_key, key_id=key_id, team_id=team_id ) - assert "status" in regenerate_result and regenerate_result["status"] != 200, "User should not be able to regenerate keys for team" + assert "status" in regenerate_result and regenerate_result["status"] == 401, "User should not be able to regenerate keys they don't own (even with /key/regenerate permission, ownership is required)" @pytest.mark.asyncio() async def test_create_permissions(): @@ -475,13 +481,16 @@ async def test_create_permissions(): key=user_key, key_id=key_id ) - assert "status" in delete_result and delete_result["status"] != 200, "User should not be able to delete keys for team" + assert "status" in delete_result and delete_result["status"] == 403, "User should not be able to delete keys for team" # User tries regenerating a key with team_id + # User doesn't have /key/regenerate permission, so should get 401 (team member permission error) regenerate_result = await regenerate_key( session=session, key=user_key, key_id=key_id, team_id=team_id ) - assert "status" in regenerate_result and regenerate_result["status"] != 200, "User should not be able to regenerate keys for team" \ No newline at end of file + assert "status" in regenerate_result and regenerate_result["status"] == 401, "User should not be able to regenerate keys for team (no /key/regenerate permission)" + error_data = json.loads(regenerate_result["error"]) + assert error_data["error"]["type"] == ProxyErrorTypes.team_member_permission_error.value, "Error should be a team member permission error" \ No newline at end of file diff --git a/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py b/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py index 581f1d19793..97a1f2eecc7 100644 --- a/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py +++ b/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py @@ -105,7 +105,7 @@ def test_create_anthropic_response_logging_payload(mock_logging_obj, metadata_pa kwargs={ "litellm_params": { "metadata": { - "user_api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key": "sk-test-mock-api-key-123", "user_api_key_user_id": "default_user_id", "user_api_key_team_id": None, "user_api_key_end_user_id": ("test" if metadata_params else ""), diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py new file mode 100644 index 00000000000..3bcacdfc05d --- /dev/null +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -0,0 +1,382 @@ +""" +Unit tests for CheckResponsesCost class +""" + +import asyncio +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, Mock, patch + +import pytest + +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + +class TestCheckResponsesCost: + """Test suite for CheckResponsesCost class""" + + @pytest.fixture + def mock_prisma_client(self): + """Create a mock Prisma client""" + client = MagicMock() + client.db = MagicMock() + client.db.litellm_managedobjecttable = MagicMock() + return client + + @pytest.fixture + def mock_proxy_logging_obj(self): + """Create a mock ProxyLogging object""" + logging_obj = MagicMock() + logging_obj.get_proxy_hook = MagicMock(return_value=None) + return logging_obj + + @pytest.fixture + def mock_llm_router(self): + """Create a mock LLM Router""" + router = MagicMock() + router.aget_responses = AsyncMock() + router.get_deployment = MagicMock() + return router + + @pytest.fixture + def check_responses_cost_instance( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Create a CheckResponsesCost instance with mocked dependencies""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + return CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + def test_initialization(self, check_responses_cost_instance): + """Test that CheckResponsesCost initializes correctly""" + assert check_responses_cost_instance.proxy_logging_obj is not None + assert check_responses_cost_instance.prisma_client is not None + assert check_responses_cost_instance.llm_router is not None + + @pytest.mark.asyncio + async def test_check_responses_cost_no_jobs( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost when there are no jobs to process""" + # Mock empty job list + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[] + ) + + # Should not raise any errors + await check_responses_cost_instance.check_responses_cost() + + # Verify find_many was called with correct parameters + mock_prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( + where={ + "status": {"in": ["queued", "in_progress"]}, + "file_purpose": "response", + } + ) + + @pytest.mark.asyncio + async def test_check_responses_cost_with_completed_response( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Test check_responses_cost with a completed response""" + # Mock job with response ID + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_123" + mock_job.created_by = "test-user" + mock_job.id = "job-123" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock completed response + mock_response = ResponsesAPIResponse( + id="resp_123", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check with mocked litellm.aget_responses + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + + await check_responses_cost_instance.check_responses_cost() + + # Verify the job was marked as completed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + assert call_args[1]["data"]["status"] == "completed" + assert call_args[1]["where"]["id"]["in"] == ["job-123"] + + @pytest.mark.asyncio + async def test_check_responses_cost_with_failed_response( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Test check_responses_cost with a failed response""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_456" + mock_job.created_by = "test-user" + mock_job.id = "job-456" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock failed response + mock_response = ResponsesAPIResponse( + id="resp_456", + object="response", + status="failed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + + await check_responses_cost_instance.check_responses_cost() + + # Verify the job was marked as completed (even though response failed) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + assert call_args[1]["data"]["status"] == "completed" + + @pytest.mark.asyncio + async def test_check_responses_cost_with_cancelled_response( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with a cancelled response""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_789" + mock_job.created_by = "test-user" + mock_job.id = "job-789" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock cancelled response + mock_response = ResponsesAPIResponse( + id="resp_789", + object="response", + status="cancelled", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + + await check_responses_cost_instance.check_responses_cost() + + # Verify the job was marked as completed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_in_progress_response( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with a response still in progress""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_in_progress" + mock_job.created_by = "test-user" + mock_job.id = "job-in-progress" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock in-progress response + mock_response = ResponsesAPIResponse( + id="resp_in_progress", + object="response", + status="in_progress", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + + await check_responses_cost_instance.check_responses_cost() + + # Verify no updates were made (response still in progress) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_queued_response( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with a queued response""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_queued" + mock_job.created_by = "test-user" + mock_job.id = "job-queued" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock queued response + mock_response = ResponsesAPIResponse( + id="resp_queued", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = mock_response + + await check_responses_cost_instance.check_responses_cost() + + # Verify no updates were made (response still queued) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_exception( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost handles exceptions gracefully""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_error" + mock_job.created_by = "test-user" + mock_job.id = "job-error" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check with mocked exception + with patch( + "litellm.aget_responses", + new_callable=AsyncMock, + side_effect=Exception("Provider error"), + ): + # Should not raise, just skip the job + await check_responses_cost_instance.check_responses_cost() + + # Verify no updates were made (job was skipped due to error) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_multiple_jobs( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with multiple jobs""" + # Mock multiple jobs + mock_job1 = MagicMock() + mock_job1.unified_object_id = "resp_test_1" + mock_job1.created_by = "user1" + mock_job1.id = "job-1" + + mock_job2 = MagicMock() + mock_job2.unified_object_id = "resp_test_2" + mock_job2.created_by = "user2" + mock_job2.id = "job-2" + + mock_job3 = MagicMock() + mock_job3.unified_object_id = "resp_test_3" + mock_job3.created_by = "user3" + mock_job3.id = "job-3" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job1, mock_job2, mock_job3] + ) + + # Mock responses - 2 completed, 1 in progress + mock_response1 = ResponsesAPIResponse( + id="resp_1", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + mock_response2 = ResponsesAPIResponse( + id="resp_2", + object="response", + status="in_progress", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + mock_response3 = ResponsesAPIResponse( + id="resp_3", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=200, + output_tokens=100, + total_tokens=300, + ), + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.side_effect = [mock_response1, mock_response2, mock_response3] + + await check_responses_cost_instance.check_responses_cost() + + # Verify only the 2 completed jobs were marked as complete + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + assert len(call_args[1]["where"]["id"]["in"]) == 2 + assert "job-1" in call_args[1]["where"]["id"]["in"] + assert "job-3" in call_args[1]["where"]["id"]["in"] + assert "job-2" not in call_args[1]["where"]["id"]["in"] diff --git a/tests/proxy_unit_tests/test_db_schema_migration.py b/tests/proxy_unit_tests/test_db_schema_migration.py index b3178183759..a8fa3242129 100644 --- a/tests/proxy_unit_tests/test_db_schema_migration.py +++ b/tests/proxy_unit_tests/test_db_schema_migration.py @@ -21,7 +21,7 @@ def test_aaaasschema_migration_check(schema_setup, monkeypatch): """Test to check if schema requires migration""" # Set test database URL test_db_url = f"postgresql://{schema_setup.info.user}:@{schema_setup.info.host}:{schema_setup.info.port}/{schema_setup.info.dbname}" - # test_db_url = "postgresql://neondb_owner:npg_JiZPS0DAhRn4@ep-delicate-wave-a55cvbuc.us-east-2.aws.neon.tech/neondb?sslmode=require" + # test_db_url = "postgresql://test-user:test-password@test-host.example.com/test-db?sslmode=require" monkeypatch.setenv("DATABASE_URL", test_db_url) deploy_dir = Path("./litellm-proxy-extras/litellm_proxy_extras") diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index 57434993977..2af61aa2653 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -1266,7 +1266,7 @@ def test_user_api_key_auth_jwt_hashing(): from litellm.proxy.auth.handle_jwt import JWTHandler # Test with a JWT token (3 parts separated by dots) - jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.SflKxwRJSMeKKF2QT4fwpMeJf36POk6yJV_adQssw5c" + jwt_token = "test-jwt-token-header.payload.signature" # Create UserAPIKeyAuth instance with JWT user_auth = UserAPIKeyAuth(api_key=jwt_token) @@ -1303,7 +1303,7 @@ def test_jwt_handler_is_jwt_static_method(): from litellm.proxy.auth.handle_jwt import JWTHandler # Test with valid JWT format - valid_jwt = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.SflKxwRJSMeKKF2QT4fwpMeJf36POk6yJV_adQssw5c" + valid_jwt = "test-jwt-token-header.payload.signature" assert JWTHandler.is_jwt(valid_jwt) == True # Test with invalid JWT format (only 2 parts) diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index c88efe2ebe2..2e5cfff8bf0 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -238,7 +238,7 @@ def test_dynamic_logging_metadata_key_and_team_metadata(callback_vars): proxy_config = ProxyConfig() user_api_key_dict = UserAPIKeyAuth( - token="6f8688eaff1d37555bb9e9a6390b6d7032b3ab2526ba0152da87128eab956432", + token="sk-test-mock-token-789", key_name="sk-...63Fg", key_alias=None, spend=0.000111, @@ -287,7 +287,7 @@ def test_dynamic_logging_metadata_key_and_team_metadata(callback_vars): end_user_rpm_limit=None, end_user_max_budget=None, last_refreshed_at=1726101560.967527, - api_key="7c305cc48fe72272700dc0d67dc691c2d1f2807490ef5eb2ee1d3a3ca86e12b1", + api_key="sk-test-mock-api-key-202", user_role=LitellmUserRoles.INTERNAL_USER, allowed_model_region=None, parent_otel_span=None, @@ -320,7 +320,7 @@ def test_dynamic_turn_off_message_logging(callback_vars): proxy_config = ProxyConfig() user_api_key_dict = UserAPIKeyAuth( - token="6f8688eaff1d37555bb9e9a6390b6d7032b3ab2526ba0152da87128eab956432", + token="sk-test-mock-token-789", key_name="sk-...63Fg", key_alias=None, spend=0.000111, @@ -368,7 +368,7 @@ def test_dynamic_turn_off_message_logging(callback_vars): end_user_rpm_limit=None, end_user_max_budget=None, last_refreshed_at=1726101560.967527, - api_key="7c305cc48fe72272700dc0d67dc691c2d1f2807490ef5eb2ee1d3a3ca86e12b1", + api_key="sk-test-mock-api-key-202", user_role=LitellmUserRoles.INTERNAL_USER, allowed_model_region=None, parent_otel_span=None, @@ -1267,7 +1267,7 @@ def test_litellm_verification_token_view_response_with_budget_table( from litellm.proxy._types import LiteLLM_VerificationTokenView args: Dict[str, Any] = { - "token": "78b627d4d14bc3acf5571ae9cb6834e661bc8794d1209318677387add7621ce1", + "token": "sk-test-mock-token-303", "key_name": "sk-...if_g", "key_alias": None, "soft_budget_cooldown": False, diff --git a/tests/proxy_unit_tests/test_skills_db.py b/tests/proxy_unit_tests/test_skills_db.py new file mode 100644 index 00000000000..ec72087849d --- /dev/null +++ b/tests/proxy_unit_tests/test_skills_db.py @@ -0,0 +1,257 @@ +""" +Test LiteLLM Skills SDK with custom_llm_provider=litellm_proxy + +Tests the SDK-level skills methods when using the LiteLLM database backend: +1. Create a skill using SDK and verify it was stored correctly +2. List skills using SDK +3. Get a skill by ID using SDK +4. Delete a skill using SDK +5. Skills injection hook correctly resolves skills from database +""" + +import os +import sys +import zipfile +from contextlib import contextmanager +from io import BytesIO +from pathlib import Path + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +from litellm.caching.caching import DualCache +from litellm.proxy import proxy_server +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.types.utils import LlmProviders + +proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + + +@contextmanager +def create_skill_zip(skill_name: str): + """ + Helper context manager to create a zip file for a skill. + + Args: + skill_name: Name of the skill directory in test_skills_data/ + + Yields: + Tuple of (file handle, file content bytes) + + The zip file is automatically cleaned up after use. + """ + test_dir = Path(__file__).parent.parent / "llm_translation" / "test_skills_data" + skill_dir = test_dir / skill_name + + # Create a zip file containing the skill directory + zip_path = test_dir / f"{skill_name}.zip" + with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zip_file: + zip_file.write(skill_dir, arcname=skill_name) + zip_file.write(skill_dir / "SKILL.md", arcname=f"{skill_name}/SKILL.md") + + try: + with open(zip_path, "rb") as f: + content = f.read() + f.seek(0) + yield f, content + finally: + # Clean up zip file + if zip_path.exists(): + zip_path.unlink() + + +@pytest.fixture +def prisma_client(): + """Set up prisma client for tests.""" + from litellm.proxy.proxy_cli import append_query_params + + params = {"connection_limit": 100, "pool_timeout": 60} + database_url = os.getenv("DATABASE_URL") + modified_url = append_query_params(database_url, params) + os.environ["DATABASE_URL"] = modified_url + + prisma_client = PrismaClient( + database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj + ) + + return prisma_client + + +@pytest.mark.asyncio +async def test_create_skill_sdk(prisma_client): + """ + Test creating a skill using SDK with custom_llm_provider=litellm_proxy. + + Verifies that: + - Skill is created with correct display_title + - Skill ID is generated and returned + - Skill response has correct type + """ + setattr(proxy_server, "prisma_client", prisma_client) + await proxy_server.prisma_client.connect() + + from litellm.skills.main import acreate_skill, adelete_skill + + # Create a skill using SDK + skill = await acreate_skill( + display_title="SDK Test Skill", + extra_body={ + "description": "A test skill created via SDK", + "instructions": "Use this skill for SDK testing", + }, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + # Verify skill was created correctly + assert skill is not None + assert skill.id is not None + assert skill.id.startswith("skill_") + assert skill.display_title == "SDK Test Skill" + assert skill.type == "skill" + assert skill.source == "custom" + + # Clean up + await adelete_skill( + skill_id=skill.id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + +@pytest.mark.asyncio +async def test_list_skills_sdk(prisma_client): + """ + Test listing skills using SDK with custom_llm_provider=litellm_proxy. + + Verifies that: + - Multiple skills can be created + - List returns the created skills + """ + setattr(proxy_server, "prisma_client", prisma_client) + await proxy_server.prisma_client.connect() + + from litellm.skills.main import acreate_skill, adelete_skill, alist_skills + + # Create multiple skills + created_skill_ids = [] + for i in range(3): + skill = await acreate_skill( + display_title=f"List Test Skill {i}", + extra_body={ + "description": f"Test skill {i} for list test", + }, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + created_skill_ids.append(skill.id) + + # List skills using SDK + response = await alist_skills( + limit=10, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + # Verify we got skills back + assert response is not None + assert response.data is not None + assert len(response.data) >= 3 + + # Verify our created skills are in the list + skill_ids_in_list = [s.id for s in response.data] + for created_id in created_skill_ids: + assert created_id in skill_ids_in_list + + # Clean up + for skill_id in created_skill_ids: + await adelete_skill( + skill_id=skill_id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + +@pytest.mark.asyncio +async def test_get_skill_sdk(prisma_client): + """ + Test getting a skill by ID using SDK with custom_llm_provider=litellm_proxy. + + Verifies that: + - Skill can be retrieved by ID + - Retrieved skill has correct data + """ + setattr(proxy_server, "prisma_client", prisma_client) + await proxy_server.prisma_client.connect() + + from litellm.skills.main import acreate_skill, adelete_skill, aget_skill + + # Create a skill + created_skill = await acreate_skill( + display_title="Get Test Skill", + extra_body={ + "description": "A skill for get test", + }, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + # Get the skill by ID using SDK + retrieved_skill = await aget_skill( + skill_id=created_skill.id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + # Verify retrieved skill matches created skill + assert retrieved_skill is not None + assert retrieved_skill.id == created_skill.id + assert retrieved_skill.display_title == "Get Test Skill" + + # Clean up + await adelete_skill( + skill_id=created_skill.id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + +@pytest.mark.asyncio +async def test_delete_skill_sdk(prisma_client): + """ + Test deleting a skill using SDK with custom_llm_provider=litellm_proxy. + + Verifies that: + - Skill can be deleted by ID + - Deleted skill cannot be retrieved + """ + setattr(proxy_server, "prisma_client", prisma_client) + await proxy_server.prisma_client.connect() + + from litellm.skills.main import acreate_skill, adelete_skill, aget_skill + + # Create a skill + created_skill = await acreate_skill( + display_title="Delete Test Skill", + extra_body={ + "description": "A skill to be deleted", + }, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + + # Verify skill exists + retrieved = await aget_skill( + skill_id=created_skill.id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + assert retrieved is not None + + # Delete the skill using SDK + result = await adelete_skill( + skill_id=created_skill.id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) + assert result.id == created_skill.id + assert result.type == "skill_deleted" + + # Verify skill no longer exists + with pytest.raises(Exception): + await aget_skill( + skill_id=created_skill.id, + custom_llm_provider=LlmProviders.LITELLM_PROXY.value, + ) diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index ec61c7305bb..72d13aadad3 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -696,7 +696,7 @@ def test_is_allowed_route(): "request": request, "request_data": {"input": ["hello world"], "model": "embedding-small"}, "valid_token": UserAPIKeyAuth( - token="9644159bc181998825c44c788b1526341ed2e825d1b6f562e23173759e14bb86", + token="sk-test-mock-token-101", key_name="sk-...CJjQ", key_alias=None, spend=0.0, diff --git a/tests/test_callbacks_on_proxy.py b/tests/test_callbacks_on_proxy.py index 831ca449f83..3bc07da8db1 100644 --- a/tests/test_callbacks_on_proxy.py +++ b/tests/test_callbacks_on_proxy.py @@ -26,7 +26,7 @@ async def config_update(session, routing_strategy=None): }, "general_settings": { "alert_to_webhook_url": { - "llm_exceptions": "https://hooks.slack.com/services/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" + "llm_exceptions": "example-slack-webhook-url" }, "alert_types": ["llm_exceptions", "db_exceptions"], }, diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 2ef27396585..2e7a64df8be 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -142,6 +142,84 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag print("✓ Tool result with image correctly transformed to Responses API format") +def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text(): + """ + Test that tool messages with text content are correctly transformed to Responses API format. + + This is a regression test for the issue where tool results were being transformed + with type='output_text' instead of type='input_text', which caused OpenAI's Responses API + to reject the request with "Invalid value: 'output_text'". + + Chat Completion format: + {"role": "tool", "tool_call_id": "call_abc123", "content": "15 degrees"} + + Responses API format should use input_text, not output_text: + {"type": "function_call_output", "call_id": "call_abc123", "output": [{"type": "input_text", "text": "15 degrees"}]} + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + + # Chat Completion format with tool result containing text + messages = [ + { + "role": "user", + "content": "What is the weather like in San Francisco?", + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "San Francisco, CA", "unit": "celsius"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_abc123", + "content": "15 degrees", + }, + ] + + response, _ = handler.convert_chat_completion_messages_to_responses_api(messages) + + # Find the function_call_output item + function_call_output = None + for item in response: + if item.get("type") == "function_call_output": + function_call_output = item + break + + assert ( + function_call_output is not None + ), "function_call_output not found in response" + assert function_call_output["call_id"] == "call_abc123" + + # Check that the output is correctly transformed to use input_text, not output_text + output = function_call_output["output"] + assert isinstance(output, list), "output should be a list" + assert len(output) == 1, "output should have one item" + + text_item = output[0] + # Should be transformed to use input_text for tool results in Responses API format + assert ( + text_item["type"] == "input_text" + ), f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" + assert ( + text_item["text"] == "15 degrees" + ), f"Expected text '15 degrees', got '{text_item.get('text')}'" + + print("✓ Tool result with text correctly transformed to use input_text for Responses API format") + + def test_openai_responses_chunk_parser_reasoning_summary(): from litellm.completion_extras.litellm_responses_transformation.transformation import ( OpenAiResponsesToChatCompletionStreamIterator, @@ -717,3 +795,213 @@ def test_text_plus_tool_calls_sequence(): assert ( completed_result.choices[0].finish_reason == "stop" ), "response.completed should have finish_reason='stop'" + + +# ============================================================================= +# Tests for issue #18201: Tool calls transformation fixes +# ============================================================================= + + +def test_tool_message_output_is_string_not_list(): + """ + Test that tool message content is converted to a string, not a list. + + This is a regression test for a bug where tool results were transformed to: + {"type": "function_call_output", "output": [{"type": "output_text", "text": "..."}]} + + But the Responses API expects: + {"type": "function_call_output", "output": "..."} + + The incorrect format caused OpenAI to reject with: + "Invalid value: 'output_text'. Supported values are: 'input_text', 'input_image', and 'input_file'." + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + + messages = [ + {"role": "user", "content": "What's the weather?"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Paris"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_abc123", + "content": '{"temperature": 15, "condition": "sunny"}', + }, + ] + + response, _ = handler.convert_chat_completion_messages_to_responses_api(messages) + + # Find the function_call_output item + function_call_output = None + for item in response: + if item.get("type") == "function_call_output": + function_call_output = item + break + + assert function_call_output is not None, "function_call_output not found" + assert function_call_output["call_id"] == "call_abc123" + + # The output should be a string, NOT a list + output = function_call_output["output"] + assert isinstance(output, str), f"output should be a string, got {type(output)}" + assert output == '{"temperature": 15, "condition": "sunny"}' + + print("✓ Tool message output is correctly a string, not a list") + + +def test_multiple_tool_calls_in_single_choice(): + """ + Test that multiple tool calls are grouped into a single choice. + + This is a regression test for a bug where each tool call was put in its own + Choice with separate indices: + choices = [ + {"index": 0, "message": {"tool_calls": [tc1]}}, + {"index": 1, "message": {"tool_calls": [tc2]}}, + {"index": 2, "message": {"tool_calls": [tc3]}}, + ] + + But Chat Completions API expects all tool calls in a single choice: + choices = [ + {"index": 0, "message": {"tool_calls": [tc1, tc2, tc3]}}, + ] + """ + from unittest.mock import Mock + + from openai.types.responses import ResponseFunctionToolCall + + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + from litellm.types.llms.openai import ( + InputTokensDetails, + OutputTokensDetails, + ResponseAPIUsage, + ResponsesAPIResponse, + ) + from litellm.types.utils import ModelResponse, Usage + + handler = LiteLLMResponsesTransformationHandler() + + # Create multiple function tool calls (simulating parallel tool calls) + tool_call_1 = ResponseFunctionToolCall( + id="fc_1", + type="function_call", + status="completed", + arguments='{"location": "Paris"}', + call_id="call_paris", + name="get_weather", + ) + tool_call_2 = ResponseFunctionToolCall( + id="fc_2", + type="function_call", + status="completed", + arguments='{"location": "Tokyo"}', + call_id="call_tokyo", + name="get_weather", + ) + tool_call_3 = ResponseFunctionToolCall( + id="fc_3", + type="function_call", + status="completed", + arguments='{"sign": "Leo"}', + call_id="call_horoscope", + name="get_horoscope", + ) + + usage = ResponseAPIUsage( + input_tokens=50, + input_tokens_details=InputTokensDetails(cached_tokens=0), + output_tokens=100, + output_tokens_details=OutputTokensDetails(reasoning_tokens=0), + total_tokens=150, + ) + + raw_response = ResponsesAPIResponse( + id="resp_test", + created_at=1234567890, + error=None, + incomplete_details=None, + instructions=None, + metadata={}, + model="gpt-4o", + object="response", + output=[tool_call_1, tool_call_2, tool_call_3], + parallel_tool_calls=True, + temperature=1.0, + tool_choice="auto", + tools=[], + top_p=1.0, + max_output_tokens=None, + previous_response_id=None, + reasoning=None, + status="completed", + text=None, + truncation="disabled", + usage=usage, + user=None, + store=True, + background=False, + ) + + model_response = ModelResponse( + id="chatcmpl-test", + created=1234567890, + model=None, + object="chat.completion", + choices=[], + usage=Usage(completion_tokens=0, prompt_tokens=0, total_tokens=0), + ) + + logging_obj = Mock() + + result = handler.transform_response( + model="gpt-4o", + raw_response=raw_response, + model_response=model_response, + logging_obj=logging_obj, + request_data={"model": "gpt-4o"}, + messages=[{"role": "user", "content": "test"}], + optional_params={}, + litellm_params={}, + encoding=Mock(), + ) + + # Should have exactly ONE choice + assert len(result.choices) == 1, f"Expected 1 choice, got {len(result.choices)}" + + choice = result.choices[0] + assert choice.index == 0 + assert choice.finish_reason == "tool_calls" + + # That one choice should have ALL THREE tool calls + tool_calls = choice.message.tool_calls + assert tool_calls is not None, "tool_calls should not be None" + assert len(tool_calls) == 3, f"Expected 3 tool_calls, got {len(tool_calls)}" + + # Verify each tool call + assert tool_calls[0]["id"] == "call_paris" + assert tool_calls[0]["function"]["name"] == "get_weather" + + assert tool_calls[1]["id"] == "call_tokyo" + assert tool_calls[1]["function"]["name"] == "get_weather" + + assert tool_calls[2]["id"] == "call_horoscope" + assert tool_calls[2]["function"]["name"] == "get_horoscope" + + print("✓ Multiple tool calls are correctly grouped in a single choice") diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py index 586ab433502..e45db8df106 100644 --- a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py +++ b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py @@ -46,7 +46,7 @@ class TestCloudZeroHourlyExport: { "team_id": ["a3d6b0bb-098f-4260-81d6-fabae695b622"], "key_alias": ["key_1"], - "token": ["c1465c9a821f420927b3d81972323fb516745bc93a4a54ceca0ce6ddf6100c39"], + "token": ["sk-test-cloudzero-token-010"], } ) diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py b/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py index 70e97381082..5389cdf7377 100644 --- a/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py +++ b/tests/test_litellm/integrations/langfuse/test_langfuse_prompt_management.py @@ -27,3 +27,24 @@ class TestLangfusePromptManagement: mock_get_prompt_from_id.assert_called_once() assert mock_get_prompt_from_id.call_args.kwargs["prompt_version"] == 4 + + def test_log_failure_event_runs_async_logger(self): + langfuse_prompt_management = LangfusePromptManagement() + with patch( + "litellm.integrations.langfuse.langfuse_prompt_management.run_async_function" + ) as mock_run_async: + kwargs = {"standard_callback_dynamic_params": {}} + start_time, end_time = 1, 2 + + langfuse_prompt_management.log_failure_event( + kwargs=kwargs, + response_obj=None, + start_time=start_time, + end_time=end_time, + ) + + mock_run_async.assert_called_once() + assert ( + mock_run_async.call_args[0][0] + == langfuse_prompt_management.async_log_failure_event + ) diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 21206ec9482..a719d102a7c 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -383,3 +383,156 @@ class TestGuardrailLoggingAggregation: assert isinstance(info, list) assert len(info) == 2 assert info[1]["guardrail_name"] == "test_guardrail" + + +class TestCustomGuardrailPassthroughSupport: + """Tests for passthrough endpoint guardrail support - Issue fixes.""" + + @pytest.mark.asyncio + async def test_async_post_call_success_deployment_hook_with_httpx_response(self): + """ + Test that async_post_call_success_deployment_hook handles raw httpx.Response objects + from passthrough endpoints without crashing with TypeError. + + This tests Fix #3: TypeError: TypedDict does not support instance and class checks + """ + import httpx + + custom_guardrail = CustomGuardrail() + + # Mock the async_post_call_success_hook to return None (guardrail didn't modify response) + custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) + + # Create a mock httpx.Response object (typical passthrough response) + mock_response = AsyncMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.text = "Mock response" + + request_data = { + "guardrails": ["test_guardrail"], + "user_api_key_user_id": "test_user", + "user_api_key_team_id": "test_team", + "user_api_key_end_user_id": "test_end_user", + "user_api_key_hash": "test_hash", + "user_api_key_request_route": "passthrough_route", + } + + # This should not raise TypeError: TypedDict does not support instance and class checks + result = await custom_guardrail.async_post_call_success_deployment_hook( + request_data=request_data, + response=mock_response, + call_type=CallTypes.allm_passthrough_route, + ) + + # When result is None, should return the original response + assert result == mock_response + + @pytest.mark.asyncio + async def test_async_post_call_success_deployment_hook_with_none_call_type(self): + """ + Test that async_post_call_success_deployment_hook handles None call_type gracefully. + + This ensures that even if call_type is None (before fix #1), the guardrail doesn't crash. + """ + custom_guardrail = CustomGuardrail() + + # Mock the async_post_call_success_hook to return None + custom_guardrail.async_post_call_success_hook = AsyncMock(return_value=None) + + mock_response = AsyncMock() + + request_data = { + "guardrails": ["test_guardrail"], + "user_api_key_user_id": "test_user", + } + + # Call with None call_type - should not crash + result = await custom_guardrail.async_post_call_success_deployment_hook( + request_data=request_data, + response=mock_response, + call_type=None, + ) + + # Should return the original response when result is None + assert result == mock_response + + def test_is_valid_response_type_with_none(self): + """ + Test _is_valid_response_type helper method correctly identifies None as invalid. + + This is part of Fix #3: Safely handling TypedDict types that don't support isinstance checks. + """ + custom_guardrail = CustomGuardrail() + + # None should be invalid + assert custom_guardrail._is_valid_response_type(None) is False + + def test_is_valid_response_type_with_typeddict_error(self): + """ + Test _is_valid_response_type gracefully handles TypeError from TypedDict. + + This tests Fix #3: When isinstance() is called with TypedDict types, it raises TypeError. + The method should catch this and allow the response through. + """ + from litellm.types.utils import ModelResponse + + custom_guardrail = CustomGuardrail() + + # Create a valid LiteLLM response object + response = ModelResponse( + id="test-id", + choices=[], + created=0, + model="test-model", + object="chat.completion", + ) + + # This should return True (it's a valid response type or TypeError is caught) + result = custom_guardrail._is_valid_response_type(response) + assert result is True + + +class TestPassthroughCallTypeHandling: + """Tests for passthrough call type handling in common_request_processing.""" + + def test_get_pre_call_type_with_allm_passthrough_route(self): + """ + Test that _get_pre_call_type correctly maps allm_passthrough_route. + + This tests Fix #1: allm_passthrough_route was not being handled, causing call_type to be None. + """ + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + # Test the mapping + result = ProxyBaseLLMRequestProcessing._get_pre_call_type( + route_type="allm_passthrough_route" + ) + + # Should return allm_passthrough_route, not None + assert result == "allm_passthrough_route" + + def test_get_pre_call_type_preserves_standard_mappings(self): + """ + Test that _get_pre_call_type still correctly maps standard route types. + + Ensures Fix #1 didn't break existing functionality. + """ + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + # Test standard mappings are preserved + assert ( + ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="acompletion") + == "completion" + ) + assert ( + ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="aembedding") + == "embeddings" + ) + assert ( + ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="aresponses") + == "responses" + ) diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/test_litellm/integrations/test_responses_background_cost.py new file mode 100644 index 00000000000..6f1e7e96103 --- /dev/null +++ b/tests/test_litellm/integrations/test_responses_background_cost.py @@ -0,0 +1,513 @@ +""" +Integration tests for responses API background cost tracking +""" + +import asyncio +import os +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, Mock, patch + +import pytest + +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + +class TestResponsesBackgroundCostTracking: + """Integration tests for responses API background cost tracking""" + + @pytest.fixture + def mock_managed_files_obj(self): + """Create a mock managed files object""" + managed_files = MagicMock() + managed_files.store_unified_object_id = AsyncMock() + return managed_files + + @pytest.fixture + def mock_proxy_logging_obj(self, mock_managed_files_obj): + """Create a mock proxy logging object""" + logging_obj = MagicMock() + logging_obj.get_proxy_hook = MagicMock(return_value=mock_managed_files_obj) + return logging_obj + + @pytest.fixture + def mock_llm_router(self): + """Create a mock LLM router""" + router = MagicMock() + return router + + @pytest.mark.asyncio + async def test_store_response_in_managed_objects_table( + self, mock_managed_files_obj, mock_proxy_logging_obj, mock_llm_router + ): + """Test that background responses are stored in managed objects table""" + # Create a mock response with queued status and hidden params + response = ResponsesAPIResponse( + id="resp_bGl0ZWxsbTpjdXN0b21fbGxtX3Byb3ZpZGVyOm9wZW5haTttb2RlbF9pZDpncHQtNDtsbGxfcmVzcG9uc2VfaWQ6cmVzcF8xMjM", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Add hidden params with model_id (simulating what base_process_llm_request does) + response._hidden_params = { + "model_id": "model-deployment-id-123" + } + + # Mock request data + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + # Mock user_api_key_dict + user_api_key_dict = MagicMock() + user_api_key_dict.user_id = "test-user" + + # Simulate the storage logic from endpoints.py + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + # Get model_id from hidden params + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if model_id: + # Store in managed objects table using response.id directly + await mock_managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + + # Verify store_unified_object_id was called + mock_managed_files_obj.store_unified_object_id.assert_called_once() + call_args = mock_managed_files_obj.store_unified_object_id.call_args + + # Verify the arguments - unified_object_id should be response.id + assert call_args[1]["unified_object_id"] == response.id + assert call_args[1]["model_object_id"] == response.id + assert call_args[1]["file_purpose"] == "response" + assert call_args[1]["user_api_key_dict"] == user_api_key_dict + + @pytest.mark.asyncio + async def test_no_storage_for_non_background_requests( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that non-background requests are not stored""" + # Create a mock response + response = ResponsesAPIResponse( + id="resp_456", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + # Mock request data without background flag + data = { + "model": "gpt-4", + "input": "Test input", + "background": False, + } + + # Simulate the storage logic + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + await mock_managed_files_obj.store_unified_object_id() + + # Verify store_unified_object_id was NOT called + mock_managed_files_obj.store_unified_object_id.assert_not_called() + + @pytest.mark.asyncio + async def test_no_storage_for_completed_responses( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that completed responses are not stored""" + # Create a mock response with completed status + response = ResponsesAPIResponse( + id="resp_789", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + # Mock request data with background flag + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + # Simulate the storage logic + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + await mock_managed_files_obj.store_unified_object_id() + + # Verify store_unified_object_id was NOT called (status is completed) + mock_managed_files_obj.store_unified_object_id.assert_not_called() + + @pytest.mark.asyncio + async def test_no_storage_without_model_id( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that responses without model_id in hidden params are not stored""" + # Create a mock response without hidden params + response = ResponsesAPIResponse( + id="resp_no_model", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock request data with background flag + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + user_api_key_dict = MagicMock() + + # Simulate the storage logic + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if model_id: # This will be False + await mock_managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + + # Verify store_unified_object_id was NOT called (no model_id) + mock_managed_files_obj.store_unified_object_id.assert_not_called() + + @pytest.mark.asyncio + async def test_error_handling_in_storage( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that errors during storage are handled gracefully""" + # Mock store_unified_object_id to raise an exception + mock_managed_files_obj.store_unified_object_id = AsyncMock( + side_effect=Exception("Database error") + ) + + response = ResponsesAPIResponse( + id="resp_error", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + response._hidden_params = {"model_id": "test-model-id"} + + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + user_api_key_dict = MagicMock() + user_api_key_dict.user_id = "test-user" + + # Try to store - should not raise (error is caught in endpoints.py) + try: + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if model_id: + await mock_managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + except Exception: + # Exception should be caught and logged, not raised + pass + + # Verify the method was called (even though it raised) + assert mock_managed_files_obj.store_unified_object_id.called + + +class TestCheckResponsesCost: + """Tests for the CheckResponsesCost polling class""" + + @pytest.fixture + def mock_prisma_client(self): + """Create a mock Prisma client""" + client = MagicMock() + client.db = MagicMock() + client.db.litellm_managedobjecttable = MagicMock() + return client + + @pytest.fixture + def mock_proxy_logging_obj(self): + """Create a mock proxy logging object""" + return MagicMock() + + @pytest.fixture + def mock_llm_router(self): + """Create a mock LLM router""" + return MagicMock() + + @pytest.mark.asyncio + async def test_check_responses_cost_initialization( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test CheckResponsesCost initialization""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + assert checker.proxy_logging_obj == mock_proxy_logging_obj + assert checker.prisma_client == mock_prisma_client + assert checker.llm_router == mock_llm_router + + @pytest.mark.asyncio + async def test_check_responses_cost_no_jobs( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling when there are no jobs""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Mock find_many to return empty list + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[] + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + # Should not raise any errors + await checker.check_responses_cost() + + # Verify find_many was called with correct parameters + mock_prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( + where={ + "status": {"in": ["queued", "in_progress"]}, + "file_purpose": "response", + } + ) + + @pytest.mark.asyncio + async def test_check_responses_cost_with_completed_job( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling with a completed job""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-123" + mock_job.unified_object_id = "resp_test_id" + mock_job.created_by = "test-user" + + # Mock find_many to return the job + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Create a completed response + completed_response = ResponsesAPIResponse( + id="resp_test_id", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + # Mock litellm.aget_responses to return completed response + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = completed_response + + await checker.check_responses_cost() + + # Verify update_many was called to mark job as completed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = ( + mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + ) + assert call_args[1]["where"]["id"]["in"] == ["job-123"] + assert call_args[1]["data"]["status"] == "completed" + + @pytest.mark.asyncio + async def test_check_responses_cost_with_failed_job( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling with a failed job""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-456" + mock_job.unified_object_id = "resp_failed" + mock_job.created_by = "test-user" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Create a failed response + failed_response = ResponsesAPIResponse( + id="resp_failed", + object="response", + status="failed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = failed_response + + await checker.check_responses_cost() + + # Verify job was marked as completed even though it failed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_in_progress_job( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling with a job still in progress""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-789" + mock_job.unified_object_id = "resp_in_progress" + mock_job.created_by = "test-user" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Create an in-progress response + in_progress_response = ResponsesAPIResponse( + id="resp_in_progress", + object="response", + status="in_progress", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = in_progress_response + + await checker.check_responses_cost() + + # Verify update_many was NOT called (job still in progress) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_error_handling( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test that errors when querying responses are handled gracefully""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-error" + mock_job.unified_object_id = "resp_error" + mock_job.created_by = "test-user" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + # Mock litellm.aget_responses to raise an exception + with patch( + "litellm.aget_responses", + new_callable=AsyncMock, + side_effect=Exception("API error"), + ): + # Should not raise - errors are caught and logged + await checker.check_responses_cost() + + # Verify update_many was NOT called (error occurred) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() diff --git a/tests/test_litellm/llms/bedrock/image/test_amazon_nova_canvas_transformation.py b/tests/test_litellm/llms/bedrock/image/test_amazon_nova_canvas_transformation.py index 0dd0b80f36f..122d3e44364 100644 --- a/tests/test_litellm/llms/bedrock/image/test_amazon_nova_canvas_transformation.py +++ b/tests/test_litellm/llms/bedrock/image/test_amazon_nova_canvas_transformation.py @@ -1,5 +1,5 @@ import pytest -from litellm.llms.bedrock.image.amazon_nova_canvas_transformation import AmazonNovaCanvasConfig +from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import AmazonNovaCanvasConfig from litellm.types.utils import ImageResponse def test_transform_request_body_text_to_image(): diff --git a/tests/test_litellm/llms/bedrock/image/test_amazon_stability3_transformation.py b/tests/test_litellm/llms/bedrock/image/test_amazon_stability3_transformation.py index 1cf1747b8c7..a758202d74f 100644 --- a/tests/test_litellm/llms/bedrock/image/test_amazon_stability3_transformation.py +++ b/tests/test_litellm/llms/bedrock/image/test_amazon_stability3_transformation.py @@ -10,7 +10,7 @@ sys.path.insert( ) # Adds the parent directory to the system path from unittest.mock import MagicMock, patch -from litellm.llms.bedrock.image.amazon_stability3_transformation import ( +from litellm.llms.bedrock.image_generation.amazon_stability3_transformation import ( AmazonStability3Config, ) diff --git a/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py b/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py index b348c1193c7..5e0b3995470 100644 --- a/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py +++ b/tests/test_litellm/llms/bedrock/image/test_bedrock_image_bearer_token.py @@ -23,7 +23,7 @@ class TestBedrockImageGeneration: model = "bedrock/stability.sd3-large-v1:0" prompt = "A cute baby sea otter" - with patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: + with patch("litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: # Setup mock response mock_image_response_obj = litellm.ImageResponse() mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] @@ -55,7 +55,7 @@ class TestBedrockImageGeneration: # Mock the environment variable with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": test_api_key}), \ - patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: + patch("litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: mock_image_response_obj = litellm.ImageResponse() mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] @@ -85,7 +85,7 @@ class TestBedrockImageGeneration: model = "bedrock/stability.sd3-large-v1:0" prompt = "A cute baby sea otter" - with patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.async_image_generation") as mock_async_bedrock_image_gen: + with patch("litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.async_image_generation") as mock_async_bedrock_image_gen: mock_image_response_obj = litellm.ImageResponse() mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] mock_async_bedrock_image_gen.return_value = mock_image_response_obj @@ -114,7 +114,7 @@ class TestBedrockImageGeneration: model = "bedrock/stability.sd3-large-v1:0" prompt = "A cute baby sea otter" - with patch("litellm.llms.bedrock.image.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: + with patch("litellm.llms.bedrock.image_generation.image_handler.BedrockImageGeneration.image_generation") as mock_bedrock_image_gen: mock_image_response_obj = litellm.ImageResponse() mock_image_response_obj.data = [{"url": "https://example.com/image.jpg"}] mock_bedrock_image_gen.return_value = mock_image_response_obj diff --git a/tests/test_litellm/llms/bedrock/image/test_bedrock_image_prepare_request.py b/tests/test_litellm/llms/bedrock/image/test_bedrock_image_prepare_request.py index 22dc0cc8a48..6c56ccc1ef7 100644 --- a/tests/test_litellm/llms/bedrock/image/test_bedrock_image_prepare_request.py +++ b/tests/test_litellm/llms/bedrock/image/test_bedrock_image_prepare_request.py @@ -1,6 +1,6 @@ from unittest.mock import patch, MagicMock -from litellm.llms.bedrock.image.image_handler import BedrockImageGeneration +from litellm.llms.bedrock.image_generation.image_handler import BedrockImageGeneration def test_bedrock_image_prepare_request_with_arn() -> None: dummy_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abcdefghi123" diff --git a/tests/test_litellm/llms/vertex_ai/test_bge_embedding.py b/tests/test_litellm/llms/vertex_ai/test_bge_embedding.py index 4a06e9ea1aa..1f0f3346c2a 100644 --- a/tests/test_litellm/llms/vertex_ai/test_bge_embedding.py +++ b/tests/test_litellm/llms/vertex_ai/test_bge_embedding.py @@ -193,7 +193,7 @@ def test_vertex_ai_bge_psc_endpoint_url_construction(): client = HTTPHandler() def mock_auth_token(*args, **kwargs): - return "fake-token", "gen-lang-client-0682925754" + return "test-token-123", "test-gcp-project-id-123" with patch.object(client, "post") as mock_post, patch( "litellm.llms.vertex_ai.vertex_embeddings.embedding_handler.VertexEmbedding._ensure_access_token", @@ -212,7 +212,7 @@ def test_vertex_ai_bge_psc_endpoint_url_construction(): model="vertex_ai/bge/378943383978115072", input=["The food was delicious and the waiter.."], api_base="http://10.128.16.2", - vertex_project="gen-lang-client-0682925754", + vertex_project="test-gcp-project-id-123", vertex_location="us-central1", client=client, use_psc_endpoint_format=True # Enable PSC endpoint format for this test @@ -239,7 +239,7 @@ def test_vertex_ai_bge_psc_endpoint_url_construction(): print("="*50 + "\n") # Verify the URL is constructed correctly - expected_url = "http://10.128.16.2/v1/projects/gen-lang-client-0682925754/locations/us-central1/endpoints/378943383978115072:predict" + expected_url = "http://10.128.16.2/v1/projects/test-gcp-project-id-123/locations/us-central1/endpoints/378943383978115072:predict" assert api_url_called == expected_url, f"Expected URL: {expected_url}, Got: {api_url_called}" # Verify bge/ prefix is NOT in the URL diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index a5eee9e37b1..f850b53e12b 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -984,6 +984,7 @@ async def test_vertex_ai_token_counter_routes_partner_models(): to the partner models token counter instead of the Gemini token counter. """ from unittest.mock import AsyncMock, patch + from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter from litellm.types.utils import TokenCountResponse @@ -1027,6 +1028,7 @@ async def test_vertex_ai_token_counter_routes_gemini_models(): to the Gemini token counter (not partner models). """ from unittest.mock import AsyncMock, patch + from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter from litellm.types.utils import TokenCountResponse @@ -1124,3 +1126,73 @@ def test_vertex_ai_moonshot_uses_openai_handler(): assert VertexAIPartnerModels.should_use_openai_handler( "moonshotai/kimi-k2-thinking-maas" ) + + +def test_build_vertex_schema_empty_properties(): + """ + Test _build_vertex_schema handles empty properties objects correctly. + + This test verifies the fix for the issue where Gemini rejects schemas + with empty properties objects like {"properties": {}, "type": "object"}. + + Error from Gemini: "GenerateContentRequest.generation_config.response_schema + .properties[\"action\"].items.any_of[0].properties[\"go_back\"].properties: + should be non-empty for OBJECT type" + + The fix removes empty properties objects and their associated type/required fields. + """ + from litellm.llms.vertex_ai.common_utils import _build_vertex_schema + + # Input: Schema with empty properties (the problematic case from real request) + input_schema = { + "properties": { + "action": { + "description": "List of actions to execute", + "items": { + "anyOf": [ + { + "properties": { + "go_back": { + "properties": {}, + "type": "object", + "additionalProperties": False, + "description": "Go back", + "required": [] + } + }, + "required": ["go_back"], + "type": "object", + "additionalProperties": False + } + ] + }, + "type": "array" + } + }, + "type": "object", + "additionalProperties": False + } + + # Apply the transformation + result = _build_vertex_schema(input_schema) + + # Verify the transformation removed empty properties + # Navigate to the go_back schema + go_back_schema = result["properties"]["action"]["items"]["anyOf"][0]["properties"]["go_back"] + + # Verify empty properties was removed + assert "properties" not in go_back_schema, "Empty properties should be removed" + + # Verify type was also removed (since object without properties is invalid in Gemini) + assert "type" not in go_back_schema, "Type should be removed when properties is empty" + + # Verify required was also removed + assert "required" not in go_back_schema, "Required should be removed when properties is empty" + + # Verify description is preserved + assert go_back_schema.get("description") == "Go back", "Description should be preserved" + + # Verify parent schema still has proper structure + parent_schema = result["properties"]["action"]["items"]["anyOf"][0] + assert parent_schema["type"] == "object", "Parent schema should still have object type" + assert "go_back" in parent_schema["properties"], "go_back should still be in parent properties" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 21782f42189..e1e4b3a8b6d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -332,7 +332,7 @@ class TestMCPRequestHandler: async def mock_user_api_key_auth(api_key, request): return UserAPIKeyAuth( token=( - "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + "test-token-sha256-empty-hash" if api_key else None ), @@ -691,7 +691,7 @@ class TestMCPCustomHeaderName: # Create an async mock for user_api_key_auth async def mock_user_api_key_auth(api_key, request): return UserAPIKeyAuth( - token="e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", + token="test-token-sha256-empty-hash", api_key=api_key, user_id="test-user-id", team_id="test-team-id", @@ -866,7 +866,7 @@ class TestMCPAccessGroupsE2E: # Create an async mock for user_api_key_auth async def mock_user_api_key_auth(api_key, request): return UserAPIKeyAuth( - token="e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", + token="test-token-sha256-empty-hash", api_key=api_key, user_id="test-user-id", team_id="test-team-id", @@ -917,7 +917,7 @@ class TestMCPAccessGroupsE2E: # Create an async mock for user_api_key_auth async def mock_user_api_key_auth(api_key, request): return UserAPIKeyAuth( - token="e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", + token="test-token-sha256-empty-hash", api_key=api_key, user_id="test-user-id", team_id="test-team-id", diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index eeae0ece02c..f3de89d6d6c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -43,8 +43,8 @@ def mock_user_api_key_dict(): team_id="test-team", team_alias=None, user_role=None, - api_key="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", - token="88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + api_key="a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456", + token="a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456", permissions={}, models=[], spend=0.0, @@ -71,7 +71,7 @@ def mock_request_data_input(): ], "litellm_call_id": "test-call-id", "metadata": { - "user_api_key_hash": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "user_api_key_hash": "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456", "user_api_key_user_id": "default_user_id", "user_api_key_user_email": "test@example.com", "user_api_key_team_id": "test-team", @@ -197,7 +197,7 @@ class TestMetadataExtraction: # Verify metadata was extracted from request_data["metadata"] assert ( request_metadata["user_api_key_hash"] - == "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b" + == "a1b2c3d4e5f6789012345678901234567890abcdef1234567890abcdef123456" ) assert request_metadata["user_api_key_user_id"] == "default_user_id" assert request_metadata["user_api_key_user_email"] == "test@example.com" diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 249612cd58a..194c1d87f6e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -20,6 +20,7 @@ from litellm.proxy._types import ( LiteLLM_TeamTableCachedObj, LiteLLM_VerificationToken, LitellmUserRoles, + Member, ProxyException, UpdateKeyRequest, ) @@ -32,6 +33,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( _persist_deleted_verification_tokens, _save_deleted_verification_token_records, _transform_verification_tokens_to_deleted_records, + can_modify_verification_token, check_org_key_model_specific_limits, check_team_key_model_specific_limits, delete_verification_tokens, @@ -2973,3 +2975,706 @@ async def test_delete_key_fn_persists_deleted_keys(monkeypatch): ) assert result["deleted_keys"] == ["sk-token-1"] + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_proxy_admin_team_key(monkeypatch): + """Test that team admin can delete team keys from their own team.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="other-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="team-admin-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-123", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="team-admin-user", role="admin"), + Member(user_id="other-user", role="user"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_team_admin_different_team(monkeypatch): + """Test that team admin cannot delete team keys from a different team.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="other-user", + team_id="test-team-456", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="team-admin-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-456", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="different-admin", role="admin"), + Member(user_id="other-user", role="user"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_key_owner_team_key(monkeypatch): + """Test that key owner can delete their own team key.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="key-owner-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-123", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="key-owner-user", role="user"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_key_owner_personal_key(monkeypatch): + """Test that key owner can delete their own personal key.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="key-owner-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_other_user_team_key(monkeypatch): + """Test that other user cannot delete team keys they don't own and aren't admin for.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="other-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-123", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="key-owner-user", role="user"), + Member(user_id="other-user", role="user"), + Member(user_id="team-admin-user", role="admin"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_other_user_personal_key(monkeypatch): + """Test that other user cannot delete personal keys they don't own.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="other-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_team_key_no_team_found(monkeypatch): + """Test that deletion fails when team is not found in database.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id="non-existent-team", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="key-owner-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return None + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_delete_verification_token_personal_key_no_user_id(monkeypatch): + """Test that deletion fails for personal key when key has no user_id.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id=None, + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="some-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + +@pytest.mark.asyncio +async def test_can_modify_verification_token_proxy_admin_team_key(monkeypatch): + """Test that proxy admin can modify any team key.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="other-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin-user", + api_key="sk-admin", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_proxy_admin_personal_key(monkeypatch): + """Test that proxy admin can modify any personal key.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="other-user", + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin-user", + api_key="sk-admin", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_team_admin_own_team(monkeypatch): + """Test that team admin can modify team keys from their own team.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="other-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="team-admin-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-123", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="team-admin-user", role="admin"), + Member(user_id="other-user", role="user"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_team_admin_different_team(monkeypatch): + """Test that team admin cannot modify team keys from a different team.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="other-user", + team_id="test-team-456", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="team-admin-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-456", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="different-admin", role="admin"), + Member(user_id="other-user", role="user"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_key_owner_team_key(monkeypatch): + """Test that key owner can modify their own team key.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="key-owner-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-123", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="key-owner-user", role="user"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_key_owner_personal_key(monkeypatch): + """Test that key owner can modify their own personal key.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="key-owner-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is True + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_other_user_team_key(monkeypatch): + """Test that other user cannot modify team keys they don't own and aren't admin for.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id="test-team-123", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="other-user", + api_key="sk-user", + ) + + team_table = LiteLLM_TeamTableCachedObj( + team_id="test-team-123", + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="key-owner-user", role="user"), + Member(user_id="other-user", role="user"), + Member(user_id="team-admin-user", role="admin"), + ], + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_other_user_personal_key(monkeypatch): + """Test that other user cannot modify personal keys they don't own.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="other-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_team_key_no_team_found(monkeypatch): + """Test that modification fails when team is not found in database.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id="key-owner-user", + team_id="non-existent-team", + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="key-owner-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + async def mock_get_team_object(*args, **kwargs): + return None + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False + + +@pytest.mark.asyncio +async def test_can_modify_verification_token_personal_key_no_user_id(monkeypatch): + """Test that modification fails for personal key when key has no user_id.""" + key_info = LiteLLM_VerificationToken( + token="test-token", + user_id=None, + team_id=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="some-user", + api_key="sk-user", + ) + + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + + result = await can_modify_verification_token( + key_info=key_info, + user_api_key_cache=mock_user_api_key_cache, + user_api_key_dict=user_api_key_dict, + prisma_client=mock_prisma_client, + ) + + assert result is False diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index a08fc2cba67..fa8157ff658 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -2050,7 +2050,7 @@ class TestProcessSSOJWTAccessToken: @pytest.fixture def sample_jwt_token(self): """Create a sample JWT token string""" - return "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.SflKxwRJSMeKKF2QT4fwpMeJf36POk6yJV_adQssw5c" + return "test-jwt-token-header.payload.signature" @pytest.fixture def sample_jwt_payload(self): diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 90d958e711d..5f03ef18171 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -180,7 +180,7 @@ class TestProxyInitializationHelpers: test_env = { "DATABASE_HOST": "localhost:5432", "DATABASE_USERNAME": "user@with+special", - "DATABASE_PASSWORD": "pass&word!@#$%", + "DATABASE_PASSWORD": "test-password-special-chars", "DATABASE_NAME": "db_name/test", } @@ -205,7 +205,7 @@ class TestProxyInitializationHelpers: database_url = f"postgresql://{database_username_enc}:{database_password_enc}@{database_host}/{database_name_enc}" # Assert the correct URL was constructed with properly escaped characters - expected_url = "postgresql://user%40with%2Bspecial:pass%26word%21%40%23%24%25@localhost:5432/db_name%2Ftest" + expected_url = "postgresql://user%40with%2Bspecial:test-password-special-chars@localhost:5432/db_name%2Ftest" assert database_url == expected_url # Test appending query parameters @@ -381,13 +381,13 @@ class TestProxyInitializationHelpers: test_env_special = { "DATABASE_HOST": "localhost:5432", "DATABASE_USERNAME": "user@with+special", - "DATABASE_PASSWORD": "pass&word!@#$%", + "DATABASE_PASSWORD": "test-password-special-chars", "DATABASE_NAME": "db_name/test", } with patch.dict(os.environ, test_env_special): result = construct_database_url_from_env_vars() - expected_url = "postgresql://user%40with%2Bspecial:pass%26word%21%40%23%24%25@localhost:5432/db_name%2Ftest" + expected_url = "postgresql://user%40with%2Bspecial:test-password-special-chars@localhost:5432/db_name%2Ftest" assert result == expected_url # Test without password (should still work) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 1f81026b537..6b8342968ad 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -559,7 +559,7 @@ async def test_aaaproxy_startup_master_key(mock_prisma, monkeypatch, tmp_path): assert master_key == test_master_key # Test Case 2: Master key from environment variable - test_env_master_key = "sk-67890" + test_env_master_key = "sk-test-67890" # Create empty config empty_config = {"general_settings": {}} diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index b98354032fe..f697ad9abb2 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -644,7 +644,7 @@ class TestIsAllowedToCallVectorStoreEndpoint: mock_request.method = "GET" mock_request.url.path = "/azure_ai/indexes/dall-e-4/docs/search" mock_user_api_key = UserAPIKeyAuth( - token="b637312ebffb9745321224644430ba9e4916a291c8281f293d21182c5e80bc5a", + token="sk-test-mock-token-404", key_name="sk-...plNQ", metadata={ "allowed_vector_store_indexes": [ diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py b/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py index bd6bab9d61e..b0a232a7bf4 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py @@ -27,7 +27,7 @@ async def test_get_chat_completion_message_history_for_previous_response_id(): { "request_id": "chatcmpl-935b8dad-fdc2-466e-a8ca-e26e5a8a21bb", "call_type": "aresponses", - "api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "api_key": "sk-test-mock-api-key-123", "spend": 0.004803, "total_tokens": 329, "prompt_tokens": 11, @@ -68,7 +68,7 @@ async def test_get_chat_completion_message_history_for_previous_response_id(): { "request_id": "chatcmpl-370760c9-39fa-4db7-b034-d1f8d933c935", "call_type": "aresponses", - "api_key": "88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b", + "api_key": "sk-test-mock-api-key-123", "spend": 0.010437, "total_tokens": 967, "prompt_tokens": 339, diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index c26801ac3f6..69e0f04e5e1 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -855,7 +855,7 @@ def test_azure_image_generation_cost_calculator(): ImageObject( b64_json=None, revised_prompt="A futuristic, techno-inspired green duck wearing cool modern sunglasses. The duck has a sleek, metallic appearance with glowing neon green accents, standing on a high-tech urban background with holographic billboards and illuminated city lights in the distance. The duck's feathers have a glossy, high-tech sheen, resembling a robotic design but still maintaining its avian features. The scene has a vibrant, cyberpunk aesthetic with a neon color palette.", - url="https://dalleprodsec.blob.core.windows.net/private/images/caa17dc4-357d-4257-8938-eeea9baa8d0a/generated_00.png?se=2025-10-31T00%3A47%3A59Z&sig=KHRjLz3vMahbw94JtxL02S6t2AueeRMaiqj4z35HKDM%3D&ske=2025-11-05T00%3A26%3A20Z&skoid=e52d5ed7-0657-4f62-bc12-7e5dbb260a96&sks=b&skt=2025-10-29T00%3A26%3A20Z&sktid=33e01921-4d64-4f8c-a055-5bdaffd5e33d&skv=2020-10-02&sp=r&spr=https&sr=b&sv=2020-10-02", + url="test-azure-blob-url-with-sas-token", ) ], output_format=None, diff --git a/tests/test_spend_logs.py b/tests/test_spend_logs.py index 80dd8c9bcca..8aec1d5cc60 100644 --- a/tests/test_spend_logs.py +++ b/tests/test_spend_logs.py @@ -198,7 +198,7 @@ async def get_predict_spend_logs(session): { "date": "2024-03-09", "spend": 200000, - "api_key": "f19bdeb945164278fc11c1020d8dfd70465bffd931ed3cb2e1efa6326225b8b7", + "api_key": "sk-test-mock-api-key-456", } ] } diff --git a/tests/vector_store_tests/rag/test_rag_openai.py b/tests/vector_store_tests/rag/test_rag_openai.py index d077ebe0cb6..a9cffa3776c 100644 --- a/tests/vector_store_tests/rag/test_rag_openai.py +++ b/tests/vector_store_tests/rag/test_rag_openai.py @@ -42,4 +42,110 @@ class TestRAGOpenAI(BaseRAGTest): return search_response return None + @pytest.mark.asyncio + async def test_rag_query_basic(self): + """Test basic RAG query flow.""" + import asyncio + + litellm._turn_on_debug() + + # First ingest a document + filename, unique_id = self.get_unique_filename("rag_query") + text_content = ( + f"LiteLLM is a unified interface for 100+ LLMs. ID: {unique_id}".encode() + ) + + ingest_response = await litellm.rag.aingest( + ingest_options=self.get_base_ingest_options(), + file_data=(filename, text_content, "text/plain"), + ) + + # Check if ingestion succeeded + if ingest_response["status"] != "completed": + pytest.fail( + f"Ingestion failed with status: {ingest_response['status']}, " + f"error: {ingest_response.get('error', 'Unknown')}" + ) + + vector_store_id = ingest_response["vector_store_id"] + assert vector_store_id, "vector_store_id should not be empty" + + # Wait for indexing + await asyncio.sleep(10) + + # Query with RAG + response = await litellm.rag.aquery( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "What is LiteLLM?"}], + retrieval_config={ + "vector_store_id": vector_store_id, + "custom_llm_provider": "openai", + "top_k": 5, + }, + ) + + print(f"RAG Query Response: {response}") + + assert response.choices[0].message.content + assert ( + "search_results" in response.choices[0].message.provider_specific_fields + ) + + @pytest.mark.asyncio + async def test_rag_query_with_rerank(self): + """Test RAG query with reranking.""" + import asyncio + + litellm._turn_on_debug() + + # First ingest a document + filename, unique_id = self.get_unique_filename("rag_query_rerank") + text_content = ( + f"LiteLLM is a unified interface for 100+ LLMs. ID: {unique_id}".encode() + ) + + ingest_response = await litellm.rag.aingest( + ingest_options=self.get_base_ingest_options(), + file_data=(filename, text_content, "text/plain"), + ) + + # Check if ingestion succeeded + if ingest_response["status"] != "completed": + pytest.fail( + f"Ingestion failed with status: {ingest_response['status']}, " + f"error: {ingest_response.get('error', 'Unknown')}" + ) + + vector_store_id = ingest_response["vector_store_id"] + assert vector_store_id, "vector_store_id should not be empty" + + # Wait for indexing + await asyncio.sleep(10) + + # Query with RAG and rerank + response = await litellm.rag.aquery( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "What is LiteLLM?"}], + retrieval_config={ + "vector_store_id": vector_store_id, + "custom_llm_provider": "openai", + "top_k": 5, + }, + rerank={ + "enabled": True, + "model": "cohere/rerank-english-v3.0", + "top_n": 3, + }, + ) + + print(f"RAG Query Response with Rerank: {response.model_dump_json(indent=4)}") + + assert response.choices[0].message.content + assert ( + "search_results" in response.choices[0].message.provider_specific_fields + ) + assert ( + "rerank_results" in response.choices[0].message.provider_specific_fields + ) + \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/cache_dashboard.tsx b/ui/litellm-dashboard/src/components/cache_dashboard.tsx index 38c0f1a8f41..65f02874cf1 100644 --- a/ui/litellm-dashboard/src/components/cache_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/cache_dashboard.tsx @@ -162,13 +162,13 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole /* Data looks like this - [{"api_key":"147dba2181f28914eea90eb484926c293cdcf7f5b5c9c3dd6a004d9e0f9fdb21","call_type":"acompletion","model":"llama3-8b-8192","total_rows":13,"cache_hit_true_rows":0}, - {"api_key":"8c23f021d0535c2e59abb7d83d0e03ccfb8db1b90e231ff082949d95df419e86","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b","call_type":"acompletion","model":"gpt-3.5-turbo","total_rows":19,"cache_hit_true_rows":0}, - {"api_key":"88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b","call_type":"aimage_generation","model":"","total_rows":3,"cache_hit_true_rows":0}, - {"api_key":"0ad4b3c03dcb6de0b5b8f761db798c6a8ae80be3fd1e2ea30c07ce6d5e3bf870","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"034224b36e9769bc50e2190634abc3f97cad789b17ca80ac43b82f46cd5579b3","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"4f9c71cce0a2bb9a0b62ce6f0ebb3245b682702a8851d26932fa7e3b8ebfc755","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, + [{"api_key":"sk-test-mock-key-001","call_type":"acompletion","model":"llama3-8b-8192","total_rows":13,"cache_hit_true_rows":0}, + {"api_key":"sk-test-mock-key-002","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, + {"api_key":"sk-test-mock-key-123","call_type":"acompletion","model":"gpt-3.5-turbo","total_rows":19,"cache_hit_true_rows":0}, + {"api_key":"sk-test-mock-key-123","call_type":"aimage_generation","model":"","total_rows":3,"cache_hit_true_rows":0}, + {"api_key":"sk-test-mock-key-003","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, + {"api_key":"sk-test-mock-key-004","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, + {"api_key":"sk-test-mock-key-005","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, */ // What data we need for bar chat