diff --git a/.circleci/config.yml b/.circleci/config.yml index 0bfbbcb4405..8ae399c5c5f 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1041,6 +1041,49 @@ jobs: paths: - llm_responses_api_coverage.xml - llm_responses_api_coverage + ocr_testing: + docker: + - image: cimg/python:3.11 + auth: + username: ${DOCKERHUB_USERNAME} + password: ${DOCKERHUB_PASSWORD} + working_directory: ~/project + + steps: + - checkout + - setup_google_dns + - run: + name: Install Dependencies + command: | + python -m pip install --upgrade pip + python -m pip install -r requirements.txt + pip install "pytest==7.3.1" + pip install "pytest-retry==1.6.3" + pip install "pytest-cov==5.0.0" + pip install "pytest-asyncio==0.21.1" + pip install "respx==0.22.0" + # Run pytest and generate JUnit XML report + - run: + name: Run tests + command: | + pwd + ls + python -m pytest -vv tests/ocr_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5 + no_output_timeout: 120m + - run: + name: Rename the coverage files + command: | + mv coverage.xml ocr_coverage.xml + mv .coverage ocr_coverage + + # Store test results + - store_test_results: + path: test-results + - persist_to_workspace: + root: . + paths: + - ocr_coverage.xml + - ocr_coverage litellm_mapped_tests: docker: - image: cimg/python:3.11 @@ -2741,7 +2784,7 @@ jobs: python -m venv venv . venv/bin/activate pip install coverage - coverage combine llm_translation_coverage llm_responses_api_coverage mcp_coverage logging_coverage litellm_router_coverage local_testing_coverage litellm_assistants_api_coverage auth_ui_unit_tests_coverage langfuse_coverage caching_coverage litellm_proxy_unit_tests_coverage image_gen_coverage pass_through_unit_tests_coverage batches_coverage litellm_security_tests_coverage guardrails_coverage + coverage combine llm_translation_coverage llm_responses_api_coverage ocr_coverage mcp_coverage logging_coverage litellm_router_coverage local_testing_coverage litellm_assistants_api_coverage auth_ui_unit_tests_coverage langfuse_coverage caching_coverage litellm_proxy_unit_tests_coverage image_gen_coverage pass_through_unit_tests_coverage batches_coverage litellm_security_tests_coverage guardrails_coverage coverage xml - codecov/upload: file: ./coverage.xml @@ -3289,6 +3332,12 @@ workflows: only: - main - /litellm_.*/ + - ocr_testing: + filters: + branches: + only: + - main + - /litellm_.*/ - litellm_mapped_enterprise_tests: filters: branches: @@ -3338,6 +3387,7 @@ workflows: - google_generate_content_endpoint_testing - guardrails_testing - llm_responses_api_testing + - ocr_testing - litellm_mapped_tests - litellm_mapped_enterprise_tests - batches_testing @@ -3400,6 +3450,7 @@ workflows: - mcp_testing - google_generate_content_endpoint_testing - llm_responses_api_testing + - ocr_testing - litellm_mapped_tests - litellm_mapped_enterprise_tests - batches_testing diff --git a/README.md b/README.md index c785ee82ffa..812b20e6986 100644 --- a/README.md +++ b/README.md @@ -347,6 +347,7 @@ curl 'http://0.0.0.0:4000/key/generate' \ | [Nebius AI Studio](https://docs.litellm.ai/docs/providers/nebius) | ✅ | ✅ | ✅ | ✅ | ✅ | | | [Heroku](https://docs.litellm.ai/docs/providers/heroku) | ✅ | ✅ | | | | | | [OVHCloud AI Endpoints](https://docs.litellm.ai/docs/providers/ovhcloud) | ✅ | ✅ | | | | | +| [CometAPI](https://docs.litellm.ai/docs/providers/cometapi) | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | [**Read the Docs**](https://docs.litellm.ai/docs/) diff --git a/cookbook/LiteLLM_CometAPI.ipynb b/cookbook/LiteLLM_CometAPI.ipynb new file mode 100644 index 00000000000..bdd916c5bfe --- /dev/null +++ b/cookbook/LiteLLM_CometAPI.ipynb @@ -0,0 +1,474 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": { + "id": "iFEmsVJI_2BR" + }, + "source": [ + "# LiteLLM CometAPI Cookbook" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "metadata": { + "id": "cBlUhCEP_xj4" + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Requirement already satisfied: litellm in /Users/xmx/.miniforge3/lib/python3.12/site-packages (1.78.2)\n", + "Requirement already satisfied: aiohttp>=3.10 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (3.11.18)\n", + "Requirement already satisfied: click in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (8.3.0)\n", + "Requirement already satisfied: fastuuid>=0.13.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (0.13.3)\n", + "Requirement already satisfied: httpx>=0.23.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (0.28.1)\n", + "Requirement already satisfied: importlib-metadata>=6.8.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (8.6.1)\n", + "Requirement already satisfied: jinja2<4.0.0,>=3.1.2 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (3.1.6)\n", + "Requirement already satisfied: jsonschema<5.0.0,>=4.22.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (4.25.1)\n", + "Requirement already satisfied: openai>=1.99.5 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (1.109.1)\n", + "Requirement already satisfied: pydantic<3.0.0,>=2.5.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (2.11.10)\n", + "Requirement already satisfied: python-dotenv>=0.2.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (1.1.1)\n", + "Requirement already satisfied: tiktoken>=0.7.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (0.12.0)\n", + "Requirement already satisfied: tokenizers in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from litellm) (0.22.1)\n", + "Requirement already satisfied: aiohappyeyeballs>=2.3.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (2.6.1)\n", + "Requirement already satisfied: aiosignal>=1.1.2 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (1.4.0)\n", + "Requirement already satisfied: attrs>=17.3.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (25.3.0)\n", + "Requirement already satisfied: frozenlist>=1.1.1 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (1.6.0)\n", + "Requirement already satisfied: multidict<7.0,>=4.5 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (6.6.3)\n", + "Requirement already satisfied: propcache>=0.2.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (0.3.1)\n", + "Requirement already satisfied: yarl<2.0,>=1.17.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from aiohttp>=3.10->litellm) (1.20.0)\n", + "Requirement already satisfied: anyio in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from httpx>=0.23.0->litellm) (4.11.0)\n", + "Requirement already satisfied: certifi in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from httpx>=0.23.0->litellm) (2025.10.5)\n", + "Requirement already satisfied: httpcore==1.* in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from httpx>=0.23.0->litellm) (1.0.9)\n", + "Requirement already satisfied: idna in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from httpx>=0.23.0->litellm) (3.10)\n", + "Requirement already satisfied: h11>=0.16 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from httpcore==1.*->httpx>=0.23.0->litellm) (0.16.0)\n", + "Requirement already satisfied: zipp>=3.20 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from importlib-metadata>=6.8.0->litellm) (3.21.0)\n", + "Requirement already satisfied: MarkupSafe>=2.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from jinja2<4.0.0,>=3.1.2->litellm) (3.0.3)\n", + "Requirement already satisfied: jsonschema-specifications>=2023.03.6 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from jsonschema<5.0.0,>=4.22.0->litellm) (2025.9.1)\n", + "Requirement already satisfied: referencing>=0.28.4 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from jsonschema<5.0.0,>=4.22.0->litellm) (0.36.2)\n", + "Requirement already satisfied: rpds-py>=0.7.1 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from jsonschema<5.0.0,>=4.22.0->litellm) (0.27.1)\n", + "Requirement already satisfied: distro<2,>=1.7.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from openai>=1.99.5->litellm) (1.9.0)\n", + "Requirement already satisfied: jiter<1,>=0.4.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from openai>=1.99.5->litellm) (0.11.0)\n", + "Requirement already satisfied: sniffio in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from openai>=1.99.5->litellm) (1.3.1)\n", + "Requirement already satisfied: tqdm>4 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from openai>=1.99.5->litellm) (4.67.1)\n", + "Requirement already satisfied: typing-extensions<5,>=4.11 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from openai>=1.99.5->litellm) (4.15.0)\n", + "Requirement already satisfied: annotated-types>=0.6.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from pydantic<3.0.0,>=2.5.0->litellm) (0.7.0)\n", + "Requirement already satisfied: pydantic-core==2.33.2 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from pydantic<3.0.0,>=2.5.0->litellm) (2.33.2)\n", + "Requirement already satisfied: typing-inspection>=0.4.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from pydantic<3.0.0,>=2.5.0->litellm) (0.4.2)\n", + "Requirement already satisfied: regex>=2022.1.18 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from tiktoken>=0.7.0->litellm) (2025.9.18)\n", + "Requirement already satisfied: requests>=2.26.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from tiktoken>=0.7.0->litellm) (2.32.2)\n", + "Requirement already satisfied: huggingface-hub<2.0,>=0.16.4 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from tokenizers->litellm) (0.25.2)\n", + "Requirement already satisfied: filelock in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from huggingface-hub<2.0,>=0.16.4->tokenizers->litellm) (3.15.4)\n", + "Requirement already satisfied: fsspec>=2023.5.0 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from huggingface-hub<2.0,>=0.16.4->tokenizers->litellm) (2025.9.0)\n", + "Requirement already satisfied: packaging>=20.9 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from huggingface-hub<2.0,>=0.16.4->tokenizers->litellm) (25.0)\n", + "Requirement already satisfied: pyyaml>=5.1 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from huggingface-hub<2.0,>=0.16.4->tokenizers->litellm) (6.0.3)\n", + "Requirement already satisfied: charset-normalizer<4,>=2 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from requests>=2.26.0->tiktoken>=0.7.0->litellm) (3.4.0)\n", + "Requirement already satisfied: urllib3<3,>=1.21.1 in /Users/xmx/.miniforge3/lib/python3.12/site-packages (from requests>=2.26.0->tiktoken>=0.7.0->litellm) (1.26.20)\n" + ] + } + ], + "source": [ + "!pip install litellm" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Completion" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "p-MQqWOT_1a7" + }, + "outputs": [], + "source": [ + "import os\n", + "\n", + "os.environ['COMETAPI_KEY'] = \"Your_CometAPI_Key_Here\"\n", + "api_key = os.getenv('COMETAPI_KEY')" + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "Ze8JqMqWAARO", + "outputId": "64f3e836-69fa-4f8e-fb35-088a913bbe98" + }, + "outputs": [ + { + "data": { + "text/plain": [ + "ModelResponse(id='msg_017L3DDDit8AkEgHRe2DQBc9', created=1760589916, model='claude-sonnet-4-5-20250929', object='chat.completion', system_fingerprint=None, choices=[Choices(finish_reason='stop', index=0, message=Message(content='I\\'ll create a simple Python script that says hi.\\n\\n\\nhello.py\\n#!/usr/bin/env python3\\n\"\"\"\\nA simple script that says hi!\\n\"\"\"\\n\\ndef say_hi(name=None):\\n \"\"\"Say hi to someone, or just say hi generally.\"\"\"\\n if name:\\n print(f\"Hi, {name}!\")\\n else:\\n print(\"Hi!\")\\n\\nif __name__ == \"__main__\":\\n # Say hi generally\\n say_hi()\\n \\n # Say hi to someone specific\\n say_hi(\"World\")\\n\\n\\n\\nI\\'ve created a simple Python script called `hello.py` that:\\n\\n1. Defines a `say_hi()` function that can optionally take a name parameter\\n2. Prints \"Hi!\" if no name is provided\\n3. Prints \"Hi, [name]!\" if a name is provided\\n4. Demonstrates both usages when run\\n\\nYou can run it with:\\n```bash\\npython hello.py\\n```\\n\\nThis will output:\\n```\\nHi!\\nHi, World!\\n```\\n\\nWould you like me to create versions in other programming languages, or modify this in any way?', role='assistant', tool_calls=None, function_call=None, provider_specific_fields=None), provider_specific_fields={})], usage=Usage(completion_tokens=290, prompt_tokens=26, total_tokens=316, completion_tokens_details=None, prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=None, cached_tokens=None, text_tokens=None, image_tokens=None, cached_tokens_details={})))" + ] + }, + "execution_count": 3, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "from litellm import completion\n", + "response = completion(\n", + " model=\"cometapi/claude-sonnet-4-5-20250929\",\n", + " messages=[{\"role\": \"user\", \"content\": \"write code for saying hi\"}]\n", + ")\n", + "response" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "-LnhELrnAM_J", + "outputId": "d51c7ab7-d761-4bd1-f849-1534d9df4cd0" + }, + "outputs": [ + { + "data": { + "text/plain": [ + "ModelResponse(id='chatcmpl-CRA9Uo6nsQ9C7kMJv1J4kyNDFJym7', created=1760589916, model='gpt-5-chat-latest', object='chat.completion', system_fingerprint='fp_2da73a467a', choices=[Choices(finish_reason='stop', index=0, message=Message(content='Sure! I can help you write a simple code that prints out \"Hi\" in different programming languages. \\n\\nHere’s an example in **Python**:\\n\\n```python\\n# Simple Python program to say \"Hi\"\\nprint(\"Hi\")\\n```\\n\\nExample in **JavaScript**:\\n\\n```javascript\\n// Simple JavaScript program to say \"Hi\"\\nconsole.log(\"Hi\");\\n```\\n\\nExample in **C**:\\n\\n```c\\n#include \\n\\nint main() {\\n printf(\"Hi\\\\n\");\\n return 0;\\n}\\n```\\n\\nExample in **Java**:\\n\\n```java\\npublic class SayHi {\\n public static void main(String[] args) {\\n System.out.println(\"Hi\");\\n }\\n}\\n```\\n\\nWhich language would you like me to focus on, or do you want me to make it interactive so the program greets the user by name?', role='assistant', tool_calls=None, function_call=None, provider_specific_fields={'refusal': None}, annotations=[]), provider_specific_fields={})], usage=Usage(completion_tokens=174, prompt_tokens=12, total_tokens=186, completion_tokens_details=CompletionTokensDetailsWrapper(accepted_prediction_tokens=0, audio_tokens=0, reasoning_tokens=0, rejected_prediction_tokens=0, text_tokens=None), prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=0, cached_tokens=0, text_tokens=None, image_tokens=None)))" + ] + }, + "execution_count": 4, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "response = completion(\n", + " model=\"cometapi/gpt-5-chat-latest\",\n", + " messages=[{\"role\": \"user\", \"content\": \"write code for saying hi\"}]\n", + ")\n", + "response" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "dJBOUYdwCEn1", + "outputId": "ffa18679-ec15-4dad-fe2b-68665cdf36b0" + }, + "outputs": [ + { + "data": { + "text/plain": [ + "ModelResponse(id='02176058998406949c23b3bf3d52941de23f13565a086747738f0', created=1760589991, model='deepseek-v3.2-exp', object='chat.completion', system_fingerprint=None, choices=[Choices(finish_reason='stop', index=0, message=Message(content='Here are several ways to say \"hi\" in different programming languages:\\n\\n## Python\\n```python\\nprint(\"Hi!\")\\n```\\n\\n## JavaScript (Browser)\\n```javascript\\nconsole.log(\"Hi!\");\\n// or\\nalert(\"Hi!\");\\n```\\n\\n## JavaScript (Node.js)\\n```javascript\\nconsole.log(\"Hi!\");\\n```\\n\\n## Java\\n```java\\npublic class Hello {\\n public static void main(String[] args) {\\n System.out.println(\"Hi!\");\\n }\\n}\\n```\\n\\n## C\\n```c\\n#include \\n\\nint main() {\\n printf(\"Hi!\\\\n\");\\n return 0;\\n}\\n```\\n\\n## C++\\n```cpp\\n#include \\n\\nint main() {\\n std::cout << \"Hi!\" << std::endl;\\n return 0;\\n}\\n```\\n\\n## C#\\n```csharp\\nusing System;\\n\\nclass Program {\\n static void Main() {\\n Console.WriteLine(\"Hi!\");\\n }\\n}\\n```\\n\\n## PHP\\n```php\\n\\n```\\n\\n## Ruby\\n```ruby\\nputs \"Hi!\"\\n```\\n\\n## Go\\n```go\\npackage main\\n\\nimport \"fmt\"\\n\\nfunc main() {\\n fmt.Println(\"Hi!\")\\n}\\n```\\n\\n## Rust\\n```rust\\nfn main() {\\n println!(\"Hi!\");\\n}\\n```\\n\\n## Swift\\n```swift\\nprint(\"Hi!\")\\n```\\n\\n## Kotlin\\n```kotlin\\nfun main() {\\n println(\"Hi!\")\\n}\\n```\\n\\n## HTML (webpage)\\n```html\\n\\n\\n\\n Hi Page\\n\\n\\n

Hi!

\\n\\n\\n```\\n\\nThe Python version is probably the simplest if you\\'re just getting started!', role='assistant', tool_calls=None, function_call=None, provider_specific_fields={'refusal': None}), provider_specific_fields={})], usage=Usage(completion_tokens=347, prompt_tokens=10, total_tokens=357, completion_tokens_details=CompletionTokensDetailsWrapper(accepted_prediction_tokens=None, audio_tokens=None, reasoning_tokens=0, rejected_prediction_tokens=None, text_tokens=None), prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=None, cached_tokens=0, text_tokens=None, image_tokens=None)), service_tier='default')" + ] + }, + "execution_count": 5, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "response = completion(\n", + " model=\"cometapi/deepseek-v3.2-exp\",\n", + " messages=[{\"role\": \"user\", \"content\": \"write code for saying hi\"}]\n", + ")\n", + "response" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Streaming" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Streaming Responses" + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "I'm\n", + " doing\n", + " well\n", + " —\n", + " thanks\n", + " for\n", + " asking\n", + "!\n", + " How\n", + " can\n", + " I\n", + " help\n", + " you\n", + " today\n", + "?\n", + "\n" + ] + } + ], + "source": [ + "messages = [{\"role\": \"user\", \"content\": \"Hey, how's it going?\"}]\n", + "response = completion(model=\"cometapi/gpt-5-mini\", messages=messages, stream=True)\n", + "for part in response:\n", + " print(part.choices[0].delta.content or \"\")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Async Completion" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "ModelResponse(id='chatcmpl-CRAAkmfczlmCEnCM55D9CKexbRenn', created=1760589994, model='gpt-5-mini-2025-08-07', object='chat.completion', system_fingerprint=None, choices=[Choices(finish_reason='stop', index=0, message=Message(content=\"I'm doing well, thanks — how are you? How can I help today?\", role='assistant', tool_calls=None, function_call=None, provider_specific_fields={'refusal': None}, annotations=[]), provider_specific_fields={'content_filter_results': {'hate': {'filtered': False, 'severity': 'safe'}, 'protected_material_code': {'filtered': False, 'detected': False}, 'protected_material_text': {'filtered': False, 'detected': False}, 'self_harm': {'filtered': False, 'severity': 'safe'}, 'sexual': {'filtered': False, 'severity': 'safe'}, 'violence': {'filtered': False, 'severity': 'safe'}}})], usage=Usage(completion_tokens=26, prompt_tokens=12, total_tokens=38, completion_tokens_details=CompletionTokensDetailsWrapper(accepted_prediction_tokens=0, audio_tokens=0, reasoning_tokens=0, rejected_prediction_tokens=0, text_tokens=None), prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=0, cached_tokens=0, text_tokens=None, image_tokens=None)), prompt_filter_results=[{'prompt_index': 0, 'content_filter_results': {'hate': {'filtered': False, 'severity': 'safe'}, 'jailbreak': {'filtered': False, 'detected': False}, 'self_harm': {'filtered': False, 'severity': 'safe'}, 'sexual': {'filtered': False, 'severity': 'safe'}, 'violence': {'filtered': False, 'severity': 'safe'}}}])\n" + ] + } + ], + "source": [ + "from litellm import acompletion\n", + "import asyncio\n", + "\n", + "async def test_get_response():\n", + " user_message = \"Hello, how are you?\"\n", + " messages = [{\"content\": user_message, \"role\": \"user\"}]\n", + " response = await acompletion(model=\"cometapi/gpt-5-mini\", messages=messages)\n", + " return response\n", + "\n", + "response = await test_get_response()\n", + "print(response)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Async Streaming" + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "test acompletion + streaming\n", + "response: \n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content='Hi', role='assistant', function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' —', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' I', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content='’m', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' doing', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' well', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=',', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' thanks', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content='!', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' How', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' are', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' you', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content='?', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' What', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' can', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' I', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' help', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' you', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' with', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content=' today', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(provider_specific_fields=None, content='?', role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None, citations=None)\n", + "ModelResponseStream(id='chatcmpl-CRAAl9VMDBB5skZt638Qx86K9h1Hb', created=1760589996, model='gpt-5-mini', object='chat.completion.chunk', system_fingerprint=None, choices=[StreamingChoices(finish_reason='stop', index=0, delta=Delta(provider_specific_fields=None, content=None, role=None, function_call=None, tool_calls=None, audio=None), logprobs=None)], provider_specific_fields=None)\n" + ] + } + ], + "source": [ + "from litellm import acompletion\n", + "import asyncio, os, traceback\n", + "\n", + "async def completion_call():\n", + " try:\n", + " print(\"test acompletion + streaming\")\n", + " response = await acompletion(\n", + " model=\"cometapi/gpt-5-mini\", \n", + " messages=[{\"content\": \"Hello, how are you?\", \"role\": \"user\"}], \n", + " stream=True\n", + " )\n", + " print(f\"response: {response}\")\n", + " async for chunk in response:\n", + " print(chunk)\n", + " except:\n", + " print(f\"error occurred: {traceback.format_exc()}\")\n", + " pass\n", + "\n", + "await completion_call()" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Embedding" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "EmbeddingResponse(model='text-embedding-3-small', data=[{'object': 'embedding', 'index': 0, 'embedding': [-0.018048199, 0.0047550877, -0.013976435, -0.021936804, -0.038773336, -0.03708264, 0.03854791, -0.0007172257, 0.026473511, -0.0027438616, -0.019823432, -0.011947598, -0.013426959, -0.0059914105, 0.020485623, 0.04269012, -0.028276922, -0.015216281, -0.03325039, 0.045057096, 0.0037477135, 0.015793936, -0.005188329, -0.0071713766, 0.008446445, 0.0070938864, 0.0027632343, 0.025656339, 0.022091785, -0.026797561, -0.029840818, -0.0542714, 0.017907308, -0.03454659, -0.014582269, 0.0429719, 0.03575826, 0.007939235, -0.010123054, -0.029587213, 0.018921727, 0.022556728, 0.019598005, 0.008819807, 0.01655475, 0.043310042, -0.034321167, 0.004441604, 0.032686826, 0.047226828, -0.0043253684, 0.006681779, 0.008995921, 0.06593721, 0.0066078105, -0.023190739, -0.00777721, 0.049875587, -0.028023317, -0.019386668, -0.0013252605, 0.009214303, -0.0055828253, -0.007432026, -0.01199691, 0.0054384116, -0.024247425, 0.047255006, -0.013743965, 0.012799991, 0.009883538, -0.005579303, -0.012616833, -0.014709071, -0.004473305, -0.022204498, 0.010454148, -0.02202134, 0.034631126, 0.0097567355, 0.041478455, -0.010750021, -0.0005340668, -0.041450277, -0.04516981, 0.014709071, -0.06142869, 0.021640932, -0.0008752884, -0.0012671428, -0.045254346, -0.004216178, -0.028220564, -0.011074071, 0.010693664, -0.0029481545, -0.039590508, -0.030178957, -0.019062618, -0.03539194, 0.019978413, 0.019161243, -0.025924034, 0.0058329077, -0.029305428, 0.028840488, 0.02731886, -0.0008048426, -0.056582022, 0.003043256, -0.08459125, -0.012560476, 0.00081100664, 0.0015771041, 0.015188103, -0.0049206354, 0.046888687, -0.027107522, 0.019245777, -0.011412211, 0.039336905, -0.018611765, 0.0429719, -0.04302826, -0.018132735, -0.0074883825, -0.02035882, -0.073038146, -0.042380158, 0.04485985, 0.03361671, -0.033842135, -0.0017259207, -0.0017091898, -0.049199305, -0.024148801, -0.044803493, 0.040943068, -0.03093977, 0.045057096, -0.0186963, -0.014053926, -0.009848315, -0.0070480965, -0.0060054995, -0.02813603, -0.022105874, 0.010235767, -0.004476827, 0.027220234, -0.026290352, 0.0216832, -0.05559578, 0.042464696, 0.0253182, -0.0031594916, 0.02266944, -0.030798879, -0.02896729, -0.0005891025, 0.004061197, -0.042295624, -0.008383043, 0.017146494, -0.025797231, -0.046522368, 0.044127215, 0.021105545, -0.04302826, 0.024430584, -0.014934498, -0.01556851, -0.075743265, -0.022415835, -0.028586883, 0.032743182, 0.03539194, -0.0034483192, -0.04669144, 0.051932603, 0.021626843, 0.03127791, 0.015610777, -0.016470214, -0.0056744046, -0.012743635, 0.060132485, 0.017428277, -0.0039942735, -0.017118316, 0.025106862, 0.008974788, 0.018513141, 0.0016035212, -0.0049593803, 0.0017514573, 0.031221554, -0.034856554, -0.011461522, 0.04621241, 0.044239927, 0.034715664, -0.0121941585, -0.012053267, -0.08526753, -0.011813751, -0.025769053, 0.0125182085, -0.0046670306, 0.038266126, 0.1187997, 0.005787118, -0.030038064, -0.054553185, -0.041506633, -0.024078354, 0.0071854657, -0.013391736, 0.03192601, -0.059625275, 0.0023458432, 0.027924692, 0.09163582, 0.030967949, 0.017639615, 0.01489223, 0.029559033, 0.042943723, 0.0003306547, 0.0047198646, 0.029897174, 0.00012603184, 0.0046811197, -0.0427183, 0.01789322, 0.018175002, -0.0081857955, 0.02581132, 0.009890582, 0.03840702, -0.094453655, -0.01076411, 0.06858598, 0.041647524, 0.033532172, 0.007467249, -0.008235107, -0.030967949, 0.0151317455, 0.027361127, 0.011834885, 0.008707094, 0.008178751, -0.022458103, 0.02844599, 0.003605061, -0.02399382, 0.05212985, 0.06041427, -0.023317542, -0.013335379, -0.044099037, -0.040802173, -0.0047656544, -0.023909286, 0.017315563, 0.017428277, 0.00736158, -0.0016070436, -0.055454887, -0.038012523, -0.020626513, 0.018273626, 0.03260229, -0.016991513, 0.038463376, -0.022458103, -0.0109613575, -0.021810003, 0.04846667, -0.042521052, 0.008601425, -0.019259866, -0.0040048403, -0.03308132, -0.02499415, 0.026783472, -0.032884073, 0.021824092, 0.013145176, -0.009186125, -0.01769597, 0.03240504, -0.015343083, -0.012539342, 0.03578644, -0.012299826, 0.011898286, 0.035730083, 0.058441788, 0.032010544, 0.048720278, 0.012926794, -0.0015207474, -0.03313768, 0.014540002, 0.020189751, 0.00029058868, 0.011531969, -0.022514459, 0.019752987, -0.037956167, 0.005272864, -0.042295624, -0.08521117, -0.03494109, 0.053313337, 0.029981708, -0.008150573, -0.053200625, 0.059681635, -0.035476476, 0.034828376, 0.00087881065, 0.025712697, 0.018668123, 0.03212326, 0.008474623, 0.017836861, 0.004910068, 0.016174342, -0.059681635, 0.04004136, 0.00753065, 0.008150573, -0.038012523, 0.0051178834, 0.012525253, 0.022119964, 0.030235313, 0.008242152, -0.01835816, -0.003150686, 0.010714797, -0.0033162334, -0.028882755, -0.06836055, 0.056159347, 0.013624207, 0.0077349427, 0.0066183778, 0.018583586, -0.008883208, -0.046550546, 0.046945043, -0.07393985, -0.017343743, 0.029530855, -0.010982491, 0.008129438, 0.009700378, -0.024613742, -0.0030097943, 0.0078053884, -0.006438741, 0.04770586, -0.008221018, 0.01654066, 0.02498006, -0.015793936, -0.010827511, 0.02399382, 0.03192601, 0.022923045, -0.029192716, 0.006724046, -0.04601516, 0.038519733, 0.031531516, 0.019443026, 0.000109466084, 0.03525105, -0.027248414, -0.038125236, 0.011771483, -0.007467249, 0.010285079, 0.01670973, -0.007861745, 0.026868006, 0.052327096, 0.026374886, -0.03905512, 0.031193376, -0.053566944, 0.04257741, -0.004670553, 0.0168788, 0.035166513, -0.057878222, 0.07095295, -0.009749691, -0.0137510095, -0.02151413, -0.02431787, 0.010073741, -0.05176353, -0.02083785, 0.003959051, -0.02682574, 0.062104966, -0.011461522, 0.04170388, 0.0076363184, 0.026177637, 0.0144413775, 0.014821785, -0.00046890447, 0.0050544823, 0.00032228927, -0.038970586, -0.011355854, -0.056300238, -0.04302826, -0.003545182, 0.04021043, 0.0051108385, -0.048438493, 0.00252548, -0.07692675, -0.0012433673, 0.0054278444, 0.029305428, -0.016188432, -0.003263399, -0.046156053, 6.031917e-05, 0.060977835, -0.016611107, -0.010637308, -0.012602744, 0.016442036, -0.051509928, -0.016991513, 0.0019407802, 0.019161243, 0.045282524, 0.031869654, -0.036941748, -0.035814617, -0.017850952, -0.027192056, -0.049734693, -0.020964652, 0.0228526, -0.025050506, 0.023472521, 0.025740875, -0.017738238, -0.009813092, -0.030883415, -0.012405495, -0.03277136, -0.029502677, 0.016780175, -0.04421175, -0.0020816717, 0.010341435, 0.059230782, -0.041901127, -0.04119667, 0.025924034, 0.02334572, -0.0008435878, 0.020654691, -0.022753974, 0.010700708, -0.013856677, -0.0121941585, -0.011391076, 0.006590199, 0.0050227814, -0.007960369, 0.0008418266, -0.0198657, 0.10781016, -0.0384352, -0.019147152, 0.0057237167, -0.0038181592, -0.047424074, -0.009341106, 0.018499052, -0.016906979, 0.005642704, -0.01837225, -0.038125236, -0.024895526, -0.010285079, -0.055708494, 0.014173684, -0.019724809, 0.00024215724, 0.04500074, 0.048804812, 0.009777869, -0.006572588, 0.008277375, 0.012328005, -0.012609788, 0.026079014, -0.012990195, 0.017963665, -0.007312268, -0.0015682983, 0.05446865, -0.01258161, 0.00035376972, -0.011299497, -0.036321826, -0.0071854657, 0.012969062, 0.026558045, -0.051819887, -0.0029146927, -0.044606246, -0.010383703, -0.03919601, 0.013624207, 0.0030978515, 0.0121941585, -0.0022225631, 0.0512845, -0.0029780937, -0.025191398, -0.015751667, -0.021006921, 0.0039520063, 0.04418357, 0.020570157, 0.00083390146, 0.020541979, -0.004807922, -0.0114263, -0.036152754, 0.018428607, -0.032658648, -0.002035882, 0.013828499, 0.03144698, -0.0003275727, -0.029756282, 0.008488712, 0.0041879993, 0.027826069, 0.0007273523, -0.018949905, -0.0029023646, 0.007861745, 0.011398122, 0.0125322975, 0.014976765, 0.006318983, -0.0066536004, -0.042915545, 0.025867676, -0.015272637, 0.034602948, 0.050241902, 0.014582269, 0.005987888, 0.015244459, -0.050213724, -0.003212326, 0.01315222, -0.022866689, -0.004772699, 0.035673723, -0.024796901, -0.00699174, -0.002072866, 0.022077696, 0.021147812, 0.005093227, -0.039618686, -0.0049241576, 0.012264604, -0.062104966, -0.0022613083, -0.004339458, 0.065486364, -0.0033320836, 0.02944632, 0.017498722, 0.0033039053, -0.020260196, -0.0154980635, -0.05460954, -0.03626547, 0.0072629564, 0.0028900367, 1.2954037e-05, 0.01769597, -0.0045930627, 0.022260854, 0.0027192058, -0.0010566862, -0.0005212985, 0.012158935, -0.0017312041, -0.035110157, -0.0036032998, -0.02317665, -0.01639977, 0.010327346, 0.018259536, -0.011095204, -0.00061904197, -0.023134382, -0.011989865, 0.0025924034, 0.0056708823, 0.03110884, -0.013462181, -0.021105545, 0.010376658, -0.010017385, -0.025106862, 0.026093103, 0.018456785, -0.02134506, 0.0066993902, 0.011891241, -0.010017385, -0.012687278, -0.017132405, 0.04717047, 0.012475941, -0.018752657, -0.008657781, 0.005276386, -0.02582541, 0.02913636, -0.0193444, -0.01101067, 0.029305428, 0.011736261, 0.043140974, 0.02135915, 0.00089422066, 0.009827181, 0.013638296, 0.013884856, -0.014004614, 0.010285079, 0.008108305, -0.04035132, -0.02978446, 0.008481667, -0.022289034, 0.01621661, -0.0057941624, -0.019090796, -0.01852723, -0.022923045, 0.0077208537, -0.039985005, -0.017428277, -0.009460864, 0.018301804, 0.0014397348, 0.04815671, -0.012187114, 0.018879458, 0.021739556, 0.018414518, -0.013462181, -0.06368295, 0.0057096276, 0.013088819, 0.0061640027, 0.031193376, -0.008728228, -0.019245777, 0.010735931, 0.012454808, 0.0397314, -0.017597347, 0.012278693, 0.0130465515, -0.025473181, -0.03215144, -0.0053292206, -0.0027068777, 0.014068015, -0.028079674, 0.016498392, 0.015159924, -0.009207259, -0.02334572, -0.0013710503, 0.008488712, 0.0012231142, 0.0020464489, -0.025149131, -0.021063277, -0.014427288, -0.035222873, 0.051030897, 0.016103897, 0.0063401167, -0.03093977, -0.004684642, -0.0070199184, -0.008495756, -0.0038674714, -0.012222337, -0.022556728, 0.0036015387, -0.040943068, -0.011362898, 0.016794264, 0.017766416, -0.014194817, 0.0011755633, -0.039759576, 0.011384032, -0.0006318103, 0.008298509, 0.04449353, 0.0004838742, 0.016935157, -0.011341765, 0.016864711, -0.00027892113, 0.0009140335, -0.031306088, -0.049452912, -0.0068367594, -0.00011216283, -0.005079138, -0.014420244, 0.01803411, 0.03984411, -0.026276262, -0.0011077593, -0.00063313113, -0.006301372, -0.019992502, 0.0064316965, -0.024289692, 0.0120039545, 0.0068649375, -0.017724149, -0.015667133, -0.0036490895, -0.007953324, 0.024627833, 0.024402406, 0.021810003, -0.015977094, 0.010524594, -0.0060054995, 0.0414221, -0.048551206, 0.01472316, 0.015427617, 0.0029217373, 0.012666144, 0.0048995013, 0.007326357, -0.04187295, -0.0064176074, -0.00674518, 0.0047762212, -0.053059734, -0.09541172, 0.022063607, 0.029530855, 0.01556851, 0.011292453, -0.0038709936, -0.0055370354, -0.016005272, -0.0035170037, -0.0572583, 0.038632445, 0.007981502, -0.005434889, -0.023895197, 0.0021380284, -0.0015084195, 0.016117986, 0.005434889, -0.014694982, -0.007104453, 0.011595369, -0.055229463, 0.0036455672, 0.0027104, -0.010052607, -0.023697948, -0.016315235, -0.002757951, 0.039505973, 0.011095204, 0.0002681341, 0.058948997, -0.0074883825, 0.0050122146, 0.040604927, 0.012912705, -0.025078684, 0.040464036, -0.008925476, -0.00876345, -0.040633105, -0.009024099, 0.024796901, 0.03592733, 0.03626547, -0.029474499, -0.00055431994, 0.0010839839, 0.016737908, 0.013286067, -0.005441934, 0.0059420983, -0.0121941585, 0.015089478, -0.010186454, -0.03477202, -0.0076363184, -0.0087141385, 0.0018439173, 0.028065585, -0.022331301, 0.0029516767, -0.045789734, 0.0010672531, 0.018287715, -0.015948916, 0.04849485, 0.0057589393, 0.0066219, 0.002196146, -0.047255006, 0.012116668, 0.02085194, 0.025924034, -0.0036737456, -0.02877004, 0.016906979, -0.037336245, -0.016258877, 0.010883868, -0.003765325, -0.0049523357, -0.002613537, -0.03263047, 0.023204828, 0.0049946033, -0.007692675, -0.034236632, 0.034095738, 0.020133393, 0.019259866, -0.014103238, 0.024599653, 0.005889264, 0.02430378, 0.0111233825, -0.018780835, -0.00040550332, 0.020232018, 0.03806888, 0.009890582, 0.032376863, 0.031052483, 0.01871039, 0.03891423, -0.0009739124, 0.002759712, 0.017498722, -0.01158128, -0.0045578396, 0.02744566, 0.06497915, 0.024853257, 0.004709298, 0.016667463, -0.00066263025, -0.018132735, -0.013138131, -0.01124314, -0.0125182085, -0.0038111147, 0.03361671, -0.007270001, 0.0012011, -0.01771006, -0.00039999973, 0.024021998, 0.0027896515, 0.0024744067, 0.0013965869, -0.05939985, 0.0014150789, -0.0052517303, 0.052524347, 0.015779847, -0.03327857, 0.042633764, 0.0059420983, -0.023387987, 0.0039097387, -0.028023317, -0.011863063, 0.004378203, 0.02052789, -0.063626595, -0.014864052, 0.014293442, -0.00015938349, -0.007932191, -0.0010954313, 0.023528878, -0.007467249, 0.0059667546, 0.017132405, 0.005730761, -0.00020495309, -0.032038722, 0.0036631785, 0.042915545, -0.029925352, 0.015667133, 0.018935816, -0.0072065997, 0.01556851, -0.025473181, 0.017625526, -0.0026698937, -0.007446115, -0.008622559, -0.043422755, -0.020133393, -0.0039801844, 0.01489223, -0.021655021, 0.015357172, -0.03640636, -0.005663838, -0.028530527, 0.0022648307, -0.00043015933, 0.043591827, -0.015526242, 0.011870108, -0.02530411, -0.016315235, -0.00032316984, -0.030150779, -0.0052552526, 0.020372909, 0.0075024716, 0.0104330145, -0.00055608107, -0.026248084, -0.015202192, -0.03341946, 0.031559695, -0.0012046222, 0.07185466, -0.039590508, 0.022979401, 0.05810365, 0.014025748, -0.029756282, -0.022866689, 0.0073897582, 0.037618026, -0.004180955, -0.0051566283, 0.009728557, -0.03604004, 0.040633105, 0.0026963109, -0.0054172776, 0.034095738, -0.00595971, 0.040943068, -0.031390622, 0.055962097, 0.02117599, -0.012912705, -0.019626183, 0.055877563, 0.017343743, -0.0035416598, 0.013257889, -0.0186963, 0.01656884, -0.06396473, -0.0055405577, 0.020767406, -0.0046564634, 0.045085277, -0.009221348, 0.013645341, 0.008777539, 0.004730432, -0.018625854, -0.011067026, 0.021500042, -0.015047211, 0.004600107, -0.0014344514, -0.0023740216, -0.016188432, 0.006209792, 0.0011993388, 0.004180955, -0.017160583, 0.014497734, 0.015371261, 0.018259536, -0.028333278, -0.008390088, 0.041929305, 0.003923828, 0.02550136, -0.003300383, -0.008058993, -0.010418925, 0.058216363, 0.01885128, -0.02020384, 0.002858336, -0.009806047, -0.022274945, 0.0070445742, 0.026670758, 0.008213974, -0.035307407, -0.027713355, 0.042915545, -0.039675042, -0.0029217373, 0.012053267, -0.003853382, 0.01133472, -0.010073741, 0.005878697, 0.0070938864, -0.035673723, 0.024205158, 0.005896309, 0.030573452, 0.02416289, -0.0072911344, 0.01738601, 0.017005602, -0.02846008, 0.0030344503, 0.018794924, -0.0148076955, -0.0344057, 0.025430914, 0.033503994, -0.0050580045, 0.0077138087, 0.03243322, 0.01372283, -0.005441934, 0.0073404466, -0.0007832686, -0.04767768, 0.0070480965, 0.015145835, 0.026233995, -0.01670973, -0.019513471, -0.014849963, 0.007953324, -0.0032176094, 0.006572588, -0.0012477703, 0.004230267, 0.004476827, -0.021810003, -0.030009886, -0.019273955, -0.0030414949, -0.002918215, 0.060639694, 0.024641922, 0.010327346, 0.026558045, 0.018921727, -0.025867676, -0.016117986, 0.023881108, 0.025360467, 0.009770825, 0.03792799, -0.022429924, 0.033363104, -0.0018914682, 0.04040768, 0.018484963, 0.0070199184, -0.017583257, 0.016258877, 0.010954313, -0.008939565, -0.024148801, -0.02498006, -0.007889924, 0.02748793, 0.0307707, 0.029756282, 0.0051425393, 0.0045719286, -0.03046074, 0.013596028, 0.025684519, -0.0033197557, 0.006967084, 0.03677268, 0.0120039545, -0.0032792494, -0.0032211316, -0.02399382, -0.026924362, -0.013920079, -0.0042197, 0.025346378, -0.0027015943, -0.016991513, 0.0031594916, -0.007579962, 0.018978084, 0.017681882, 0.0126591, 0.028939111, 0.008833896, 0.10183637, 0.0059632324, -0.05196078, -0.023697948, 0.011045893, -0.008777539, -0.013807366, 0.019273955, -0.025346378, 0.0074742935, 0.009961028, -0.010813422, 0.018597675, 0.009636978, 0.014948587, 0.024064265, -0.008693005, -0.020570157, 0.014194817, -0.026219906, -0.02299349, 0.011067026, 0.032066904, 0.013391736, -0.05148175, -0.009489042, -0.03062981, -0.0012847543, 0.07286908, 0.026529867, -0.00025008238, 0.013638296, 0.016089808, 0.018654034, -0.0020394044, -0.024543297, 0.0147795165, -0.009601755, 0.0018791402, -0.040520392, -0.003360262, 0.02216223, 0.0137650985, 0.0059914105, 0.0048361, -0.0009844792, 0.016977424, -0.00934815, 0.024233336, -0.013088819, -0.017555078, -0.0050263037, 0.010595039, -0.027516108, 0.0071537653, -0.023247095, -0.0017655464, -0.015948916, 0.058160007, -0.025966302, 0.0121941585, -0.012384362, -0.0015612538, 0.009946939, 0.00628376, 0.011327676, 0.0109613575, 0.008601425, -0.018329982, 0.055680316, -0.012778858, -0.0100807855, -0.011067026, -0.0036490895, -0.01356785, 0.0073193125, -0.014272308, -0.027403394, -0.030742522, 0.02862915, 0.03062981, -0.014596358, 0.021697288, 0.0042408337, 0.027572464, 0.0019601528, -0.037138995, -0.031306088, 0.041929305, 0.017738238, 0.004857234, 0.008256241, 0.0118278405, 0.021753646, -0.00160176, -0.0018333505, -0.0047374764, 0.042239267, 0.0058329077, -0.026459422, 0.015075389, 0.021147812, -0.005212985, 0.01281408, 0.017738238, 0.008242152, -0.020372909, -0.011081115, -0.011017715, 0.007706764, 0.01834407, 0.01954165, 0.037477136, -0.010278034, 0.015808025, 0.00031590514, -0.017681882, -0.008967743, -0.020612424, -0.025416825, -0.0037970257, -0.029868996, 0.01720285, -0.0144554665, 0.026727116, 0.00414221, 0.0040154075, 0.05838543, 0.0005622451, -0.025219576, 0.004180955, -0.002932304, -0.0090663675, 0.011574236, 0.02450103, -0.012553431, -0.020612424, -0.032095082, 0.015526242, 0.008974788, 0.0053151315, -0.0003112821, -0.017935487, -0.0076222294, 0.03358853, 0.029474499, -0.011496745, -0.012835215, -0.020739228, -0.012482986, -0.037871633, 0.0052517303, -0.012926794, -0.0025237189, 0.0020323596, 0.045113456, -0.04835396, -0.027755624, -0.0079955915, 0.007896968, 0.0072559114, 0.015047211, -0.0014573464, -0.014032792, 0.021091456, -0.0046071517, -0.0065232757, -0.02582541, -0.035870973, -0.015343083, 0.03254593, -0.028431902, -0.003286294, 0.014328664, 0.008840941, 0.015948916, 0.012835215, 0.019400757, -0.012342094, -0.010693664, 0.004772699, -0.03254593, 0.010707753, -0.016822444, -0.0032827717, 0.021246437, -0.04485985, -0.04384543, -0.015906649, -0.009707424, 0.02299349, 0.019513471, -0.010151232, 0.018963994, -0.0057976847, 0.05739919, -0.019922055, -0.029108182, -0.0106232185, 0.021077367, 0.0036455672, -0.026614401, 0.04497256, -0.04446535, -0.0004556959, -0.004578973, 0.003962573, -0.004910068, 0.015089478, -0.0301226, 0.007664497, 0.008375999, 0.031982366, 0.006135824, 0.02152822, -0.015469885, -0.007210122, 0.034715664, -0.01233505, 0.0004490916, -0.0144413775, -0.003150686, -0.02003477, -0.027924692, -0.0015850292, -0.009376328, -0.0035997776, -0.03240504, -0.010912046, 0.0031999978, 0.022303123, -0.008988877, 0.00024633997, -0.0035698381, 0.0070974086, -0.002599448, -0.042267445, -0.016935157, -0.0002481011, -0.041393917, 0.014483645, 0.019006262, -0.02813603, 0.0072030774, -7.3032425e-05, 0.01802002, -0.017188761, 0.015991183, 0.020401087, 0.03542012, 0.04469078, 0.04071764, 0.011095204, -0.031390622, -0.03254593, 0.014187773, 0.016272966, -0.009721513, -0.026388975, -0.014849963, -0.005642704, -0.022556728, 0.0064457855, -0.043450933, 0.010834555, -0.015977094, 0.020880118, -0.02385293, -0.054806788, 0.03789981, 0.0013516777, -0.026431242, -0.015540331, 0.016695641, -0.037167173, -0.021190079, 0.023881108, -0.0045860177, 0.0064105624, -0.007763121, -0.013053596, 0.024472851, -0.0004962022, -0.00976378, 0.060019772, -0.0057624616, -0.04384543, 0.010313257, 0.0076715415, 0.0025888812, -0.03589915, 0.008791628, -0.012785902, 0.01042597, 0.015653044, 0.04767768, -0.009869449, 0.0064457855, -0.010947268, -0.0077349427, -0.032715004, -0.023867019, -0.011327676, -0.00046274049, -0.036998104, 0.013913034, 0.012250515, -0.009996251, 0.021204168, 0.020091126, -0.003740669, -0.0049769916, -0.0140891485, 0.024064265, 0.0038815604, 0.025684519, 0.041788414, -0.013553761, 0.006681779, -0.0050826604, -0.018175002, 0.008228063, -0.006230926, -0.018907638, 0.0154839745, -0.028713685, -0.015047211, -0.019682541, 0.02516322, 0.040802173, 0.007213644, 0.011743305, -0.015963005, -0.03818159, 0.01191942, -0.031728763, -0.011863063, 0.023881108, 0.0053116092, -0.020992832, -0.017991843, -0.00405063, -0.017780505, -0.0057659843, 0.02978446, 0.031165197, 0.0014221234, 0.021316882, 0.026008569, -0.0018544842, -0.032658648, 0.028474169, 0.013109953, 0.018076377, 0.0007991189, -0.0042373114, 0.028910933, -0.0029358263, 0.021866359, 0.024472851, -0.002576553, -0.033532172, 0.01920351, -0.0095665315, -0.03093977, 0.0034817809, 0.018654034, -0.0074038478, 0.021443684, 0.0038604268, -0.02745975, 0.031587873, 0.0061146906, 0.022711707, -0.019795254, -0.016991513, -0.04471896, -0.007875834, -0.0034941088, -0.043789074, 0.021091456, 0.024909616, -0.013194487, -0.0042690123, 0.027896514, -0.018414518, -0.023303451, -0.025797231, -0.009524264]}], object='list', usage=Usage(completion_tokens=0, prompt_tokens=3, total_tokens=3, completion_tokens_details=None, prompt_tokens_details=None))\n" + ] + } + ], + "source": [ + "import litellm\n", + "\n", + "\n", + "async def main():\n", + " response = await litellm.aembedding(\n", + " model=\"cometapi/text-embedding-3-small\", # The model name must include prefix \"openai\" + the model name from ai/ml api\n", + " api_key=api_key, # your aiml api-key\n", + " api_base=\"https://api.cometapi.com/v1\", # 👈 the URL has changed from v2 to v1\n", + " input=\"Your text string\",\n", + " )\n", + " print(response)\n", + "\n", + "await main()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "EmbeddingResponse(model='text-embedding-3-small', data=[{'object': 'embedding', 'index': 0, 'embedding': [-0.018048199, 0.0047550877, -0.013976435, -0.021936804, -0.038773336, -0.03708264, 0.03854791, -0.0007172257, 0.026473511, -0.0027438616, -0.019823432, -0.011947598, -0.013426959, -0.0059914105, 0.020485623, 0.04269012, -0.028276922, -0.015216281, -0.03325039, 0.045057096, 0.0037477135, 0.015793936, -0.005188329, -0.0071713766, 0.008446445, 0.0070938864, 0.0027632343, 0.025656339, 0.022091785, -0.026797561, -0.029840818, -0.0542714, 0.017907308, -0.03454659, -0.014582269, 0.0429719, 0.03575826, 0.007939235, -0.010123054, -0.029587213, 0.018921727, 0.022556728, 0.019598005, 0.008819807, 0.01655475, 0.043310042, -0.034321167, 0.004441604, 0.032686826, 0.047226828, -0.0043253684, 0.006681779, 0.008995921, 0.06593721, 0.0066078105, -0.023190739, -0.00777721, 0.049875587, -0.028023317, -0.019386668, -0.0013252605, 0.009214303, -0.0055828253, -0.007432026, -0.01199691, 0.0054384116, -0.024247425, 0.047255006, -0.013743965, 0.012799991, 0.009883538, -0.005579303, -0.012616833, -0.014709071, -0.004473305, -0.022204498, 0.010454148, -0.02202134, 0.034631126, 0.0097567355, 0.041478455, -0.010750021, -0.0005340668, -0.041450277, -0.04516981, 0.014709071, -0.06142869, 0.021640932, -0.0008752884, -0.0012671428, -0.045254346, -0.004216178, -0.028220564, -0.011074071, 0.010693664, -0.0029481545, -0.039590508, -0.030178957, -0.019062618, -0.03539194, 0.019978413, 0.019161243, -0.025924034, 0.0058329077, -0.029305428, 0.028840488, 0.02731886, -0.0008048426, -0.056582022, 0.003043256, -0.08459125, -0.012560476, 0.00081100664, 0.0015771041, 0.015188103, -0.0049206354, 0.046888687, -0.027107522, 0.019245777, -0.011412211, 0.039336905, -0.018611765, 0.0429719, -0.04302826, -0.018132735, -0.0074883825, -0.02035882, -0.073038146, -0.042380158, 0.04485985, 0.03361671, -0.033842135, -0.0017259207, -0.0017091898, -0.049199305, -0.024148801, -0.044803493, 0.040943068, -0.03093977, 0.045057096, -0.0186963, -0.014053926, -0.009848315, -0.0070480965, -0.0060054995, -0.02813603, -0.022105874, 0.010235767, -0.004476827, 0.027220234, -0.026290352, 0.0216832, -0.05559578, 0.042464696, 0.0253182, -0.0031594916, 0.02266944, -0.030798879, -0.02896729, -0.0005891025, 0.004061197, -0.042295624, -0.008383043, 0.017146494, -0.025797231, -0.046522368, 0.044127215, 0.021105545, -0.04302826, 0.024430584, -0.014934498, -0.01556851, -0.075743265, -0.022415835, -0.028586883, 0.032743182, 0.03539194, -0.0034483192, -0.04669144, 0.051932603, 0.021626843, 0.03127791, 0.015610777, -0.016470214, -0.0056744046, -0.012743635, 0.060132485, 0.017428277, -0.0039942735, -0.017118316, 0.025106862, 0.008974788, 0.018513141, 0.0016035212, -0.0049593803, 0.0017514573, 0.031221554, -0.034856554, -0.011461522, 0.04621241, 0.044239927, 0.034715664, -0.0121941585, -0.012053267, -0.08526753, -0.011813751, -0.025769053, 0.0125182085, -0.0046670306, 0.038266126, 0.1187997, 0.005787118, -0.030038064, -0.054553185, -0.041506633, -0.024078354, 0.0071854657, -0.013391736, 0.03192601, -0.059625275, 0.0023458432, 0.027924692, 0.09163582, 0.030967949, 0.017639615, 0.01489223, 0.029559033, 0.042943723, 0.0003306547, 0.0047198646, 0.029897174, 0.00012603184, 0.0046811197, -0.0427183, 0.01789322, 0.018175002, -0.0081857955, 0.02581132, 0.009890582, 0.03840702, -0.094453655, -0.01076411, 0.06858598, 0.041647524, 0.033532172, 0.007467249, -0.008235107, -0.030967949, 0.0151317455, 0.027361127, 0.011834885, 0.008707094, 0.008178751, -0.022458103, 0.02844599, 0.003605061, -0.02399382, 0.05212985, 0.06041427, -0.023317542, -0.013335379, -0.044099037, -0.040802173, -0.0047656544, -0.023909286, 0.017315563, 0.017428277, 0.00736158, -0.0016070436, -0.055454887, -0.038012523, -0.020626513, 0.018273626, 0.03260229, -0.016991513, 0.038463376, -0.022458103, -0.0109613575, -0.021810003, 0.04846667, -0.042521052, 0.008601425, -0.019259866, -0.0040048403, -0.03308132, -0.02499415, 0.026783472, -0.032884073, 0.021824092, 0.013145176, -0.009186125, -0.01769597, 0.03240504, -0.015343083, -0.012539342, 0.03578644, -0.012299826, 0.011898286, 0.035730083, 0.058441788, 0.032010544, 0.048720278, 0.012926794, -0.0015207474, -0.03313768, 0.014540002, 0.020189751, 0.00029058868, 0.011531969, -0.022514459, 0.019752987, -0.037956167, 0.005272864, -0.042295624, -0.08521117, -0.03494109, 0.053313337, 0.029981708, -0.008150573, -0.053200625, 0.059681635, -0.035476476, 0.034828376, 0.00087881065, 0.025712697, 0.018668123, 0.03212326, 0.008474623, 0.017836861, 0.004910068, 0.016174342, -0.059681635, 0.04004136, 0.00753065, 0.008150573, -0.038012523, 0.0051178834, 0.012525253, 0.022119964, 0.030235313, 0.008242152, -0.01835816, -0.003150686, 0.010714797, -0.0033162334, -0.028882755, -0.06836055, 0.056159347, 0.013624207, 0.0077349427, 0.0066183778, 0.018583586, -0.008883208, -0.046550546, 0.046945043, -0.07393985, -0.017343743, 0.029530855, -0.010982491, 0.008129438, 0.009700378, -0.024613742, -0.0030097943, 0.0078053884, -0.006438741, 0.04770586, -0.008221018, 0.01654066, 0.02498006, -0.015793936, -0.010827511, 0.02399382, 0.03192601, 0.022923045, -0.029192716, 0.006724046, -0.04601516, 0.038519733, 0.031531516, 0.019443026, 0.000109466084, 0.03525105, -0.027248414, -0.038125236, 0.011771483, -0.007467249, 0.010285079, 0.01670973, -0.007861745, 0.026868006, 0.052327096, 0.026374886, -0.03905512, 0.031193376, -0.053566944, 0.04257741, -0.004670553, 0.0168788, 0.035166513, -0.057878222, 0.07095295, -0.009749691, -0.0137510095, -0.02151413, -0.02431787, 0.010073741, -0.05176353, -0.02083785, 0.003959051, -0.02682574, 0.062104966, -0.011461522, 0.04170388, 0.0076363184, 0.026177637, 0.0144413775, 0.014821785, -0.00046890447, 0.0050544823, 0.00032228927, -0.038970586, -0.011355854, -0.056300238, -0.04302826, -0.003545182, 0.04021043, 0.0051108385, -0.048438493, 0.00252548, -0.07692675, -0.0012433673, 0.0054278444, 0.029305428, -0.016188432, -0.003263399, -0.046156053, 6.031917e-05, 0.060977835, -0.016611107, -0.010637308, -0.012602744, 0.016442036, -0.051509928, -0.016991513, 0.0019407802, 0.019161243, 0.045282524, 0.031869654, -0.036941748, -0.035814617, -0.017850952, -0.027192056, -0.049734693, -0.020964652, 0.0228526, -0.025050506, 0.023472521, 0.025740875, -0.017738238, -0.009813092, -0.030883415, -0.012405495, -0.03277136, -0.029502677, 0.016780175, -0.04421175, -0.0020816717, 0.010341435, 0.059230782, -0.041901127, -0.04119667, 0.025924034, 0.02334572, -0.0008435878, 0.020654691, -0.022753974, 0.010700708, -0.013856677, -0.0121941585, -0.011391076, 0.006590199, 0.0050227814, -0.007960369, 0.0008418266, -0.0198657, 0.10781016, -0.0384352, -0.019147152, 0.0057237167, -0.0038181592, -0.047424074, -0.009341106, 0.018499052, -0.016906979, 0.005642704, -0.01837225, -0.038125236, -0.024895526, -0.010285079, -0.055708494, 0.014173684, -0.019724809, 0.00024215724, 0.04500074, 0.048804812, 0.009777869, -0.006572588, 0.008277375, 0.012328005, -0.012609788, 0.026079014, -0.012990195, 0.017963665, -0.007312268, -0.0015682983, 0.05446865, -0.01258161, 0.00035376972, -0.011299497, -0.036321826, -0.0071854657, 0.012969062, 0.026558045, -0.051819887, -0.0029146927, -0.044606246, -0.010383703, -0.03919601, 0.013624207, 0.0030978515, 0.0121941585, -0.0022225631, 0.0512845, -0.0029780937, -0.025191398, -0.015751667, -0.021006921, 0.0039520063, 0.04418357, 0.020570157, 0.00083390146, 0.020541979, -0.004807922, -0.0114263, -0.036152754, 0.018428607, -0.032658648, -0.002035882, 0.013828499, 0.03144698, -0.0003275727, -0.029756282, 0.008488712, 0.0041879993, 0.027826069, 0.0007273523, -0.018949905, -0.0029023646, 0.007861745, 0.011398122, 0.0125322975, 0.014976765, 0.006318983, -0.0066536004, -0.042915545, 0.025867676, -0.015272637, 0.034602948, 0.050241902, 0.014582269, 0.005987888, 0.015244459, -0.050213724, -0.003212326, 0.01315222, -0.022866689, -0.004772699, 0.035673723, -0.024796901, -0.00699174, -0.002072866, 0.022077696, 0.021147812, 0.005093227, -0.039618686, -0.0049241576, 0.012264604, -0.062104966, -0.0022613083, -0.004339458, 0.065486364, -0.0033320836, 0.02944632, 0.017498722, 0.0033039053, -0.020260196, -0.0154980635, -0.05460954, -0.03626547, 0.0072629564, 0.0028900367, 1.2954037e-05, 0.01769597, -0.0045930627, 0.022260854, 0.0027192058, -0.0010566862, -0.0005212985, 0.012158935, -0.0017312041, -0.035110157, -0.0036032998, -0.02317665, -0.01639977, 0.010327346, 0.018259536, -0.011095204, -0.00061904197, -0.023134382, -0.011989865, 0.0025924034, 0.0056708823, 0.03110884, -0.013462181, -0.021105545, 0.010376658, -0.010017385, -0.025106862, 0.026093103, 0.018456785, -0.02134506, 0.0066993902, 0.011891241, -0.010017385, -0.012687278, -0.017132405, 0.04717047, 0.012475941, -0.018752657, -0.008657781, 0.005276386, -0.02582541, 0.02913636, -0.0193444, -0.01101067, 0.029305428, 0.011736261, 0.043140974, 0.02135915, 0.00089422066, 0.009827181, 0.013638296, 0.013884856, -0.014004614, 0.010285079, 0.008108305, -0.04035132, -0.02978446, 0.008481667, -0.022289034, 0.01621661, -0.0057941624, -0.019090796, -0.01852723, -0.022923045, 0.0077208537, -0.039985005, -0.017428277, -0.009460864, 0.018301804, 0.0014397348, 0.04815671, -0.012187114, 0.018879458, 0.021739556, 0.018414518, -0.013462181, -0.06368295, 0.0057096276, 0.013088819, 0.0061640027, 0.031193376, -0.008728228, -0.019245777, 0.010735931, 0.012454808, 0.0397314, -0.017597347, 0.012278693, 0.0130465515, -0.025473181, -0.03215144, -0.0053292206, -0.0027068777, 0.014068015, -0.028079674, 0.016498392, 0.015159924, -0.009207259, -0.02334572, -0.0013710503, 0.008488712, 0.0012231142, 0.0020464489, -0.025149131, -0.021063277, -0.014427288, -0.035222873, 0.051030897, 0.016103897, 0.0063401167, -0.03093977, -0.004684642, -0.0070199184, -0.008495756, -0.0038674714, -0.012222337, -0.022556728, 0.0036015387, -0.040943068, -0.011362898, 0.016794264, 0.017766416, -0.014194817, 0.0011755633, -0.039759576, 0.011384032, -0.0006318103, 0.008298509, 0.04449353, 0.0004838742, 0.016935157, -0.011341765, 0.016864711, -0.00027892113, 0.0009140335, -0.031306088, -0.049452912, -0.0068367594, -0.00011216283, -0.005079138, -0.014420244, 0.01803411, 0.03984411, -0.026276262, -0.0011077593, -0.00063313113, -0.006301372, -0.019992502, 0.0064316965, -0.024289692, 0.0120039545, 0.0068649375, -0.017724149, -0.015667133, -0.0036490895, -0.007953324, 0.024627833, 0.024402406, 0.021810003, -0.015977094, 0.010524594, -0.0060054995, 0.0414221, -0.048551206, 0.01472316, 0.015427617, 0.0029217373, 0.012666144, 0.0048995013, 0.007326357, -0.04187295, -0.0064176074, -0.00674518, 0.0047762212, -0.053059734, -0.09541172, 0.022063607, 0.029530855, 0.01556851, 0.011292453, -0.0038709936, -0.0055370354, -0.016005272, -0.0035170037, -0.0572583, 0.038632445, 0.007981502, -0.005434889, -0.023895197, 0.0021380284, -0.0015084195, 0.016117986, 0.005434889, -0.014694982, -0.007104453, 0.011595369, -0.055229463, 0.0036455672, 0.0027104, -0.010052607, -0.023697948, -0.016315235, -0.002757951, 0.039505973, 0.011095204, 0.0002681341, 0.058948997, -0.0074883825, 0.0050122146, 0.040604927, 0.012912705, -0.025078684, 0.040464036, -0.008925476, -0.00876345, -0.040633105, -0.009024099, 0.024796901, 0.03592733, 0.03626547, -0.029474499, -0.00055431994, 0.0010839839, 0.016737908, 0.013286067, -0.005441934, 0.0059420983, -0.0121941585, 0.015089478, -0.010186454, -0.03477202, -0.0076363184, -0.0087141385, 0.0018439173, 0.028065585, -0.022331301, 0.0029516767, -0.045789734, 0.0010672531, 0.018287715, -0.015948916, 0.04849485, 0.0057589393, 0.0066219, 0.002196146, -0.047255006, 0.012116668, 0.02085194, 0.025924034, -0.0036737456, -0.02877004, 0.016906979, -0.037336245, -0.016258877, 0.010883868, -0.003765325, -0.0049523357, -0.002613537, -0.03263047, 0.023204828, 0.0049946033, -0.007692675, -0.034236632, 0.034095738, 0.020133393, 0.019259866, -0.014103238, 0.024599653, 0.005889264, 0.02430378, 0.0111233825, -0.018780835, -0.00040550332, 0.020232018, 0.03806888, 0.009890582, 0.032376863, 0.031052483, 0.01871039, 0.03891423, -0.0009739124, 0.002759712, 0.017498722, -0.01158128, -0.0045578396, 0.02744566, 0.06497915, 0.024853257, 0.004709298, 0.016667463, -0.00066263025, -0.018132735, -0.013138131, -0.01124314, -0.0125182085, -0.0038111147, 0.03361671, -0.007270001, 0.0012011, -0.01771006, -0.00039999973, 0.024021998, 0.0027896515, 0.0024744067, 0.0013965869, -0.05939985, 0.0014150789, -0.0052517303, 0.052524347, 0.015779847, -0.03327857, 0.042633764, 0.0059420983, -0.023387987, 0.0039097387, -0.028023317, -0.011863063, 0.004378203, 0.02052789, -0.063626595, -0.014864052, 0.014293442, -0.00015938349, -0.007932191, -0.0010954313, 0.023528878, -0.007467249, 0.0059667546, 0.017132405, 0.005730761, -0.00020495309, -0.032038722, 0.0036631785, 0.042915545, -0.029925352, 0.015667133, 0.018935816, -0.0072065997, 0.01556851, -0.025473181, 0.017625526, -0.0026698937, -0.007446115, -0.008622559, -0.043422755, -0.020133393, -0.0039801844, 0.01489223, -0.021655021, 0.015357172, -0.03640636, -0.005663838, -0.028530527, 0.0022648307, -0.00043015933, 0.043591827, -0.015526242, 0.011870108, -0.02530411, -0.016315235, -0.00032316984, -0.030150779, -0.0052552526, 0.020372909, 0.0075024716, 0.0104330145, -0.00055608107, -0.026248084, -0.015202192, -0.03341946, 0.031559695, -0.0012046222, 0.07185466, -0.039590508, 0.022979401, 0.05810365, 0.014025748, -0.029756282, -0.022866689, 0.0073897582, 0.037618026, -0.004180955, -0.0051566283, 0.009728557, -0.03604004, 0.040633105, 0.0026963109, -0.0054172776, 0.034095738, -0.00595971, 0.040943068, -0.031390622, 0.055962097, 0.02117599, -0.012912705, -0.019626183, 0.055877563, 0.017343743, -0.0035416598, 0.013257889, -0.0186963, 0.01656884, -0.06396473, -0.0055405577, 0.020767406, -0.0046564634, 0.045085277, -0.009221348, 0.013645341, 0.008777539, 0.004730432, -0.018625854, -0.011067026, 0.021500042, -0.015047211, 0.004600107, -0.0014344514, -0.0023740216, -0.016188432, 0.006209792, 0.0011993388, 0.004180955, -0.017160583, 0.014497734, 0.015371261, 0.018259536, -0.028333278, -0.008390088, 0.041929305, 0.003923828, 0.02550136, -0.003300383, -0.008058993, -0.010418925, 0.058216363, 0.01885128, -0.02020384, 0.002858336, -0.009806047, -0.022274945, 0.0070445742, 0.026670758, 0.008213974, -0.035307407, -0.027713355, 0.042915545, -0.039675042, -0.0029217373, 0.012053267, -0.003853382, 0.01133472, -0.010073741, 0.005878697, 0.0070938864, -0.035673723, 0.024205158, 0.005896309, 0.030573452, 0.02416289, -0.0072911344, 0.01738601, 0.017005602, -0.02846008, 0.0030344503, 0.018794924, -0.0148076955, -0.0344057, 0.025430914, 0.033503994, -0.0050580045, 0.0077138087, 0.03243322, 0.01372283, -0.005441934, 0.0073404466, -0.0007832686, -0.04767768, 0.0070480965, 0.015145835, 0.026233995, -0.01670973, -0.019513471, -0.014849963, 0.007953324, -0.0032176094, 0.006572588, -0.0012477703, 0.004230267, 0.004476827, -0.021810003, -0.030009886, -0.019273955, -0.0030414949, -0.002918215, 0.060639694, 0.024641922, 0.010327346, 0.026558045, 0.018921727, -0.025867676, -0.016117986, 0.023881108, 0.025360467, 0.009770825, 0.03792799, -0.022429924, 0.033363104, -0.0018914682, 0.04040768, 0.018484963, 0.0070199184, -0.017583257, 0.016258877, 0.010954313, -0.008939565, -0.024148801, -0.02498006, -0.007889924, 0.02748793, 0.0307707, 0.029756282, 0.0051425393, 0.0045719286, -0.03046074, 0.013596028, 0.025684519, -0.0033197557, 0.006967084, 0.03677268, 0.0120039545, -0.0032792494, -0.0032211316, -0.02399382, -0.026924362, -0.013920079, -0.0042197, 0.025346378, -0.0027015943, -0.016991513, 0.0031594916, -0.007579962, 0.018978084, 0.017681882, 0.0126591, 0.028939111, 0.008833896, 0.10183637, 0.0059632324, -0.05196078, -0.023697948, 0.011045893, -0.008777539, -0.013807366, 0.019273955, -0.025346378, 0.0074742935, 0.009961028, -0.010813422, 0.018597675, 0.009636978, 0.014948587, 0.024064265, -0.008693005, -0.020570157, 0.014194817, -0.026219906, -0.02299349, 0.011067026, 0.032066904, 0.013391736, -0.05148175, -0.009489042, -0.03062981, -0.0012847543, 0.07286908, 0.026529867, -0.00025008238, 0.013638296, 0.016089808, 0.018654034, -0.0020394044, -0.024543297, 0.0147795165, -0.009601755, 0.0018791402, -0.040520392, -0.003360262, 0.02216223, 0.0137650985, 0.0059914105, 0.0048361, -0.0009844792, 0.016977424, -0.00934815, 0.024233336, -0.013088819, -0.017555078, -0.0050263037, 0.010595039, -0.027516108, 0.0071537653, -0.023247095, -0.0017655464, -0.015948916, 0.058160007, -0.025966302, 0.0121941585, -0.012384362, -0.0015612538, 0.009946939, 0.00628376, 0.011327676, 0.0109613575, 0.008601425, -0.018329982, 0.055680316, -0.012778858, -0.0100807855, -0.011067026, -0.0036490895, -0.01356785, 0.0073193125, -0.014272308, -0.027403394, -0.030742522, 0.02862915, 0.03062981, -0.014596358, 0.021697288, 0.0042408337, 0.027572464, 0.0019601528, -0.037138995, -0.031306088, 0.041929305, 0.017738238, 0.004857234, 0.008256241, 0.0118278405, 0.021753646, -0.00160176, -0.0018333505, -0.0047374764, 0.042239267, 0.0058329077, -0.026459422, 0.015075389, 0.021147812, -0.005212985, 0.01281408, 0.017738238, 0.008242152, -0.020372909, -0.011081115, -0.011017715, 0.007706764, 0.01834407, 0.01954165, 0.037477136, -0.010278034, 0.015808025, 0.00031590514, -0.017681882, -0.008967743, -0.020612424, -0.025416825, -0.0037970257, -0.029868996, 0.01720285, -0.0144554665, 0.026727116, 0.00414221, 0.0040154075, 0.05838543, 0.0005622451, -0.025219576, 0.004180955, -0.002932304, -0.0090663675, 0.011574236, 0.02450103, -0.012553431, -0.020612424, -0.032095082, 0.015526242, 0.008974788, 0.0053151315, -0.0003112821, -0.017935487, -0.0076222294, 0.03358853, 0.029474499, -0.011496745, -0.012835215, -0.020739228, -0.012482986, -0.037871633, 0.0052517303, -0.012926794, -0.0025237189, 0.0020323596, 0.045113456, -0.04835396, -0.027755624, -0.0079955915, 0.007896968, 0.0072559114, 0.015047211, -0.0014573464, -0.014032792, 0.021091456, -0.0046071517, -0.0065232757, -0.02582541, -0.035870973, -0.015343083, 0.03254593, -0.028431902, -0.003286294, 0.014328664, 0.008840941, 0.015948916, 0.012835215, 0.019400757, -0.012342094, -0.010693664, 0.004772699, -0.03254593, 0.010707753, -0.016822444, -0.0032827717, 0.021246437, -0.04485985, -0.04384543, -0.015906649, -0.009707424, 0.02299349, 0.019513471, -0.010151232, 0.018963994, -0.0057976847, 0.05739919, -0.019922055, -0.029108182, -0.0106232185, 0.021077367, 0.0036455672, -0.026614401, 0.04497256, -0.04446535, -0.0004556959, -0.004578973, 0.003962573, -0.004910068, 0.015089478, -0.0301226, 0.007664497, 0.008375999, 0.031982366, 0.006135824, 0.02152822, -0.015469885, -0.007210122, 0.034715664, -0.01233505, 0.0004490916, -0.0144413775, -0.003150686, -0.02003477, -0.027924692, -0.0015850292, -0.009376328, -0.0035997776, -0.03240504, -0.010912046, 0.0031999978, 0.022303123, -0.008988877, 0.00024633997, -0.0035698381, 0.0070974086, -0.002599448, -0.042267445, -0.016935157, -0.0002481011, -0.041393917, 0.014483645, 0.019006262, -0.02813603, 0.0072030774, -7.3032425e-05, 0.01802002, -0.017188761, 0.015991183, 0.020401087, 0.03542012, 0.04469078, 0.04071764, 0.011095204, -0.031390622, -0.03254593, 0.014187773, 0.016272966, -0.009721513, -0.026388975, -0.014849963, -0.005642704, -0.022556728, 0.0064457855, -0.043450933, 0.010834555, -0.015977094, 0.020880118, -0.02385293, -0.054806788, 0.03789981, 0.0013516777, -0.026431242, -0.015540331, 0.016695641, -0.037167173, -0.021190079, 0.023881108, -0.0045860177, 0.0064105624, -0.007763121, -0.013053596, 0.024472851, -0.0004962022, -0.00976378, 0.060019772, -0.0057624616, -0.04384543, 0.010313257, 0.0076715415, 0.0025888812, -0.03589915, 0.008791628, -0.012785902, 0.01042597, 0.015653044, 0.04767768, -0.009869449, 0.0064457855, -0.010947268, -0.0077349427, -0.032715004, -0.023867019, -0.011327676, -0.00046274049, -0.036998104, 0.013913034, 0.012250515, -0.009996251, 0.021204168, 0.020091126, -0.003740669, -0.0049769916, -0.0140891485, 0.024064265, 0.0038815604, 0.025684519, 0.041788414, -0.013553761, 0.006681779, -0.0050826604, -0.018175002, 0.008228063, -0.006230926, -0.018907638, 0.0154839745, -0.028713685, -0.015047211, -0.019682541, 0.02516322, 0.040802173, 0.007213644, 0.011743305, -0.015963005, -0.03818159, 0.01191942, -0.031728763, -0.011863063, 0.023881108, 0.0053116092, -0.020992832, -0.017991843, -0.00405063, -0.017780505, -0.0057659843, 0.02978446, 0.031165197, 0.0014221234, 0.021316882, 0.026008569, -0.0018544842, -0.032658648, 0.028474169, 0.013109953, 0.018076377, 0.0007991189, -0.0042373114, 0.028910933, -0.0029358263, 0.021866359, 0.024472851, -0.002576553, -0.033532172, 0.01920351, -0.0095665315, -0.03093977, 0.0034817809, 0.018654034, -0.0074038478, 0.021443684, 0.0038604268, -0.02745975, 0.031587873, 0.0061146906, 0.022711707, -0.019795254, -0.016991513, -0.04471896, -0.007875834, -0.0034941088, -0.043789074, 0.021091456, 0.024909616, -0.013194487, -0.0042690123, 0.027896514, -0.018414518, -0.023303451, -0.025797231, -0.009524264]}], object='list', usage=Usage(completion_tokens=0, prompt_tokens=3, total_tokens=3, completion_tokens_details=None, prompt_tokens_details=None))\n" + ] + } + ], + "source": [ + "import litellm\n", + "\n", + "\n", + "async def main():\n", + " response = await litellm.aembedding(\n", + " model=\"cometapi/text-embedding-3-small\", # The model name must include prefix \"cometapi/\" + the model name from CometAPI\n", + " api_key=api_key, # your CometAPI api-key\n", + " api_base=\"https://api.cometapi.com/v1\",\n", + " input=\"Your text string\",\n", + " )\n", + " print(response)\n", + "\n", + "\n", + "await main()" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Async Image Generation" + ] + }, + { + "cell_type": "code", + "execution_count": 14, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "ImageResponse(created=1760591151, background=None, data=[ImageObject(b64_json=None, revised_prompt=\"Generate an image of an adorable baby sea otter. It should be floating on its back in a calm, clear ocean, playfully grasping a colorful shell in its small paws. The sun is setting in the background, casting a peaceful orange and purple hue across the sky and reflecting upon the ocean waves. The otter's fur is a deep, rich brown and appears silky and wet, with glints of sunlight catching on it. Its eyes are bright, expressing joy and curiosity as it examines its newfound treasure.\", url='https://oaidalleapiprodscus.blob.core.windows.net/private/org-OKnsK88id12jfvnKByup1O0l/user-3GxuMyEg9YMU8LFCPHi31prf/img-7PUEF8Wb6thGDAuZWLJjSnfP.png?st=2025-10-16T04%3A05%3A51Z&se=2025-10-16T06%3A05%3A51Z&sp=r&sv=2024-08-04&sr=b&rscd=inline&rsct=image/png&skoid=38e27a3b-6174-4d3e-90ac-d7d9ad49543f&sktid=a48cca56-e6da-484e-a814-9c849652bcb3&skt=2025-10-16T02%3A51%3A01Z&ske=2025-10-17T02%3A51%3A01Z&sks=b&skv=2024-08-04&sig=IZKG2VE%2B6VdOe5Tq0Zk/5bVyGK/oK/yO8g%2BDX4krpug%3D')], output_format=None, quality=None, size=None, usage=Usage(completion_tokens=0, prompt_tokens=0, total_tokens=0, completion_tokens_details=None, prompt_tokens_details=None, input_tokens=0, input_tokens_details={'image_tokens': 0, 'text_tokens': 0}, output_tokens=0))\n" + ] + } + ], + "source": [ + "import asyncio\n", + "\n", + "import litellm\n", + "\n", + "\n", + "async def main():\n", + " response = await litellm.aimage_generation(\n", + " model=\"cometapi/dall-e-3\", # The model name must include prefix \"cometapi/\" + the model name from CometAPI\n", + " api_key=api_key, # your cometapi api-key\n", + " api_base=\"https://api.cometapi.com/v1\",\n", + " prompt=\"A cute baby sea otter\",\n", + " )\n", + " print(response)\n", + "\n", + "\n", + "await main()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "colab": { + "provenance": [] + }, + "kernelspec": { + "display_name": "base", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.8" + } + }, + "nbformat": 4, + "nbformat_minor": 0 +} diff --git a/docs/my-website/docs/bedrock_converse.md b/docs/my-website/docs/bedrock_converse.md new file mode 100644 index 00000000000..cf66b1a50a6 --- /dev/null +++ b/docs/my-website/docs/bedrock_converse.md @@ -0,0 +1,151 @@ +# /converse + +Call Bedrock's `/converse` endpoint through LiteLLM Proxy. + +| Feature | Supported | +|---------|-----------| +| Cost Tracking | ✅ | +| Logging | ✅ | +| Streaming | ✅ via `/converse-stream` | +| Load Balancing | ✅ | + +## Quick Start + +### 1. Setup config.yaml + +```yaml showLineNumbers +model_list: + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID # reads from environment + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + custom_llm_provider: bedrock +``` + +Set AWS credentials in your environment: + +```bash showLineNumbers +export AWS_ACCESS_KEY_ID="your-access-key" +export AWS_SECRET_ACCESS_KEY="your-secret-key" +``` + +### 2. Start Proxy + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +### 3. Call /converse endpoint + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "messages": [ + { + "role": "user", + "content": [{"text": "Hello, how are you?"}] + } + ], + "inferenceConfig": { + "temperature": 0.5, + "maxTokens": 100 + } +}' +``` + +## Streaming + +For streaming responses, use `/converse-stream`: + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse-stream' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "messages": [ + { + "role": "user", + "content": [{"text": "Tell me a short story"}] + } + ], + "inferenceConfig": { + "temperature": 0.7, + "maxTokens": 200 + } +}' +``` + +## Load Balancing + +Define multiple deployments with the same `model_name` for automatic load balancing: + +```yaml showLineNumbers +model_list: + # Deployment 1 - us-west-2 + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + custom_llm_provider: bedrock + + # Deployment 2 - us-east-1 + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-east-1 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + custom_llm_provider: bedrock +``` + +The proxy automatically distributes requests across both regions. + +## Using boto3 SDK + +```python showLineNumbers +import boto3 +import json +import os + +# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy) +os.environ['AWS_ACCESS_KEY_ID'] = 'dummy' +os.environ['AWS_SECRET_ACCESS_KEY'] = 'dummy' +os.environ['AWS_BEARER_TOKEN_BEDROCK'] = "sk-1234" # your litellm proxy api key + +# Point boto3 to the LiteLLM proxy +bedrock_runtime = boto3.client( + service_name='bedrock-runtime', + region_name='us-west-2', + endpoint_url='http://0.0.0.0:4000/bedrock' +) + +response = bedrock_runtime.converse( + modelId='my-bedrock-model', # Your model_name from config.yaml + messages=[ + { + "role": "user", + "content": [{"text": "Hello, how are you?"}] + } + ], + inferenceConfig={ + "temperature": 0.5, + "maxTokens": 100 + } +) + +print(response['output']['message']['content'][0]['text']) +``` + +## More Info + +For complete documentation including Guardrails, Knowledge Bases, and Agents, see: +- [Full Bedrock Passthrough Docs](./pass_through/bedrock) + diff --git a/docs/my-website/docs/bedrock_invoke.md b/docs/my-website/docs/bedrock_invoke.md new file mode 100644 index 00000000000..6f29f1d51c3 --- /dev/null +++ b/docs/my-website/docs/bedrock_invoke.md @@ -0,0 +1,145 @@ +# /invoke + +Call Bedrock's `/invoke` endpoint through LiteLLM Proxy. + +| Feature | Supported | +|---------|-----------| +| Cost Tracking | ✅ | +| Logging | ✅ | +| Streaming | ✅ via `/invoke-with-response-stream` | +| Load Balancing | ✅ | + +## Quick Start + +### 1. Setup config.yaml + +```yaml showLineNumbers +model_list: + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID # reads from environment + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + custom_llm_provider: bedrock +``` + +Set AWS credentials in your environment: + +```bash showLineNumbers +export AWS_ACCESS_KEY_ID="your-access-key" +export AWS_SECRET_ACCESS_KEY="your-secret-key" +``` + +### 2. Start Proxy + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +### 3. Call /invoke endpoint + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/invoke' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "max_tokens": 100, + "messages": [ + { + "role": "user", + "content": "Hello, how are you?" + } + ], + "anthropic_version": "bedrock-2023-05-31" +}' +``` + +## Streaming + +For streaming responses, use `/invoke-with-response-stream`: + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/invoke-with-response-stream' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "max_tokens": 100, + "messages": [ + { + "role": "user", + "content": "Tell me a short story" + } + ], + "anthropic_version": "bedrock-2023-05-31" +}' +``` + +## Load Balancing + +Define multiple deployments with the same `model_name` for automatic load balancing: + +```yaml showLineNumbers +model_list: + # Deployment 1 - us-west-2 + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + custom_llm_provider: bedrock + + # Deployment 2 - us-east-1 + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-east-1 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + custom_llm_provider: bedrock +``` + +The proxy automatically distributes requests across both regions. + +## Using boto3 SDK + +```python showLineNumbers +import boto3 +import json +import os + +# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy) +os.environ['AWS_ACCESS_KEY_ID'] = 'dummy' +os.environ['AWS_SECRET_ACCESS_KEY'] = 'dummy' +os.environ['AWS_BEARER_TOKEN_BEDROCK'] = "sk-1234" # your litellm proxy api key + +# Point boto3 to the LiteLLM proxy +bedrock_runtime = boto3.client( + service_name='bedrock-runtime', + region_name='us-west-2', + endpoint_url='http://0.0.0.0:4000/bedrock' +) + +response = bedrock_runtime.invoke_model( + modelId='my-bedrock-model', # Your model_name from config.yaml + contentType='application/json', + accept='application/json', + body=json.dumps({ + "max_tokens": 100, + "messages": [{"role": "user", "content": "Hello"}], + "anthropic_version": "bedrock-2023-05-31" + }) +) + +response_body = json.loads(response['body'].read()) +print(response_body['content'][0]['text']) +``` + +## More Info + +For complete documentation including Guardrails, Knowledge Bases, and Agents, see: +- [Full Bedrock Passthrough Docs](./pass_through/bedrock) + diff --git a/docs/my-website/docs/guides/security_settings.md b/docs/my-website/docs/guides/security_settings.md index 7995f6c3c9c..d6397a7c197 100644 --- a/docs/my-website/docs/guides/security_settings.md +++ b/docs/my-website/docs/guides/security_settings.md @@ -117,10 +117,52 @@ litellm_settings: ```bash export SSL_CERTIFICATE="/path/to/certificate.pem" ``` + -## 5. Use HTTP_PROXY environment variable +## 5. Configure ECDH Curve for SSL/TLS Performance + +The `ssl_ecdh_curve` setting allows you to configure the Elliptic Curve Diffie-Hellman (ECDH) curve used for SSL/TLS key exchange. This is particularly useful for disabling Post-Quantum Cryptography (PQC) to improve performance in environments where PQC is not required. + +**Use Case:** Some OpenSSL 3.x systems enable PQC by default, which can slow down TLS handshakes. Setting the ECDH curve to `X25519` disables PQC and can significantly improve connection performance. + + + + +```python +import litellm +litellm.ssl_ecdh_curve = "X25519" # Disables PQC for better performance +``` + + + + +```yaml +litellm_settings: + ssl_ecdh_curve: "X25519" +``` + + + + +```bash +export SSL_ECDH_CURVE="X25519" +``` + + + + +**Common Valid Curves:** + +- `X25519` - Modern, fast curve (recommended for disabling PQC) +- `prime256v1` - NIST P-256 curve +- `secp384r1` - NIST P-384 curve +- `secp521r1` - NIST P-521 curve + +**Note:** If an invalid curve name is provided or if your Python/OpenSSL version doesn't support this feature, LiteLLM will log a warning and continue with default curves. + +## 6. Use HTTP_PROXY environment variable Both httpx and aiohttp libraries use `urllib.request.getproxies` from environment variables. Before client initialization, you may set proxy (and optional SSL_CERT_FILE) by setting the environment variables: diff --git a/docs/my-website/docs/ocr.md b/docs/my-website/docs/ocr.md new file mode 100644 index 00000000000..d966097b326 --- /dev/null +++ b/docs/my-website/docs/ocr.md @@ -0,0 +1,257 @@ +# /ocr + +:::tip + +LiteLLM follows the [Mistral API request/response for the OCR API](https://docs.mistral.ai/capabilities/vision/#optical-character-recognition-ocr) + +::: + +## **LiteLLM Python SDK Usage** +### Quick Start + +```python +from litellm import ocr +import os + +os.environ["MISTRAL_API_KEY"] = "sk-.." + +response = ocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + } +) + +# Access extracted text +for page in response.pages: + print(f"Page {page.index}:") + print(page.markdown) +``` + +### Async Usage + +```python +from litellm import aocr +import os, asyncio + +os.environ["MISTRAL_API_KEY"] = "sk-.." + +async def test_async_ocr(): + response = await aocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + } + ) + + # Access extracted text + for page in response.pages: + print(f"Page {page.index}:") + print(page.markdown) + +asyncio.run(test_async_ocr()) +``` + +### Using Base64 Encoded Documents + +```python +import base64 +from litellm import ocr + +# Encode PDF to base64 +with open("document.pdf", "rb") as f: + base64_pdf = base64.b64encode(f.read()).decode('utf-8') + +response = ocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": f"data:application/pdf;base64,{base64_pdf}" + } +) +``` + +### Optional Parameters + +```python +response = ocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }, + # Optional Mistral parameters + pages=[0, 1, 2], # Only process specific pages + include_image_base64=True, # Include extracted images + image_limit=10, # Max images to return + image_min_size=100 # Min image size to include +) +``` + +## **LiteLLM Proxy Usage** + +LiteLLM provides a Mistral API compatible `/ocr` endpoint for OCR calls. + +**Setup** + +Add this to your litellm proxy config.yaml + +```yaml +model_list: + - model_name: mistral-ocr + litellm_params: + model: mistral/mistral-ocr-latest + api_key: os.environ/MISTRAL_API_KEY +``` + +Start litellm + +```bash +litellm --config /path/to/config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +Test request + +```bash +curl http://0.0.0.0:4000/v1/ocr \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "mistral-ocr", + "document": { + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + } + }' +``` + + +## **Request/Response Format** + +:::info + +LiteLLM follows the **Mistral OCR API specification**. + +See the [official Mistral OCR documentation](https://docs.mistral.ai/capabilities/vision/#optical-character-recognition-ocr) for complete details. + +::: + +### Example Request + +```python +{ + "model": "mistral/mistral-ocr-latest", + "document": { + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + }, + "pages": [0, 1, 2], # Optional: specific pages to process + "include_image_base64": True, # Optional: include extracted images + "image_limit": 10, # Optional: max images to return + "image_min_size": 100 # Optional: min image size in pixels +} +``` + +### Request Parameters + +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `model` | string | Yes | The OCR model to use (e.g., `"mistral/mistral-ocr-latest"`) | +| `document` | object | Yes | Document to process. Must contain `type` and URL field | +| `document.type` | string | Yes | Either `"document_url"` for PDFs/docs or `"image_url"` for images | +| `document.document_url` | string | Conditional | URL to the document (required if `type` is `"document_url"`) | +| `document.image_url` | string | Conditional | URL to the image (required if `type` is `"image_url"`) | +| `pages` | array | No | List of specific page indices to process (0-indexed) | +| `include_image_base64` | boolean | No | Whether to include extracted images as base64 strings | +| `image_limit` | integer | No | Maximum number of images to return | +| `image_min_size` | integer | No | Minimum size (in pixels) for images to include | + +#### Document Format Examples + +**For PDFs and documents:** +```json +{ + "type": "document_url", + "document_url": "https://example.com/document.pdf" +} +``` + +**For images:** +```json +{ + "type": "image_url", + "image_url": "https://example.com/image.png" +} +``` + +**For base64-encoded content:** +```json +{ + "type": "document_url", + "document_url": "data:application/pdf;base64,JVBERi0xLjQKJ..." +} +``` + +### Response Format + +The response follows Mistral's OCR format with the following structure: + +```json +{ + "pages": [ + { + "index": 0, + "markdown": "# Document Title\n\nExtracted text content...", + "dimensions": { + "dpi": 200, + "height": 2200, + "width": 1700 + }, + "images": [ + { + "image_base64": "base64string...", + "bbox": { + "x": 100, + "y": 200, + "width": 300, + "height": 400 + } + } + ] + } + ], + "model": "mistral-ocr-2505-completion", + "usage_info": { + "pages_processed": 29, + "doc_size_bytes": 3002783 + }, + "document_annotation": null, + "object": "ocr" +} +``` + +#### Response Fields + +| Field | Type | Description | +|-------|------|-------------| +| `pages` | array | List of processed pages with extracted content | +| `pages[].index` | integer | Page number (0-indexed) | +| `pages[].markdown` | string | Extracted text in Markdown format | +| `pages[].dimensions` | object | Page dimensions (dpi, height, width in pixels) | +| `pages[].images` | array | Extracted images from the page (if `include_image_base64=true`) | +| `model` | string | The model used for OCR processing | +| `usage_info` | object | Processing statistics (pages processed, document size) | +| `document_annotation` | object | Optional document-level annotations | +| `object` | string | Always `"ocr"` for OCR responses | + + +## **Supported Providers** + +| Provider | Link to Usage | +|-------------|--------------------| +| Mistral AI | [Usage](#quick-start) | + diff --git a/docs/my-website/docs/pass_through/bedrock.md b/docs/my-website/docs/pass_through/bedrock.md index 48502864d78..b8d20d77da0 100644 --- a/docs/my-website/docs/pass_through/bedrock.md +++ b/docs/my-website/docs/pass_through/bedrock.md @@ -5,24 +5,55 @@ Pass-through endpoints for Bedrock - call provider-specific endpoint, in native | Feature | Supported | Notes | |-------|-------|-------| | Cost Tracking | ✅ | For `/invoke` and `/converse` endpoints | -| Logging | ✅ | works across all integrations | +| Load Balancing | ✅ | You can load balance `/invoke`, `/converse` routes across multiple deployments| Logging | ✅ | works across all integrations | | End-user Tracking | ❌ | [Tell us if you need this](https://github.com/BerriAI/litellm/issues/new) | | Streaming | ✅ | | Just replace `https://bedrock-runtime.{aws_region_name}.amazonaws.com` with `LITELLM_PROXY_BASE_URL/bedrock` 🚀 -#### **Example Usage** -```bash -curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \ --H 'Authorization: Bearer anything' \ +## Overview + +LiteLLM supports two ways to call Bedrock endpoints: + +### 1. **Using config.yaml** (Recommended for model endpoints) + +Define your Bedrock models in `config.yaml` and reference them by name. The proxy handles authentication and routing. + +**Use for**: `/converse`, `/converse-stream`, `/invoke`, `/invoke-with-response-stream` + +```yaml showLineNumbers +model_list: + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock +``` + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse' \ +-H 'Authorization: Bearer sk-1234' \ -H 'Content-Type: application/json' \ --d '{ - "messages": [ - {"role": "user", - "content": [{"text": "Hello"}] - } - ] -}' +-d '{"messages": [{"role": "user", "content": [{"text": "Hello"}]}]}' +``` + +### 2. **Direct passthrough** (For non-model endpoints) + +Set AWS credentials via environment variables and call Bedrock endpoints directly. + +**Use for**: Guardrails, Knowledge Bases, Agents, and other non-model endpoints + +```bash showLineNumbers +export AWS_ACCESS_KEY_ID="" +export AWS_SECRET_ACCESS_KEY="" +export AWS_REGION_NAME="us-west-2" +``` + +```bash showLineNumbers +curl "http://0.0.0.0:4000/bedrock/guardrail/my-guardrail-id/version/1/apply" \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{"contents": [{"text": {"text": "Hello"}}], "source": "INPUT"}' ``` Supports **ALL** Bedrock Endpoints (including streaming). @@ -33,39 +64,235 @@ Supports **ALL** Bedrock Endpoints (including streaming). Let's call the Bedrock [`/converse` endpoint](https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_Converse.html) -1. Add AWS Keys to your environment +1. Create a `config.yaml` file with your Bedrock model -```bash +```yaml showLineNumbers +model_list: + - model_name: my-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock +``` + +Set your AWS credentials: + +```bash showLineNumbers export AWS_ACCESS_KEY_ID="" # Access key export AWS_SECRET_ACCESS_KEY="" # Secret access key -export AWS_REGION_NAME="" # us-east-1, us-east-2, us-west-1, us-west-2 ``` 2. Start LiteLLM Proxy -```bash -litellm +```bash showLineNumbers +litellm --config config.yaml # RUNNING on http://0.0.0.0:4000 ``` 3. Test it! -Let's call the Bedrock converse endpoint +Let's call the Bedrock converse endpoint using the model name from config: -```bash -curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \ --H 'Authorization: Bearer anything' \ +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-bedrock-model/converse' \ +-H 'Authorization: Bearer sk-1234' \ -H 'Content-Type: application/json' \ -d '{ "messages": [ - {"role": "user", - "content": [{"text": "Hello"}] + { + "role": "user", + "content": [{"text": "Hello, how are you?"}] + } + ], + "inferenceConfig": { + "maxTokens": 100 } - ] }' ``` +## Setup with config.yaml + +Use config.yaml to define Bedrock models and use them via passthrough endpoints. + +### 1. Define models in config.yaml + +```yaml showLineNumbers +model_list: + - model_name: my-claude-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock + + - model_name: my-cohere-model + litellm_params: + model: bedrock/cohere.command-r-v1:0 + aws_region_name: us-east-1 + custom_llm_provider: bedrock +``` + +### 2. Start proxy with config + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +### 3. Call Bedrock Converse endpoint + +Use the `model_name` from config in the URL path: + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-claude-model/converse' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "messages": [ + { + "role": "user", + "content": [{"text": "Hello, how are you?"}] + } + ], + "inferenceConfig": { + "temperature": 0.5, + "maxTokens": 100 + } +}' +``` + +### 4. Call Bedrock Converse Stream endpoint + +For streaming responses, use the `/converse-stream` endpoint: + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-claude-model/converse-stream' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "messages": [ + { + "role": "user", + "content": [{"text": "Tell me a short story"}] + } + ], + "inferenceConfig": { + "temperature": 0.7, + "maxTokens": 200 + } +}' +``` + +### Supported Bedrock Endpoints with config.yaml + +When using models from config.yaml, you can call any Bedrock endpoint: + +| Endpoint | Description | Example | +|----------|-------------|---------| +| `/model/{model_name}/converse` | Converse API | `http://0.0.0.0:4000/bedrock/model/my-claude-model/converse` | +| `/model/{model_name}/converse-stream` | Streaming Converse | `http://0.0.0.0:4000/bedrock/model/my-claude-model/converse-stream` | +| `/model/{model_name}/invoke` | Legacy Invoke API | `http://0.0.0.0:4000/bedrock/model/my-claude-model/invoke` | +| `/model/{model_name}/invoke-with-response-stream` | Legacy Streaming | `http://0.0.0.0:4000/bedrock/model/my-claude-model/invoke-with-response-stream` | + +The proxy automatically resolves the `model_name` to the actual Bedrock model ID and region configured in your `config.yaml`. + +### Load Balancing Across Multiple Deployments + +Define multiple Bedrock deployments with the same `model_name` to enable automatic load balancing. + +#### 1. Define multiple deployments in config.yaml + +```yaml showLineNumbers +model_list: + # First deployment - us-west-2 + - model_name: my-claude-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock + + # Second deployment - us-east-1 (load balanced) + - model_name: my-claude-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-east-1 + custom_llm_provider: bedrock +``` + +#### 2. Start proxy with config + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +#### 3. Call the endpoint - requests are automatically load balanced + +```bash showLineNumbers +curl -X POST 'http://0.0.0.0:4000/bedrock/model/my-claude-model/invoke' \ +-H 'Authorization: Bearer sk-1234' \ +-H 'Content-Type: application/json' \ +-d '{ + "max_tokens": 100, + "messages": [ + { + "role": "user", + "content": "Hello, how are you?" + } + ], + "anthropic_version": "bedrock-2023-05-31" +}' +``` + +The proxy will automatically distribute requests across both `us-west-2` and `us-east-1` deployments. This works for all Bedrock endpoints: `/invoke`, `/invoke-with-response-stream`, `/converse`, and `/converse-stream`. + +#### Using boto3 SDK with load balancing + +You can also call the load-balanced endpoint using the boto3 SDK: + +```python showLineNumbers +import boto3 +import json +import os + +# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy) +os.environ['AWS_ACCESS_KEY_ID'] = 'dummy' +os.environ['AWS_SECRET_ACCESS_KEY'] = 'dummy' +os.environ['AWS_BEARER_TOKEN_BEDROCK'] = "sk-1234" # your litellm proxy api key + +# Point boto3 to the LiteLLM proxy +bedrock_runtime = boto3.client( + service_name='bedrock-runtime', + region_name='us-west-2', + endpoint_url='http://0.0.0.0:4000/bedrock' +) + +# Call the load-balanced model +response = bedrock_runtime.invoke_model( + modelId='my-claude-model', # Your model_name from config.yaml + contentType='application/json', + accept='application/json', + body=json.dumps({ + "max_tokens": 100, + "messages": [ + { + "role": "user", + "content": "Hello, how are you?" + } + ], + "anthropic_version": "bedrock-2023-05-31" + }) +) + +# Parse response +response_body = json.loads(response['body'].read()) +print(response_body['content'][0]['text']) +``` + +The proxy will automatically load balance your boto3 requests across all configured deployments. + ## Examples @@ -84,7 +311,7 @@ Key Changes: #### LiteLLM Proxy Call -```bash +```bash showLineNumbers curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \ -H 'Authorization: Bearer sk-anything' \ -H 'Content-Type: application/json' \ @@ -99,7 +326,7 @@ curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' #### Direct Bedrock API Call -```bash +```bash showLineNumbers curl -X POST 'https://bedrock-runtime.us-west-2.amazonaws.com/model/cohere.command-r-v1:0/converse' \ -H 'Authorization: AWS4-HMAC-SHA256..' \ -H 'Content-Type: application/json' \ @@ -114,9 +341,25 @@ curl -X POST 'https://bedrock-runtime.us-west-2.amazonaws.com/model/cohere.comma ### **Example 2: Apply Guardrail** +**Setup**: Set AWS credentials for direct passthrough + +```bash showLineNumbers +export AWS_ACCESS_KEY_ID="your-access-key" +export AWS_SECRET_ACCESS_KEY="your-secret-key" +export AWS_REGION_NAME="us-west-2" +``` + +Start proxy: + +```bash showLineNumbers +litellm + +# RUNNING on http://0.0.0.0:4000 +``` + #### LiteLLM Proxy Call -```bash +```bash showLineNumbers curl "http://0.0.0.0:4000/bedrock/guardrail/guardrailIdentifier/version/guardrailVersion/apply" \ -H 'Authorization: Bearer sk-anything' \ -H 'Content-Type: application/json' \ @@ -129,7 +372,7 @@ curl "http://0.0.0.0:4000/bedrock/guardrail/guardrailIdentifier/version/guardrai #### Direct Bedrock API Call -```bash +```bash showLineNumbers curl "https://bedrock-runtime.us-west-2.amazonaws.com/guardrail/guardrailIdentifier/version/guardrailVersion/apply" \ -H 'Authorization: AWS4-HMAC-SHA256..' \ -H 'Content-Type: application/json' \ @@ -142,7 +385,25 @@ curl "https://bedrock-runtime.us-west-2.amazonaws.com/guardrail/guardrailIdentif ### **Example 3: Query Knowledge Base** -```bash +**Setup**: Set AWS credentials for direct passthrough + +```bash showLineNumbers +export AWS_ACCESS_KEY_ID="your-access-key" +export AWS_SECRET_ACCESS_KEY="your-secret-key" +export AWS_REGION_NAME="us-west-2" +``` + +Start proxy: + +```bash showLineNumbers +litellm + +# RUNNING on http://0.0.0.0:4000 +``` + +#### LiteLLM Proxy Call + +```bash showLineNumbers curl -X POST "http://0.0.0.0:4000/bedrock/knowledgebases/{knowledgeBaseId}/retrieve" \ -H 'Authorization: Bearer sk-anything' \ -H 'Content-Type: application/json' \ @@ -163,7 +424,7 @@ curl -X POST "http://0.0.0.0:4000/bedrock/knowledgebases/{knowledgeBaseId}/retri #### Direct Bedrock API Call -```bash +```bash showLineNumbers curl -X POST "https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases/{knowledgeBaseId}/retrieve" \ -H 'Authorization: AWS4-HMAC-SHA256..' \ -H 'Content-Type: application/json' \ @@ -194,7 +455,7 @@ Use this, to avoid giving developers the raw AWS Keys, but still letting them us 1. Setup environment -```bash +```bash showLineNumbers export DATABASE_URL="" export LITELLM_MASTER_KEY="" export AWS_ACCESS_KEY_ID="" # Access key @@ -202,7 +463,7 @@ export AWS_SECRET_ACCESS_KEY="" # Secret access key export AWS_REGION_NAME="" # us-east-1, us-east-2, us-west-1, us-west-2 ``` -```bash +```bash showLineNumbers litellm # RUNNING on http://0.0.0.0:4000 @@ -210,7 +471,7 @@ litellm 2. Generate virtual key -```bash +```bash showLineNumbers curl -X POST 'http://0.0.0.0:4000/key/generate' \ -H 'Authorization: Bearer sk-1234' \ -H 'Content-Type: application/json' \ @@ -219,7 +480,7 @@ curl -X POST 'http://0.0.0.0:4000/key/generate' \ Expected Response -```bash +```bash showLineNumbers { ... "key": "sk-1234ewknldferwedojwojw" @@ -229,7 +490,7 @@ Expected Response 3. Test it! -```bash +```bash showLineNumbers curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' \ -H 'Authorization: Bearer sk-1234ewknldferwedojwojw' \ -H 'Content-Type: application/json' \ @@ -246,46 +507,46 @@ curl -X POST 'http://0.0.0.0:4000/bedrock/model/cohere.command-r-v1:0/converse' Call Bedrock Agents via LiteLLM proxy -```python +**Setup**: Set AWS credentials on your LiteLLM proxy server + +```bash showLineNumbers +export AWS_ACCESS_KEY_ID="your-access-key" +export AWS_SECRET_ACCESS_KEY="your-secret-key" +export AWS_REGION_NAME="us-west-2" +``` + +Start proxy: + +```bash showLineNumbers +litellm + +# RUNNING on http://0.0.0.0:4000 +``` + +**Usage from Python**: + +```python showLineNumbers import os -import boto3 -from botocore.config import Config - -# # Define your proxy endpoint -proxy_endpoint = "http://0.0.0.0:4000/bedrock" # 👈 your proxy base url - -# # Create a Config object with the proxy -# Custom headers -custom_headers = { - 'litellm_user_api_key': 'Bearer sk-1234', # 👈 your proxy api key -} - - -os.environ["AWS_ACCESS_KEY_ID"] = "my-fake-key-id" -os.environ["AWS_SECRET_ACCESS_KEY"] = "my-fake-access-key" +import boto3 +# Set dummy AWS credentials (required by boto3, but not used by LiteLLM proxy) +os.environ["AWS_ACCESS_KEY_ID"] = "dummy" +os.environ["AWS_SECRET_ACCESS_KEY"] = "dummy" +os.environ["AWS_BEARER_TOKEN_BEDROCK"] = "sk-1234" # your litellm proxy api key # Create the client runtime_client = boto3.client( service_name="bedrock-agent-runtime", region_name="us-west-2", - endpoint_url=proxy_endpoint + endpoint_url="http://0.0.0.0:4000/bedrock" ) -# Custom header injection -def inject_custom_headers(request, **kwargs): - request.headers.update(custom_headers) - -# Attach the event to inject custom headers before the request is sent -runtime_client.meta.events.register('before-send.*.*', inject_custom_headers) - - response = runtime_client.invoke_agent( - agentId="L1RT58GYRW", - agentAliasId="MFPSBCXYTW", - sessionId="12345", - inputText="Who do you know?" - ) + agentId="L1RT58GYRW", + agentAliasId="MFPSBCXYTW", + sessionId="12345", + inputText="Who do you know?" +) completion = "" @@ -294,5 +555,4 @@ for event in response.get("completion"): completion += chunk["bytes"].decode() print(completion) - ``` diff --git a/docs/my-website/docs/providers/cometapi.md b/docs/my-website/docs/providers/cometapi.md index 1245bacfad4..a7f6e65519d 100644 --- a/docs/my-website/docs/providers/cometapi.md +++ b/docs/my-website/docs/providers/cometapi.md @@ -1,6 +1,10 @@ # CometAPI LiteLLM supports all AI models from [CometAPI](https://www.cometapi.com/). CometAPI provides access to 500+ AI models through a unified API interface, including cutting-edge models like GPT-5, Claude Opus 4.1, and various other state-of-the-art language models. + + Open In Colab + + ## Authentication To use CometAPI models, you need to obtain an API key from [CometAPI Token Console](https://api.cometapi.com/console/token). CometAPI offers free tokens for new users - you can get your free API key instantly by registering. diff --git a/docs/my-website/docs/proxy/admin_ui_sso.md b/docs/my-website/docs/proxy/admin_ui_sso.md index bd18dd9c690..ae082848b6b 100644 --- a/docs/my-website/docs/proxy/admin_ui_sso.md +++ b/docs/my-website/docs/proxy/admin_ui_sso.md @@ -320,6 +320,16 @@ Okta requires the `GENERIC_CLIENT_STATE` parameter: GENERIC_CLIENT_STATE="random-string" # Required for Okta ``` +### Okta PKCE + +If your Okta application is configured to require PKCE (Proof Key for Code Exchange), enable it by setting: + +```bash +GENERIC_CLIENT_USE_PKCE="true" +``` + +This is required when your Okta app settings enforce PKCE for enhanced security. LiteLLM will automatically handle PKCE parameter generation and verification during the OAuth flow. + ### Common Configuration Issues #### Missing Protocol in Base URL diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 4e440857261..0dbecad9d4a 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -533,6 +533,7 @@ router_settings: | GENERIC_CLIENT_ID | Client ID for generic OAuth providers | GENERIC_CLIENT_SECRET | Client secret for generic OAuth providers | GENERIC_CLIENT_STATE | State parameter for generic client authentication +| GENERIC_CLIENT_USE_PKCE | Enable PKCE (Proof Key for Code Exchange) for generic OAuth providers. Set to "true" when your OAuth provider requires PKCE. **Default is false** | GENERIC_SSO_HEADERS | Comma-separated list of additional headers to add to the request - e.g. Authorization=Bearer ``, Content-Type=application/json, etc. | GENERIC_INCLUDE_CLIENT_ID | Include client ID in requests for OAuth | GENERIC_SCOPE | Scope settings for generic OAuth providers @@ -752,6 +753,7 @@ router_settings: | SPEND_LOGS_URL | URL for retrieving spend logs | SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000 | SSL_CERTIFICATE | Path to the SSL certificate file +| SSL_ECDH_CURVE | ECDH curve for SSL/TLS key exchange (e.g., 'X25519' to disable PQC). | SSL_SECURITY_LEVEL | [BETA] Security level for SSL/TLS connections. E.g. `DEFAULT@SECLEVEL=1` | SSL_VERIFY | Flag to enable or disable SSL certificate verification | SSL_CERT_FILE | Path to the SSL certificate file for custom CA bundle diff --git a/docs/my-website/docs/proxy/guardrails/pillar_security.md b/docs/my-website/docs/proxy/guardrails/pillar_security.md index c730da5b416..a5a416839f6 100644 --- a/docs/my-website/docs/proxy/guardrails/pillar_security.md +++ b/docs/my-website/docs/proxy/guardrails/pillar_security.md @@ -38,13 +38,17 @@ model_list: api_key: os.environ/OPENAI_API_KEY guardrails: - - guardrail_name: "pillar-minitor-everything" # you can change my name + - guardrail_name: "pillar-monitor-everything" # you can change my name litellm_params: guardrail: pillar mode: [pre_call, post_call] # Monitor both input and output api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "monitor" # Log threats but allow requests + persist_session: true # Keep conversations visible in Pillar dashboard + async_mode: false # Request synchronous verdicts + include_scanners: true # Return scanner category breakdown + include_evidence: true # Include detailed findings for triage default_on: true # Enable for all requests general_settings: @@ -104,10 +108,14 @@ guardrails: api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "block" # Block malicious requests + persist_session: true # Keep records for investigation + async_mode: false # Require an immediate verdict + include_scanners: true # Understand which rule triggered + include_evidence: true # Capture concrete evidence default_on: true # Enable for all requests general_settings: - master_key: "your-master-key-here" + master_key: "YOUR_LITELLM_PROXY_MASTER_KEY" litellm_settings: set_verbose: true @@ -136,10 +144,14 @@ guardrails: api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "monitor" # Log threats but allow requests + persist_session: false # Skip dashboard storage for low latency + async_mode: false # Still receive results inline + include_scanners: false # Minimal payload for performance + include_evidence: false # Omit details to keep responses light default_on: true # Enable for all requests general_settings: - master_key: "your-secure-master-key-here" + master_key: "YOUR_LITELLM_PROXY_MASTER_KEY" litellm_settings: set_verbose: true # Enable detailed logging @@ -169,10 +181,14 @@ guardrails: api_key: os.environ/PILLAR_API_KEY # Your Pillar API key api_base: os.environ/PILLAR_API_BASE # Pillar API endpoint on_flagged_action: "block" # Block threats on input and output + persist_session: true # Preserve conversations in Pillar dashboard + async_mode: false # Require synchronous approval + include_scanners: true # Inspect which scanners fired + include_evidence: true # Include detailed evidence for auditing default_on: true # Enable for all requests general_settings: - master_key: "your-secure-master-key-here" + master_key: "YOUR_LITELLM_PROXY_MASTER_KEY" litellm_settings: set_verbose: true # Enable detailed logging @@ -229,19 +245,139 @@ Logs the violation but allows the request to proceed: on_flagged_action: "monitor" ``` +## Advanced Configuration + +**Quick takeaways** +- Every request still runs *all* Pillar scanners; these options only change what comes back. +- Choose richer responses when you need audit trails, lighter responses when latency or cost matters. +- Blocking is controlled by LiteLLM’s `on_flagged_action` configuration—Pillar headers do not change block/monitor behaviour. + +Pillar Security executes the full scanner suite on each call. The settings below tune the Protect response headers LiteLLM sends, letting you balance fidelity, retention, and latency. + +### Response Control + +#### Data Retention (`persist_session`) +```yaml +persist_session: false # Default: true +``` +- **Why**: Controls whether Pillar stores session data for dashboard visibility. +- **Set false for**: Ephemeral testing, privacy-sensitive interactions. +- **Set true for**: Production monitoring, compliance, historical review (default behaviour). +- **Impact**: `false` means the conversation will *not* appear in the Pillar dashboard. + +#### Response Detail Level +The following toggles grow the payload size without changing detection behaviour. + +```yaml +include_scanners: true # → plr_scanners (default true in LiteLLM) +include_evidence: true # → plr_evidence (default true in LiteLLM) +``` + +- **Minimal response** (`include_scanners=false`, `include_evidence=false`) + ```json + { + "session_id": "abc-123", + "flagged": true + } + ``` + Use when you only care about whether Pillar detected a threat. + + > **📝 Note:** `flagged: true` means Pillar’s scanners recommend blocking. Pillar only reports this verdict—LiteLLM enforces your policy via the `on_flagged_action` configuration (no Pillar header controls it): + > - `on_flagged_action: "block"` → LiteLLM raises a 400 guardrail error + > - `on_flagged_action: "monitor"` → LiteLLM logs the threat but still returns the LLM response + +- **Scanner breakdown** (`include_scanners=true`) + ```json + { + "session_id": "abc-123", + "flagged": true, + "scanners": { + "jailbreak": true, + "prompt_injection": false, + "pii": false, + "secret": false, + "toxic_language": false + /* ... more categories ... */ + } + } + ``` + Use when you need to know which categories triggered. + +- **Full context** (both toggles true) + ```json + { + "session_id": "abc-123", + "flagged": true, + "scanners": { /* ... */ }, + "evidence": [ + { + "category": "jailbreak", + "type": "prompt_injection", + "evidence": "Ignore previous instructions", + "metadata": { "start_idx": 0, "end_idx": 28 } + } + ] + } + ``` + Ideal for debugging, audit logs, or compliance exports. + +### Processing Mode (`async_mode`) +```yaml +async_mode: true # Default: false +``` +- **Why**: Queue the request for background processing instead of waiting for a synchronous verdict. +- **Response shape**: + ```json + { + "status": "queued", + "session_id": "abc-123", + "position": 1 + } + ``` +- **Set true for**: Large batch jobs, latency-tolerant pipelines. +- **Set false for**: Real-time user flows (default). +- ⚠️ **Note**: Async mode returns only a 202 queue acknowledgment (no flagged verdict). LiteLLM treats that as “no block,” so the pre-call hook always allows the request. Use async mode only for post-call or monitor-only workflows where delayed review is acceptable. + +### Complete Examples + +```yaml +guardrails: + # Production: full fidelity & dashboard visibility + - guardrail_name: "pillar-production" + litellm_params: + guardrail: pillar + mode: [pre_call, post_call] + persist_session: true + include_scanners: true + include_evidence: true + on_flagged_action: "block" + + # Testing: lightweight, no persistence + - guardrail_name: "pillar-testing" + litellm_params: + guardrail: pillar + mode: pre_call + persist_session: false + include_scanners: false + include_evidence: false + on_flagged_action: "monitor" +``` + +Keep in mind that LiteLLM forwards these values as the documented `plr_*` headers, so any direct HTTP integrations outside the proxy can reuse the same guidance. + ## Examples -**Safe requset** +**Safe request** ```bash # Test with safe content curl -X POST "http://localhost:4000/v1/chat/completions" \ -H "Content-Type: application/json" \ - -H "Authorization: Bearer your-master-key-here" \ + -H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \ -d '{ "model": "gpt-4.1-mini", "messages": [{"role": "user", "content": "Hello! Can you tell me a joke?"}], @@ -300,7 +436,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \ ```bash curl -X POST "http://localhost:4000/v1/chat/completions" \ -H "Content-Type: application/json" \ - -H "Authorization: Bearer your-master-key-here" \ + -H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \ -d '{ "model": "gpt-4.1-mini", "messages": [ @@ -350,7 +486,7 @@ curl -X POST "http://localhost:4000/v1/chat/completions" \ ```bash curl -X POST "http://localhost:4000/v1/chat/completions" \ -H "Content-Type: application/json" \ - -H "Authorization: Bearer your-master-key-here" \ + -H "Authorization: Bearer YOUR_LITELLM_PROXY_MASTER_KEY" \ -d '{ "model": "gpt-4.1-mini", "messages": [ @@ -405,4 +541,4 @@ Feel free to contact us at support@pillar.security - [Pillar Security API Docs](https://docs.pillar.security/docs/api/introduction) - [Pillar Security Dashboard](https://app.pillar.security) - [Pillar Security Website](https://pillar.security) -- [LiteLLM Docs](https://docs.litellm.ai) \ No newline at end of file +- [LiteLLM Docs](https://docs.litellm.ai) diff --git a/docs/my-website/docs/tutorials/claude_responses_api.md b/docs/my-website/docs/tutorials/claude_responses_api.md index 098c311b30f..87d019f11b9 100644 --- a/docs/my-website/docs/tutorials/claude_responses_api.md +++ b/docs/my-website/docs/tutorials/claude_responses_api.md @@ -252,7 +252,7 @@ litellm --config /path/to/config.yaml 3. Use the MCP server in Claude Code ```bash -claude mcp add --transport http litellm_proxy http://0.0.0.0:4000 --header "Authorization: Bearer sk-LITELLM_VIRTUAL_KEY" +claude mcp add --transport http litellm_proxy http://0.0.0.0:4000/github_mcp/mcp --header "Authorization: Bearer sk-LITELLM_VIRTUAL_KEY" ``` 4. Authenticate via Claude Code diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 7de8e59f2b0..1c17216ac35 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -347,6 +347,9 @@ const sidebars = { ] }, "moderation", + "bedrock_invoke", + "bedrock_converse", + "ocr", { type: "category", label: "Pass-through Endpoints (Anthropic SDK, etc.)", @@ -536,6 +539,7 @@ const sidebars = { "providers/datarobot", "providers/ovhcloud", "providers/wandb_inference", + "providers/cometapi", ], }, { diff --git a/litellm/__init__.py b/litellm/__init__.py index 0dce8df13e7..5c0ba1c88da 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -263,6 +263,7 @@ use_client: bool = False ssl_verify: Union[str, bool] = True ssl_security_level: Optional[str] = None ssl_certificate: Optional[str] = None +ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance disable_streaming_logging: bool = False disable_token_counter: bool = False disable_add_transform_inline_image_block: bool = False @@ -1288,6 +1289,7 @@ from .llms.hyperbolic.chat.transformation import HyperbolicChatConfig from .llms.vercel_ai_gateway.chat.transformation import VercelAIGatewayConfig from .llms.ovhcloud.chat.transformation import OVHCloudChatConfig from .llms.ovhcloud.embedding.transformation import OVHCloudEmbeddingConfig +from .llms.cometapi.embed.transformation import CometAPIEmbeddingConfig from .llms.lemonade.chat.transformation import LemonadeChatConfig from .main import * # type: ignore from .integrations import * @@ -1325,6 +1327,7 @@ from .batch_completion.main import * # type: ignore from .rerank_api.main import * from .llms.anthropic.experimental_pass_through.messages.handler import * from .responses.main import * +from .ocr.main import * from .realtime_api.main import _arealtime from .fine_tuning.main import * from .files.main import * diff --git a/litellm/constants.py b/litellm/constants.py index 54ac3e6d6b8..d1858bdfd84 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -525,6 +525,7 @@ openai_compatible_providers: List = [ "vercel_ai_gateway", "aiml", "wandb", + "cometapi", ] openai_text_completion_compatible_providers: List = ( [ # providers that support `/v1/completions` @@ -849,6 +850,7 @@ BEDROCK_CONVERSE_MODELS = [ "deepseek.v3-v1:0", "openai.gpt-oss-20b-1:0", "openai.gpt-oss-120b-1:0", + "anthropic.claude-haiku-4-5-20251001-v1:0", "anthropic.claude-sonnet-4-5-20250929-v1:0", "anthropic.claude-opus-4-1-20250805-v1:0", "anthropic.claude-opus-4-20250514-v1:0", diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 626a3f3625f..2fd0b44962e 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -693,6 +693,15 @@ class CostCalculatorUtils: model=model, image_response=completion_response, ) + elif custom_llm_provider == litellm.LlmProviders.COMETAPI.value: + from litellm.llms.cometapi.image_generation.cost_calculator import ( + cost_calculator as cometapi_image_cost_calculator, + ) + + return cometapi_image_cost_calculator( + model=model, + image_response=completion_response, + ) elif custom_llm_provider == litellm.LlmProviders.GEMINI.value: from litellm.llms.gemini.image_generation.cost_calculator import ( cost_calculator as gemini_image_cost_calculator, diff --git a/litellm/llms/azure_ai/ocr/__init__.py b/litellm/llms/azure_ai/ocr/__init__.py new file mode 100644 index 00000000000..86f7e53d60b --- /dev/null +++ b/litellm/llms/azure_ai/ocr/__init__.py @@ -0,0 +1,5 @@ +"""Azure AI OCR module.""" +from .transformation import AzureAIOCRConfig + +__all__ = ["AzureAIOCRConfig"] + diff --git a/litellm/llms/azure_ai/ocr/transformation.py b/litellm/llms/azure_ai/ocr/transformation.py new file mode 100644 index 00000000000..eade2dd765f --- /dev/null +++ b/litellm/llms/azure_ai/ocr/transformation.py @@ -0,0 +1,268 @@ +""" +Azure AI OCR transformation implementation. +""" +from typing import Dict, Optional + +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.prompt_templates.image_handling import ( + async_convert_url_to_base64, + convert_url_to_base64, +) +from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData +from litellm.llms.mistral.ocr.transformation import MistralOCRConfig +from litellm.secret_managers.main import get_secret_str + + +class AzureAIOCRConfig(MistralOCRConfig): + """ + Azure AI OCR transformation configuration. + + Azure AI uses Mistral's OCR API but with a different endpoint format. + Inherits transformation logic from MistralOCRConfig since they use the same format. + + Reference: Azure AI Foundry OCR documentation + + Important: Azure AI only supports base64 data URIs (data:image/..., data:application/pdf;base64,...). + Regular URLs are not supported. + """ + + def __init__(self) -> None: + super().__init__() + + def validate_environment( + self, + headers: Dict, + model: str, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + **kwargs, + ) -> Dict: + """ + Validate environment and return headers for Azure AI OCR. + + Azure AI uses Bearer token authentication with AZURE_AI_API_KEY. + """ + # Get API key from environment if not provided + if api_key is None: + api_key = get_secret_str("AZURE_AI_API_KEY") + + if api_key is None: + raise ValueError( + "Missing Azure AI API Key - A call is being made to Azure AI but no key is set either in the environment variables or via params" + ) + + # Validate API base is provided + if api_base is None: + api_base = get_secret_str("AZURE_AI_API_BASE") + + if api_base is None: + raise ValueError( + "Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter" + ) + + headers = { + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + **headers, + } + + return headers + + def get_complete_url( + self, + api_base: Optional[str], + model: str, + optional_params: dict, + **kwargs, + ) -> str: + """ + Get complete URL for Azure AI OCR endpoint. + + Azure AI endpoint format: https:///providers/mistral/azure/ocr + + Args: + api_base: Azure AI API base URL + model: Model name (not used in URL construction) + optional_params: Optional parameters + + Returns: Complete URL for Azure AI OCR endpoint + """ + if api_base is None: + raise ValueError( + "Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter" + ) + + # Ensure no trailing slash + api_base = api_base.rstrip("/") + + # Azure AI OCR endpoint format + return f"{api_base}/providers/mistral/azure/ocr" + + def _convert_url_to_data_uri_sync(self, url: str) -> str: + """ + Synchronously convert a URL to a base64 data URI. + + Azure AI OCR doesn't have internet access, so we need to fetch URLs + and convert them to base64 data URIs. + + Args: + url: The URL to convert + + Returns: + Base64 data URI string + """ + verbose_logger.debug(f"Azure AI OCR: Converting URL to base64 data URI (sync): {url}") + + # Fetch and convert to base64 data URI + # convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..." + data_uri = convert_url_to_base64(url=url) + + verbose_logger.debug(f"Azure AI OCR: Converted URL to data URI (length: {len(data_uri)})") + + return data_uri + + async def _convert_url_to_data_uri_async(self, url: str) -> str: + """ + Asynchronously convert a URL to a base64 data URI. + + Azure AI OCR doesn't have internet access, so we need to fetch URLs + and convert them to base64 data URIs. + + Args: + url: The URL to convert + + Returns: + Base64 data URI string + """ + verbose_logger.debug(f"Azure AI OCR: Converting URL to base64 data URI (async): {url}") + + # Fetch and convert to base64 data URI asynchronously + # async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..." + data_uri = await async_convert_url_to_base64(url=url) + + verbose_logger.debug(f"Azure AI OCR: Converted URL to data URI (length: {len(data_uri)})") + + return data_uri + + def transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Transform OCR request for Azure AI, converting URLs to base64 data URIs (sync). + + Azure AI OCR doesn't have internet access, so we automatically fetch + any URLs and convert them to base64 data URIs synchronously. + + 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 + """ + verbose_logger.debug(f"Azure AI OCR transform_ocr_request (sync) - model: {model}") + + if not isinstance(document, dict): + raise ValueError(f"Expected document dict, got {type(document)}") + + # Check if we need to convert URL to base64 + doc_type = document.get("type") + transformed_document = document.copy() + + if doc_type == "document_url": + document_url = document.get("document_url", "") + # If it's not already a data URI, convert it + if document_url and not document_url.startswith("data:"): + verbose_logger.debug( + "Azure AI OCR: Converting document URL to base64 data URI (sync)" + ) + data_uri = self._convert_url_to_data_uri_sync(url=document_url) + transformed_document["document_url"] = data_uri + elif doc_type == "image_url": + image_url = document.get("image_url", "") + # If it's not already a data URI, convert it + if image_url and not image_url.startswith("data:"): + verbose_logger.debug( + "Azure AI OCR: Converting image URL to base64 data URI (sync)" + ) + data_uri = self._convert_url_to_data_uri_sync(url=image_url) + transformed_document["image_url"] = data_uri + + # Call parent's transform to build the request + return super().transform_ocr_request( + model=model, + document=transformed_document, + optional_params=optional_params, + headers=headers, + **kwargs, + ) + + async def async_transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Transform OCR request for Azure AI, converting URLs to base64 data URIs (async). + + Azure AI OCR doesn't have internet access, so we automatically fetch + any URLs and convert them to base64 data URIs asynchronously. + + 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 + """ + verbose_logger.debug(f"Azure AI OCR async_transform_ocr_request - model: {model}") + + if not isinstance(document, dict): + raise ValueError(f"Expected document dict, got {type(document)}") + + # Check if we need to convert URL to base64 + doc_type = document.get("type") + transformed_document = document.copy() + + if doc_type == "document_url": + document_url = document.get("document_url", "") + # If it's not already a data URI, convert it + if document_url and not document_url.startswith("data:"): + verbose_logger.debug( + "Azure AI OCR: Converting document URL to base64 data URI (async)" + ) + data_uri = await self._convert_url_to_data_uri_async(url=document_url) + transformed_document["document_url"] = data_uri + elif doc_type == "image_url": + image_url = document.get("image_url", "") + # If it's not already a data URI, convert it + if image_url and not image_url.startswith("data:"): + verbose_logger.debug( + "Azure AI OCR: Converting image URL to base64 data URI (async)" + ) + data_uri = await self._convert_url_to_data_uri_async(url=image_url) + transformed_document["image_url"] = data_uri + + # Call parent's transform to build the request + return super().transform_ocr_request( + model=model, + document=transformed_document, + optional_params=optional_params, + headers=headers, + **kwargs, + ) + diff --git a/litellm/llms/base_llm/ocr/__init__.py b/litellm/llms/base_llm/ocr/__init__.py new file mode 100644 index 00000000000..5965af5f2b7 --- /dev/null +++ b/litellm/llms/base_llm/ocr/__init__.py @@ -0,0 +1,22 @@ +"""Base OCR transformation module.""" +from .transformation import ( + BaseOCRConfig, + DocumentType, + OCRPage, + OCRPageDimensions, + OCRPageImage, + OCRRequestData, + OCRResponse, + OCRUsageInfo, +) + +__all__ = [ + "BaseOCRConfig", + "DocumentType", + "OCRResponse", + "OCRPage", + "OCRPageDimensions", + "OCRPageImage", + "OCRUsageInfo", + "OCRRequestData", +] diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py new file mode 100644 index 00000000000..41d7d31e6bc --- /dev/null +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -0,0 +1,207 @@ +""" +Base OCR transformation configuration. +""" +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union + +import httpx +from pydantic import BaseModel + +from litellm.llms.base_llm.chat.transformation import BaseLLMException + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +# DocumentType for OCR - Mistral format document dict +DocumentType = Dict[str, str] + + +class OCRPageDimensions(BaseModel): + """Page dimensions from OCR response.""" + dpi: Optional[int] = None + height: Optional[int] = None + width: Optional[int] = None + + +class OCRPageImage(BaseModel): + """Image extracted from OCR page.""" + image_base64: Optional[str] = None + bbox: Optional[Dict[str, Any]] = None + + model_config = {"extra": "allow"} + + +class OCRPage(BaseModel): + """Single page from OCR response.""" + index: int + markdown: str + images: Optional[List[OCRPageImage]] = None + dimensions: Optional[OCRPageDimensions] = None + + model_config = {"extra": "allow"} + + +class OCRUsageInfo(BaseModel): + """Usage information from OCR response.""" + pages_processed: Optional[int] = None + doc_size_bytes: Optional[int] = None + + model_config = {"extra": "allow"} + + +class OCRResponse(BaseModel): + """ + Standard OCR response format. + Standardized to Mistral OCR format - other providers should transform to this format. + """ + pages: List[OCRPage] + model: str + document_annotation: Optional[Any] = None + usage_info: Optional[OCRUsageInfo] = None + object: str = "ocr" + + model_config = {"extra": "allow"} + + +class OCRRequestData(BaseModel): + """OCR request data structure.""" + data: Optional[Union[Dict, bytes]] = None + files: Optional[Dict[str, Any]] = None + + +class BaseOCRConfig: + """ + Base configuration for OCR transformations. + Handles provider-agnostic OCR operations. + """ + + def __init__(self) -> None: + pass + + def get_supported_ocr_params(self, model: str) -> list: + """ + Get supported OCR parameters for this provider. + Override this method in provider-specific implementations. + """ + return [] + + def map_ocr_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + ) -> dict: + """Map OCR parameters to provider-specific parameters.""" + return optional_params + + def validate_environment( + self, + headers: Dict, + model: str, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + **kwargs, + ) -> Dict: + """ + Validate environment and return headers. + Override in provider-specific implementations. + """ + return headers + + def get_complete_url( + self, + api_base: Optional[str], + model: str, + optional_params: dict, + **kwargs, + ) -> str: + """ + Get complete URL for OCR endpoint. + Override in provider-specific implementations. + """ + raise NotImplementedError("get_complete_url must be implemented by provider") + + def transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Transform OCR request to provider-specific format. + Override in provider-specific implementations. + + Args: + model: Model name + document: Document to process (Mistral format dict, or file path, bytes, etc.) + optional_params: Optional parameters for the request + headers: Request headers + + Returns: + OCRRequestData with data and files fields + """ + raise NotImplementedError("transform_ocr_request must be implemented by provider") + + async def async_transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Async transform OCR request to provider-specific format. + Optional method - providers can override if they need async transformations + (e.g., Azure AI for URL-to-base64 conversion). + + Default implementation falls back to sync transform_ocr_request. + + Args: + model: Model name + document: Document to process (Mistral format dict, or file path, bytes, etc.) + optional_params: Optional parameters for the request + headers: Request headers + + Returns: + OCRRequestData with data and files fields + """ + # Default implementation: call sync version + 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 provider-specific OCR response to standard format. + Override in provider-specific implementations. + """ + raise NotImplementedError("transform_ocr_response must be implemented by provider") + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: dict, + ) -> Exception: + """Get appropriate error class for the provider.""" + return BaseLLMException( + status_code=status_code, + message=error_message, + headers=headers, + ) + diff --git a/litellm/llms/cometapi/embed/__init__.py b/litellm/llms/cometapi/embed/__init__.py new file mode 100644 index 00000000000..a36647f46c6 --- /dev/null +++ b/litellm/llms/cometapi/embed/__init__.py @@ -0,0 +1,3 @@ +from .transformation import CometAPIEmbeddingConfig + +__all__ = ["CometAPIEmbeddingConfig"] diff --git a/litellm/llms/cometapi/embed/transformation.py b/litellm/llms/cometapi/embed/transformation.py new file mode 100644 index 00000000000..5cfd1253149 --- /dev/null +++ b/litellm/llms/cometapi/embed/transformation.py @@ -0,0 +1,157 @@ +""" +CometAPI Embedding API support - OpenAI compatible +""" + +from typing import List, Optional, Union + +import httpx + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues +from litellm.types.utils import EmbeddingResponse, Usage + +from ..common_utils import CometAPIException + + +class CometAPIEmbeddingConfig(BaseEmbeddingConfig): + """ + Configuration class for CometAPI Embedding API. + + Since CometAPI is OpenAI-compatible, this class provides OpenAI-standard + embedding functionality with CometAPI-specific authentication and endpoints. + """ + + def __init__(self) -> None: + pass + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + Get the complete URL for the CometAPI embedding endpoint. + """ + api_base = ( + "https://api.cometapi.com/v1" if api_base is None else api_base.rstrip("/") + ) + complete_url = f"{api_base}/embeddings" + return complete_url + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """ + Validate and set up authentication headers for CometAPI. + """ + if api_key is None: + api_key = get_secret_str("COMETAPI_KEY") + + default_headers = { + "Authorization": f"Bearer {api_key}", + "accept": "application/json", + "Content-Type": "application/json", + } + + if "Authorization" in headers: + default_headers["Authorization"] = headers["Authorization"] + + return {**default_headers, **headers} + + def get_supported_openai_params(self, model: str) -> List[str]: + """ + Get the supported OpenAI parameters for embedding requests. + CometAPI supports standard OpenAI embedding parameters. + """ + return [ + "dimensions", + "encoding_format", + "user", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to CometAPI format. + """ + supported_openai_params = self.get_supported_openai_params(model) + for param, value in non_default_params.items(): + if param in supported_openai_params: + optional_params[param] = value + return optional_params + + def transform_embedding_request( + self, + model: str, + input: AllEmbeddingInputValues, + optional_params: dict, + headers: dict, + ) -> dict: + """ + Transform the embedding request into CometAPI format. + """ + return {"input": input, "model": model, **optional_params} + + def transform_embedding_response( + self, + model: str, + raw_response: httpx.Response, + model_response: EmbeddingResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + request_data: dict, + optional_params: dict, + litellm_params: dict, + ) -> EmbeddingResponse: + """ + Transform CometAPI response into standard EmbeddingResponse format. + """ + try: + raw_response_json = raw_response.json() + except Exception: + raise CometAPIException( + message=raw_response.text, + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + model_response.model = raw_response_json.get("model") + model_response.data = raw_response_json.get("data") + model_response.object = raw_response_json.get("object") + + usage = Usage( + prompt_tokens=raw_response_json.get("usage", {}).get("prompt_tokens", 0), + total_tokens=raw_response_json.get("usage", {}).get("total_tokens", 0), + ) + + model_response.usage = usage + return model_response + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + """ + Get the appropriate error class for CometAPI exceptions. + """ + return CometAPIException( + message=error_message, status_code=status_code, headers=headers + ) diff --git a/litellm/llms/cometapi/image_generation/__init__.py b/litellm/llms/cometapi/image_generation/__init__.py new file mode 100644 index 00000000000..8d7630f2b30 --- /dev/null +++ b/litellm/llms/cometapi/image_generation/__init__.py @@ -0,0 +1,13 @@ +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) + +from .transformation import CometAPIImageGenerationConfig + +__all__ = [ + "CometAPIImageGenerationConfig", +] + + +def get_cometapi_image_generation_config(model: str) -> BaseImageGenerationConfig: + return CometAPIImageGenerationConfig() diff --git a/litellm/llms/cometapi/image_generation/cost_calculator.py b/litellm/llms/cometapi/image_generation/cost_calculator.py new file mode 100644 index 00000000000..b10c9d09087 --- /dev/null +++ b/litellm/llms/cometapi/image_generation/cost_calculator.py @@ -0,0 +1,25 @@ +from typing import Any + +import litellm +from litellm.types.utils import ImageResponse + + +def cost_calculator( + model: str, + image_response: Any, +) -> float: + """ + CometAPI image generation cost calculator + """ + _model_info = litellm.get_model_info( + model=model, + custom_llm_provider=litellm.LlmProviders.COMETAPI.value, + ) + output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0 + num_images: int = 0 + if isinstance(image_response, ImageResponse): + if image_response.data: + num_images = len(image_response.data) + return output_cost_per_image * num_images + else: + raise ValueError(f"image_response must be of type ImageResponse got type={type(image_response)}") diff --git a/litellm/llms/cometapi/image_generation/transformation.py b/litellm/llms/cometapi/image_generation/transformation.py new file mode 100644 index 00000000000..bf1ca9ddde6 --- /dev/null +++ b/litellm/llms/cometapi/image_generation/transformation.py @@ -0,0 +1,170 @@ +from typing import TYPE_CHECKING, Any, List, Optional + +import httpx + +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIImageGenerationOptionalParams, +) +from litellm.types.utils import ImageObject, ImageResponse + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class CometAPIImageGenerationConfig(BaseImageGenerationConfig): + DEFAULT_BASE_URL: str = "https://api.cometapi.com" + IMAGE_GENERATION_ENDPOINT: str = "v1/images/generations" + + def get_supported_openai_params( + self, model: str + ) -> List[OpenAIImageGenerationOptionalParams]: + """ + https://api.cometapi.com/v1/images/generations + """ + return [ + "n", + "quality", + "response_format", + "size", + "style", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + supported_params = self.get_supported_openai_params(model) + + for k in non_default_params.keys(): + if k not in optional_params.keys(): + if k in supported_params: + # CometAPI uses OpenAI-compatible parameters, so we can pass them directly + optional_params[k] = non_default_params[k] + elif drop_params: + pass + else: + raise ValueError( + f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters." + ) + + return optional_params + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + """ + Get the complete url for the request + """ + complete_url: str = ( + api_base + or get_secret_str("COMETAPI_BASE_URL") + or get_secret_str("COMETAPI_API_BASE") + or self.DEFAULT_BASE_URL + ) + + complete_url = complete_url.rstrip("/") + complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}" + return complete_url + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + final_api_key: Optional[str] = ( + api_key or + get_secret_str("COMETAPI_KEY") or + get_secret_str("COMETAPI_API_KEY") + ) + if not final_api_key: + raise ValueError("COMETAPI_KEY or COMETAPI_API_KEY is not set") + + headers["Authorization"] = f"Bearer {final_api_key}" + headers["Content-Type"] = "application/json" + return headers + + def transform_image_generation_request( + self, + model: str, + prompt: str, + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transform the image generation request to the CometAPI image generation request body + + https://api.cometapi.com/v1/images/generations + """ + # CometAPI uses OpenAI-compatible format + request_body = { + "prompt": prompt, + "model": model, + **optional_params, + } + return request_body + + def transform_image_generation_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ImageResponse, + logging_obj: LiteLLMLoggingObj, + request_data: dict, + optional_params: dict, + litellm_params: dict, + encoding: Any, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> ImageResponse: + """ + Transform the image generation response to the litellm image response + + https://api.cometapi.com/v1/images/generations + """ + try: + response_data = raw_response.json() + except Exception as e: + raise self.get_error_class( + error_message=f"Error transforming image generation response: {e}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + if not model_response.data: + model_response.data = [] + + # CometAPI returns OpenAI-compatible format + # Expected format: {"created": timestamp, "data": [{"url": "...", "b64_json": "..."}]} + if "data" in response_data: + for image_data in response_data["data"]: + image_obj = ImageObject( + b64_json=image_data.get("b64_json"), + url=image_data.get("url"), + ) + model_response.data.append(image_obj) + + return model_response diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index a3ad2c67272..accdddbc4dd 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1,6 +1,7 @@ import asyncio import os import ssl +import sys import time from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, Union @@ -114,6 +115,28 @@ def get_ssl_configuration( # but falls back to widely compatible ones custom_ssl_context.set_ciphers(DEFAULT_SSL_CIPHERS) + # Configure ECDH curve for key exchange (e.g., to disable PQC and improve performance) + # Set SSL_ECDH_CURVE env var or litellm.ssl_ecdh_curve to 'X25519' to disable PQC + # Common valid curves: X25519, prime256v1, secp384r1, secp521r1 + ssl_ecdh_curve = os.getenv("SSL_ECDH_CURVE", litellm.ssl_ecdh_curve) + if ssl_ecdh_curve and isinstance(ssl_ecdh_curve, str): + try: + custom_ssl_context.set_ecdh_curve(ssl_ecdh_curve) + verbose_logger.debug(f"SSL ECDH curve set to: {ssl_ecdh_curve}") + except AttributeError: + verbose_logger.warning( + f"SSL ECDH curve configuration not supported. " + f"Python version: {sys.version.split()[0]}, OpenSSL version: {ssl.OPENSSL_VERSION}. " + f"Requested curve: {ssl_ecdh_curve}. Continuing with default curves." + ) + except ValueError as e: + # Invalid curve name + verbose_logger.warning( + f"Invalid SSL ECDH curve name: '{ssl_ecdh_curve}'. {e}. " + f"Common valid curves: X25519, prime256v1, secp384r1, secp521r1. " + f"Continuing with default curves (including PQC)." + ) + # Use our custom SSL context instead of the original ssl_verify value return custom_ssl_context diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 5037f4d8d44..d7b7987b670 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -39,6 +39,7 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) +from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig @@ -1256,6 +1257,289 @@ class BaseLLMHTTPHandler: api_key=api_key, ) + def _prepare_ocr_request( + self, + model: str, + document: Dict[str, str], + optional_params: dict, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + api_base: Optional[str], + headers: Optional[Dict[str, Any]], + provider_config: BaseOCRConfig, + litellm_params: dict, + ) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]: + """ + Shared logic for preparing OCR requests. + Returns: (headers, complete_url, data, files) + """ + from litellm.llms.base_llm.ocr.transformation import OCRRequestData + + headers = provider_config.validate_environment( + api_key=api_key, + api_base=api_base, + headers=headers or {}, + model=model, + ) + + complete_url = provider_config.get_complete_url( + api_base=api_base, + model=model, + optional_params=optional_params, + ) + + # Transform the request to get data and files + transformed_result = provider_config.transform_ocr_request( + model=model, + document=document, + optional_params=optional_params, + headers=headers, + ) + + # All providers return OCRRequestData + if not isinstance(transformed_result, OCRRequestData): + raise ValueError( + f"Provider {provider_config.__class__.__name__} must return OCRRequestData" + ) + + # Data is always a dict for Mistral OCR format + if not isinstance(transformed_result.data, dict): + raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}") + + data = transformed_result.data + + ## LOGGING + logging_obj.pre_call( + input="OCR document processing", + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "api_base": complete_url, + "headers": headers, + }, + ) + + return headers, complete_url, data, None + + async def _async_prepare_ocr_request( + self, + model: str, + document: Dict[str, str], + optional_params: dict, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + api_base: Optional[str], + headers: Optional[Dict[str, Any]], + provider_config: BaseOCRConfig, + litellm_params: dict, + ) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]: + """ + Async version of _prepare_ocr_request for providers that need async transforms. + Returns: (headers, complete_url, data, files) + """ + from litellm.llms.base_llm.ocr.transformation import OCRRequestData + + headers = provider_config.validate_environment( + api_key=api_key, + api_base=api_base, + headers=headers or {}, + model=model, + ) + + complete_url = provider_config.get_complete_url( + api_base=api_base, + model=model, + optional_params=optional_params, + ) + + # Use async transform (providers can override this method if they need async operations) + transformed_result = await provider_config.async_transform_ocr_request( + model=model, + document=document, + optional_params=optional_params, + headers=headers, + ) + + # All providers return OCRRequestData + if not isinstance(transformed_result, OCRRequestData): + raise ValueError( + f"Provider {provider_config.__class__.__name__} must return OCRRequestData" + ) + + # Data is always a dict for Mistral OCR format + if not isinstance(transformed_result.data, dict): + raise ValueError(f"Expected dict data for OCR request, got {type(transformed_result.data)}") + + data = transformed_result.data + + ## LOGGING + logging_obj.pre_call( + input="OCR document processing", + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "api_base": complete_url, + "headers": headers, + }, + ) + + return headers, complete_url, data, None + + def _transform_ocr_response( + self, + provider_config: BaseOCRConfig, + model: str, + response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> OCRResponse: + """Shared logic for transforming OCR responses.""" + return provider_config.transform_ocr_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) + + def ocr( + self, + model: str, + document: Dict[str, str], + optional_params: dict, + timeout: Union[float, httpx.Timeout], + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + api_base: Optional[str], + custom_llm_provider: str, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + aocr: bool = False, + headers: Optional[Dict[str, Any]] = None, + provider_config: Optional[BaseOCRConfig] = None, + litellm_params: Optional[dict] = None, + ) -> Union[OCRResponse, Coroutine[Any, Any, OCRResponse]]: + """ + Sync OCR handler. + """ + if provider_config is None: + raise ValueError( + f"No provider config found for model: {model} and provider: {custom_llm_provider}" + ) + + if litellm_params is None: + litellm_params = {} + + if aocr is True: + return self.async_ocr( + model=model, + document=document, + optional_params=optional_params, + timeout=timeout, + logging_obj=logging_obj, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + client=client, + headers=headers, + provider_config=provider_config, + litellm_params=litellm_params, + ) + + # Prepare the request + headers, complete_url, data, files = self._prepare_ocr_request( + model=model, + document=document, + optional_params=optional_params, + logging_obj=logging_obj, + api_key=api_key, + api_base=api_base, + headers=headers, + provider_config=provider_config, + litellm_params=litellm_params, + ) + + if client is None or not isinstance(client, HTTPHandler): + client = _get_httpx_client() + + try: + # Make the POST request with JSON data (Mistral format) + response = client.post( + url=complete_url, + headers=headers, + json=data, + timeout=timeout, + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=provider_config) + + return self._transform_ocr_response( + provider_config=provider_config, + model=model, + response=response, + logging_obj=logging_obj, + ) + + async def async_ocr( + self, + model: str, + document: Dict[str, str], + optional_params: dict, + timeout: Union[float, httpx.Timeout], + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str], + api_base: Optional[str], + custom_llm_provider: str, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + headers: Optional[Dict[str, Any]] = None, + provider_config: Optional[BaseOCRConfig] = None, + litellm_params: Optional[dict] = None, + ) -> OCRResponse: + """ + Async OCR handler. + """ + if provider_config is None: + raise ValueError( + f"No provider config found for model: {model} and provider: {custom_llm_provider}" + ) + + if litellm_params is None: + litellm_params = {} + + # Prepare the request using async prepare method + headers, complete_url, data, files = await self._async_prepare_ocr_request( + model=model, + document=document, + optional_params=optional_params, + logging_obj=logging_obj, + api_key=api_key, + api_base=api_base, + headers=headers, + provider_config=provider_config, + litellm_params=litellm_params, + ) + + if client is None or not isinstance(client, AsyncHTTPHandler): + async_httpx_client = get_async_httpx_client( + llm_provider=litellm.LlmProviders(custom_llm_provider), + ) + else: + async_httpx_client = client + + try: + # Make the async POST request with JSON data (Mistral format) + response = await async_httpx_client.post( + url=complete_url, + headers=headers, + json=data, + timeout=timeout, + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=provider_config) + + return self._transform_ocr_response( + provider_config=provider_config, + model=model, + response=response, + logging_obj=logging_obj, + ) + async def async_anthropic_messages_handler( self, model: str, @@ -2995,6 +3279,7 @@ class BaseLLMHTTPHandler: BaseGoogleGenAIGenerateContentConfig, BaseAnthropicMessagesConfig, BaseBatchesConfig, + BaseOCRConfig, "BasePassthroughConfig", ], ): diff --git a/litellm/llms/mistral/ocr/__init__.py b/litellm/llms/mistral/ocr/__init__.py new file mode 100644 index 00000000000..40cc62696be --- /dev/null +++ b/litellm/llms/mistral/ocr/__init__.py @@ -0,0 +1,2 @@ +"""Mistral OCR transformation module.""" + diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py new file mode 100644 index 00000000000..f17c872f536 --- /dev/null +++ b/litellm/llms/mistral/ocr/transformation.py @@ -0,0 +1,223 @@ +""" +Mistral OCR transformation implementation. +""" +from typing import Any, Dict, Optional + +import httpx + +from litellm._logging import verbose_logger +from litellm.llms.base_llm.ocr.transformation import ( + BaseOCRConfig, + DocumentType, + OCRRequestData, + OCRResponse, +) +from litellm.secret_managers.main import get_secret_str + + +class MistralOCRConfig(BaseOCRConfig): + """ + Mistral OCR transformation configuration. + + Reference: https://docs.mistral.ai/api/#tag/ocr + """ + + def __init__(self) -> None: + super().__init__() + + def get_supported_ocr_params(self, model: str) -> list: + """ + Get supported OCR parameters for Mistral OCR. + + Mistral OCR supports: + - pages: List of page numbers to process + - include_image_base64: Whether to include base64 encoded images + - image_limit: Maximum number of images to return + - image_min_size: Minimum size of images to include + - bbox_annotation_format: Format for bounding box annotations + - document_annotation_format: Format for document annotations + """ + return [ + "pages", + "include_image_base64", + "image_limit", + "image_min_size", + "bbox_annotation_format", + "document_annotation_format", + ] + + def map_ocr_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + ) -> dict: + """ + Map OCR parameters to Mistral-specific format. + + Mistral accepts these parameters directly, so no transformation needed. + Just filter out unsupported params. + """ + supported_params = self.get_supported_ocr_params(model=model) + + # Only include params that are in the supported list + mapped_params = {} + for param, value in non_default_params.items(): + if param in supported_params: + mapped_params[param] = value + + return mapped_params + + def validate_environment( + self, + headers: Dict, + model: str, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + **kwargs, + ) -> Dict: + """ + Validate environment and return headers for Mistral OCR. + """ + # Get API key from environment if not provided + if api_key is None: + api_key = ( + get_secret_str("MISTRAL_API_KEY") + ) + + if api_key is None: + raise ValueError( + "Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params" + ) + + headers = { + "Authorization": f"Bearer {api_key}", + **headers, + } + + # Don't set Content-Type for multipart/form-data - httpx will handle it + + return headers + + def get_complete_url( + self, + api_base: Optional[str], + model: str, + optional_params: dict, + **kwargs, + ) -> str: + """ + Get complete URL for Mistral OCR endpoint. + + Returns: https://api.mistral.ai/v1/ocr + """ + if api_base is None: + api_base = "https://api.mistral.ai/v1" + + # Ensure no trailing slash + api_base = api_base.rstrip("/") + + # Remove /v1 if it's already in the base to avoid duplication + if api_base.endswith("/v1"): + return f"{api_base}/ocr" + + return f"{api_base}/v1/ocr" + + + def transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: dict, + headers: dict, + **kwargs, + ) -> OCRRequestData: + """ + Transform OCR request to Mistral-specific format. + + Mistral OCR API accepts: + { + "model": "mistral-ocr-latest", + "document": { + "type": "document_url", + "document_url": "" + }, + "pages": [0], # optional + "include_image_base64": false, # optional + ... + } + + Args: + model: Model name (e.g., "mistral-ocr-latest") + document: Document dict from user (Mistral format) - already validated in main.py + optional_params: Already mapped optional parameters + headers: Request headers + + Returns: + OCRRequestData with JSON data + """ + verbose_logger.debug(f"Mistral OCR transform_ocr_request - model: {model}") + + # Document parameter is the Mistral-format dict from the user + # Just pass it through as-is to the Mistral API + if not isinstance(document, dict): + raise ValueError(f"Expected document dict, got {type(document)}") + + # Build request data - use document dict directly + data = { + "model": model, + "document": document, # Pass through the Mistral-format document dict + } + + # Add all optional parameters from the already-mapped optional_params + data.update(optional_params) + + # No multipart files - using JSON + return OCRRequestData(data=data, files=None) + + def transform_ocr_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: Any, + **kwargs, + ) -> OCRResponse: + """ + Return Mistral OCR response in native format. + + Mistral OCR is the standard format for LiteLLM OCR responses. + No transformation needed - return native response. + + Mistral OCR returns: + { + "pages": [ + { + "index": 0, + "markdown": "extracted text content", + "images": [...], + "dimensions": {...} + }, + ... + ], + "model": "mistral-ocr-2505-completion", + "document_annotation": null, + "usage_info": {...} + } + """ + try: + response_json = raw_response.json() + + verbose_logger.debug(f"Mistral OCR response keys: {response_json.keys()}") + + # Return native Mistral format - no transformation + return OCRResponse( + pages=response_json.get("pages", []), + model=response_json.get("model", model), + document_annotation=response_json.get("document_annotation"), + usage_info=response_json.get("usage_info"), + object="ocr", + ) + except Exception as e: + verbose_logger.error(f"Error parsing Mistral OCR response: {e}") + raise e + diff --git a/litellm/llms/openai/image_edit/__init__.py b/litellm/llms/openai/image_edit/__init__.py new file mode 100644 index 00000000000..c1898326b72 --- /dev/null +++ b/litellm/llms/openai/image_edit/__init__.py @@ -0,0 +1,26 @@ +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig + +from .dalle2_transformation import DallE2ImageEditConfig +from .transformation import OpenAIImageEditConfig + +__all__ = ["OpenAIImageEditConfig", "DallE2ImageEditConfig", "get_openai_image_edit_config"] + + +def get_openai_image_edit_config(model: str) -> BaseImageEditConfig: + """ + Get the appropriate OpenAI image edit config based on the model. + + Args: + model: The model name (e.g., "dall-e-2", "gpt-image-1") + + Returns: + The appropriate config instance for the model + """ + model_normalized = model.lower().replace("-", "").replace("_", "") + + if model_normalized == "dalle2": + return DallE2ImageEditConfig() + else: + # Default to standard OpenAI config for gpt-image-1 and other models + return OpenAIImageEditConfig() + diff --git a/litellm/llms/openai/image_edit/dalle2_transformation.py b/litellm/llms/openai/image_edit/dalle2_transformation.py new file mode 100644 index 00000000000..37e92be17a8 --- /dev/null +++ b/litellm/llms/openai/image_edit/dalle2_transformation.py @@ -0,0 +1,101 @@ +from io import BufferedReader +from typing import TYPE_CHECKING, Any, Dict, List, Tuple, cast + +from httpx._types import RequestFiles + +import litellm +from litellm.images.utils import ImageEditRequestUtils +from litellm.types.images.main import ImageEditRequestParams +from litellm.types.llms.openai import FileTypes +from litellm.types.router import GenericLiteLLMParams + +from .transformation import OpenAIImageEditConfig + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class DallE2ImageEditConfig(OpenAIImageEditConfig): + """ + DALL-E-2 specific configuration for image edit API. + + DALL-E-2 only supports editing a single image (not an array). + Uses "image" field name instead of "image[]". + """ + + 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 image edit request for DALL-E-2. + + DALL-E-2 only accepts a single image with field name "image" (not "image[]"). + """ + request = ImageEditRequestParams( + model=model, + image=image, + prompt=prompt, + **image_edit_optional_request_params, + ) + request_dict = cast(Dict, request) + + ######################################################### + # Separate images and masks as `files` and send other parameters as `data` + ######################################################### + _image_list = request_dict.get("image") + _mask = request_dict.get("mask") + data_without_files = { + k: v for k, v in request_dict.items() if k not in ["image", "mask"] + } + files_list: List[Tuple[str, Any]] = [] + + # Handle image parameter - DALL-E-2 only supports single image + if _image_list is not None: + image_list = ( + [_image_list] if not isinstance(_image_list, list) else _image_list + ) + + # Validate only one image is provided + if len(image_list) > 1: + raise litellm.BadRequestError( + message="DALL-E-2 only supports editing a single image. Please provide one image.", + model=model, + llm_provider="openai", + ) + + # Use "image" field name (singular) for DALL-E-2 + for _image in image_list: + if _image is not None: + self._add_image_to_files( + files_list=files_list, + image=_image, + field_name="image", + ) + + # Handle mask parameter if provided + if _mask is not None: + # Handle case where mask can be a list (extract first mask) + if isinstance(_mask, list): + _mask = _mask[0] if _mask else None + + if _mask is not None: + mask_content_type: str = ImageEditRequestUtils.get_image_content_type( + _mask + ) + if isinstance(_mask, BufferedReader): + files_list.append(("mask", (_mask.name, _mask, mask_content_type))) + else: + files_list.append(("mask", ("mask.png", _mask, mask_content_type))) + + return data_without_files, files_list + diff --git a/litellm/llms/openai/image_edit/transformation.py b/litellm/llms/openai/image_edit/transformation.py index be960641154..1b90d96fa92 100644 --- a/litellm/llms/openai/image_edit/transformation.py +++ b/litellm/llms/openai/image_edit/transformation.py @@ -27,6 +27,11 @@ else: class OpenAIImageEditConfig(BaseImageEditConfig): + """ + Base configuration for OpenAI image edit API. + Used for models like gpt-image-1 that support multiple images. + """ + def get_supported_openai_params(self, model: str) -> list: """ All OpenAI Image Edits params are supported @@ -57,6 +62,20 @@ class OpenAIImageEditConfig(BaseImageEditConfig): """No mapping applied since inputs are in OpenAI spec already""" return dict(image_edit_optional_params) + def _add_image_to_files( + self, + files_list: List[Tuple[str, Any]], + image: Any, + field_name: str, + ) -> None: + """Add an image to the files list with appropriate content type""" + image_content_type = ImageEditRequestUtils.get_image_content_type(image) + + if isinstance(image, BufferedReader): + files_list.append((field_name, (image.name, image, image_content_type))) + else: + files_list.append((field_name, ("image.png", image, image_content_type))) + def transform_image_edit_request( self, model: str, @@ -67,9 +86,10 @@ class OpenAIImageEditConfig(BaseImageEditConfig): headers: dict, ) -> Tuple[Dict, RequestFiles]: """ - No transform applied since inputs are in OpenAI spec already + Transform image edit request to OpenAI API format. - This handles buffered readers as images to be sent as multipart/form-data for OpenAI + Handles multipart/form-data for images. Uses "image[]" field name + to support multiple images (e.g., for gpt-image-1). """ request = ImageEditRequestParams( model=model, @@ -94,19 +114,14 @@ class OpenAIImageEditConfig(BaseImageEditConfig): image_list = ( [_image_list] if not isinstance(_image_list, list) else _image_list ) + for _image in image_list: if _image is not None: - image_content_type: str = ( - ImageEditRequestUtils.get_image_content_type(_image) + self._add_image_to_files( + files_list=files_list, + image=_image, + field_name="image[]", ) - if isinstance(_image, BufferedReader): - files_list.append( - ("image[]", (_image.name, _image, image_content_type)) - ) - else: - files_list.append( - ("image[]", ("image.png", _image, image_content_type)) - ) # Handle mask parameter if provided if _mask is not None: # Handle case where mask can be a list (extract first mask) diff --git a/litellm/main.py b/litellm/main.py index 3955a0f32f8..5c50f096460 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4754,6 +4754,33 @@ def embedding( # noqa: PLR0915 aembedding=aembedding, litellm_params={}, ) + elif custom_llm_provider == "cometapi": + api_key = ( + api_key + or litellm.cometapi_key + or get_secret_str("COMETAPI_KEY") + or litellm.api_key + ) + api_base = ( + api_base + or litellm.api_base + or get_secret_str("COMETAPI_API_BASE") + or "https://api.cometapi.com/v1" + ) + response = base_llm_http_handler.embedding( + model=model, + input=input, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + api_key=api_key, + logging_obj=logging, + timeout=timeout, + model_response=EmbeddingResponse(), + optional_params=optional_params, + client=client, + aembedding=aembedding, + litellm_params={}, + ) elif custom_llm_provider in litellm._custom_providers: custom_handler: Optional[CustomLLM] = None for item in litellm.custom_provider_map: diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c965fa2092b..8cab321c4dd 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -400,6 +400,44 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "anthropic.claude-haiku-4-5@20251001": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -810,6 +848,25 @@ "supports_tool_choice": true, "supports_vision": true }, + "apac.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "apac.anthropic.claude-3-sonnet-20240229-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -4612,6 +4669,48 @@ "supports_web_search": true, "tool_use_system_prompt_tokens": 264 }, + "claude-haiku-4-5-20251001": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_computer_use": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "claude-haiku-4-5": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_computer_use": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "claude-3-5-sonnet-20240620": { "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -7741,6 +7840,25 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "eu.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "eu.anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -9486,6 +9604,54 @@ "supports_vision": true, "supports_web_search": true }, + "gemini-2.5-flash-image": { + "cache_read_input_token_cost": 3e-08, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "max_pdf_size_mb": 30, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "image_generation", + "output_cost_per_image": 0.039, + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "rpm": 100000, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_vision": true, + "supports_web_search": true, + "tpm": 8000000 + }, "gemini-2.5-flash-image-preview": { "cache_read_input_token_cost": 7.5e-08, "input_cost_per_audio_token": 1e-06, @@ -10939,6 +11105,54 @@ "supports_web_search": true, "tpm": 8000000 }, + "gemini/gemini-2.5-flash-image": { + "cache_read_input_token_cost": 3e-08, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "max_pdf_size_mb": 30, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "image_generation", + "output_cost_per_image": 0.039, + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "rpm": 100000, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_vision": true, + "supports_web_search": true, + "tpm": 8000000 + }, "gemini/gemini-2.5-flash-image-preview": { "cache_read_input_token_cost": 7.5e-08, "input_cost_per_audio_token": 1e-06, @@ -13650,8 +13864,56 @@ "lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 32768, - "max_input_tokens": 32768, + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/gpt-oss-20b-mxfp4-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 131072, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/gpt-oss-120b-mxfp-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 131072, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/Gemma-3-4b-it-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 128000, + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/Qwen3-4B-Instruct-2507-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 262144, + "max_input_tokens": 262144, "max_output_tokens": 32768, "mode": "chat", "output_cost_per_token": 0, @@ -14466,6 +14728,25 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 346 }, + "jp.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "lambda_ai/deepseek-llama3.3-70b": { "input_cost_per_token": 2e-07, "litellm_provider": "lambda_ai", @@ -17153,6 +17434,8 @@ }, "openrouter/anthropic/claude-opus-4": { "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 1.875e-05, + "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "openrouter", "max_input_tokens": 200000, @@ -17163,6 +17446,7 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true, @@ -17170,6 +17454,9 @@ }, "openrouter/anthropic/claude-opus-4.1": { "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, + "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "openrouter", "max_input_tokens": 200000, @@ -17180,6 +17467,7 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true, @@ -17187,6 +17475,10 @@ }, "openrouter/anthropic/claude-sonnet-4": { "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_above_200k_tokens": 6e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, @@ -17199,6 +17491,31 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 159 + }, + "openrouter/anthropic/claude-sonnet-4.5": { + "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "input_cost_per_token_above_200k_tokens": 6e-06, + "output_cost_per_token_above_200k_tokens": 2.25e-05, + "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true, @@ -20097,6 +20414,25 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "us.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "us.anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -21368,6 +21704,25 @@ "supports_pdf_input": true, "supports_tool_choice": true }, + "vertex_ai/claude-haiku-4-5@20251001": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "vertex_ai/claude-3-5-sonnet": { "input_cost_per_token": 3e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -21560,8 +21915,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, @@ -21577,8 +21932,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, diff --git a/litellm/ocr/__init__.py b/litellm/ocr/__init__.py new file mode 100644 index 00000000000..53f455619d7 --- /dev/null +++ b/litellm/ocr/__init__.py @@ -0,0 +1,5 @@ +"""OCR module for LiteLLM.""" +from .main import aocr, ocr + +__all__ = ["ocr", "aocr"] + diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py new file mode 100644 index 00000000000..62172b0fbae --- /dev/null +++ b/litellm/ocr/main.py @@ -0,0 +1,301 @@ +""" +Main OCR function for LiteLLM. +""" +import asyncio +import contextvars +from functools import partial +from typing import Any, Coroutine, Dict, Optional, Union + +import httpx + +import litellm +from litellm._logging import verbose_logger +from litellm.constants import request_timeout +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.utils import ProviderConfigManager, client + +####### ENVIRONMENT VARIABLES ################### +base_llm_http_handler = BaseLLMHTTPHandler() +################################################# + + +@client +async def aocr( + model: str, + document: Dict[str, str], + api_key: Optional[str] = None, + api_base: Optional[str] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + **kwargs, +) -> OCRResponse: + """ + Async OCR function. + + Args: + model: Model name (e.g., "mistral/mistral-ocr-latest") + document: Document to process in Mistral format: + {"type": "document_url", "document_url": "https://..."} for PDFs/docs or + {"type": "image_url", "image_url": "https://..."} for images + api_key: Optional API key + api_base: Optional API base URL + timeout: Optional timeout + custom_llm_provider: Optional custom LLM provider + extra_headers: Optional extra headers + **kwargs: Additional parameters (e.g., include_image_base64, pages, image_limit) + + Returns: + OCRResponse in Mistral OCR format with pages, model, usage_info, etc. + + Example: + ```python + import litellm + + # OCR with PDF + response = await litellm.aocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + }, + include_image_base64=True + ) + + # OCR with image + response = await litellm.aocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "image_url", + "image_url": "https://example.com/image.png" + } + ) + + # OCR with base64 encoded PDF + response = await litellm.aocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": f"data:application/pdf;base64,{base64_pdf}" + } + ) + ``` + """ + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["aocr"] = True + + # Get custom llm provider + if custom_llm_provider is None: + _, custom_llm_provider, _, _ = litellm.get_llm_provider( + model=model, api_base=api_base + ) + + func = partial( + ocr, + model=model, + document=document, + api_key=api_key, + api_base=api_base, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + **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 + + if response is None: + raise ValueError( + f"Got an unexpected None response from the OCR API: {response}" + ) + + return response + except Exception as e: + raise litellm.exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def ocr( + model: str, + document: Dict[str, str], + api_key: Optional[str] = None, + api_base: Optional[str] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + **kwargs, +) -> Union[OCRResponse, Coroutine[Any, Any, OCRResponse]]: + """ + Synchronous OCR function. + + Args: + model: Model name (e.g., "mistral/mistral-ocr-latest") + document: Document to process in Mistral format: + {"type": "document_url", "document_url": "https://..."} for PDFs/docs or + {"type": "image_url", "image_url": "https://..."} for images + api_key: Optional API key + api_base: Optional API base URL + timeout: Optional timeout + custom_llm_provider: Optional custom LLM provider + extra_headers: Optional extra headers + **kwargs: Additional parameters (e.g., include_image_base64, pages, image_limit) + + Returns: + OCRResponse in Mistral OCR format with pages, model, usage_info, etc. + + Example: + ```python + import litellm + + # OCR with PDF + response = litellm.ocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + }, + include_image_base64=True + ) + + # OCR with image + response = litellm.ocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "image_url", + "image_url": "https://example.com/image.png" + } + ) + + # OCR with base64 encoded PDF + response = litellm.ocr( + model="mistral/mistral-ocr-latest", + document={ + "type": "document_url", + "document_url": f"data:application/pdf;base64,{base64_pdf}" + } + ) + + # Access pages + for page in response.pages: + print(f"Page {page.index}: {page.markdown}") + ``` + """ + local_vars = locals() + try: + litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + _is_async = kwargs.pop("aocr", False) is True + + # Validate document parameter format (Mistral spec) + if not isinstance(document, dict): + raise ValueError(f"document must be a dict with 'type' and URL field, got {type(document)}") + + doc_type = document.get("type") + if doc_type not in ["document_url", "image_url"]: + raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url' or 'image_url'") + + model, custom_llm_provider, dynamic_api_key, dynamic_api_base = ( + litellm.get_llm_provider( + model=model, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + api_key=api_key, + ) + ) + + # Update with dynamic values if available + if dynamic_api_key: + api_key = dynamic_api_key + if dynamic_api_base: + api_base = dynamic_api_base + + # Get provider config + ocr_provider_config: Optional[BaseOCRConfig] = ( + ProviderConfigManager.get_provider_ocr_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), + ) + ) + + if ocr_provider_config is None: + raise ValueError( + f"OCR is not supported for provider: {custom_llm_provider}" + ) + + verbose_logger.debug( + f"OCR call - model: {model}, provider: {custom_llm_provider}" + ) + + # Extract OCR-specific parameters from kwargs + supported_params = ocr_provider_config.get_supported_ocr_params(model=model) + non_default_params = {} + for param in supported_params: + if param in kwargs: + non_default_params[param] = kwargs.pop(param) + + # Map parameters to provider-specific format + optional_params = ocr_provider_config.map_ocr_params( + non_default_params=non_default_params, + optional_params={}, + model=model, + ) + + verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}") + + # Pre Call logging + litellm_logging_obj.update_environment_variables( + model=model, + optional_params=optional_params, + litellm_params={ + "litellm_call_id": litellm_call_id, + "api_base": api_base, + }, + custom_llm_provider=custom_llm_provider, + ) + + # Call the handler - pass document dict directly + response = base_llm_http_handler.ocr( + model=model, + document=document, # Pass the entire document dict + optional_params=optional_params, + timeout=timeout or request_timeout, + logging_obj=litellm_logging_obj, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + aocr=_is_async, + headers=extra_headers, + provider_config=ocr_provider_config, + litellm_params={ + "api_base": api_base, + "api_key": api_key, + }, + ) + + return response + except Exception as e: + raise litellm.exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index b4a76822022..cc57ceac50e 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -54,12 +54,7 @@ async def allm_passthrough_route( cookies: Optional[CookieTypes] = None, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, **kwargs, -) -> Union[ - httpx.Response, - Coroutine[Any, Any, httpx.Response], - Generator[Any, Any, Any], - AsyncGenerator[Any, Any], -]: +) -> Union[httpx.Response, AsyncGenerator[Any, Any]]: """ Async: Reranks a list of documents based on their relevance to the query """ @@ -111,23 +106,25 @@ async def allm_passthrough_route( func_with_context = partial(ctx.run, func) init_response = await loop.run_in_executor(None, func_with_context) + # Since allm_passthrough_route=True, we always get a coroutine from _async_passthrough_request if asyncio.iscoroutine(init_response): response = await init_response - try: + # Only call raise_for_status if it's a Response object (not a generator) + if isinstance(response, httpx.Response): response.raise_for_status() - except httpx.HTTPStatusError as e: - error_text = await e.response.aread() - error_text_str = error_text.decode("utf-8") - raise Exception(error_text_str) - + + return response else: - response = init_response - - return response + # This shouldn't happen when allm_passthrough_route=True, but handle it for type safety + raise Exception("Expected coroutine from async passthrough route") + except httpx.HTTPStatusError as e: + # For HTTP errors, re-raise as-is to preserve the original error details + # The caller (e.g., proxy layer) can handle conversion to appropriate response format + raise e except Exception as e: - # For passthrough routes, we need to get the provider config to properly handle errors + # For other exceptions, use provider-specific error handling from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager @@ -186,6 +183,7 @@ def llm_passthrough_route( ) -> Union[ httpx.Response, Coroutine[Any, Any, httpx.Response], + Coroutine[Any, Any, Union[httpx.Response, AsyncGenerator[Any, Any]]], Generator[Any, Any, Any], AsyncGenerator[Any, Any], ]: @@ -200,8 +198,10 @@ def llm_passthrough_route( from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager + _is_async = allm_passthrough_route + if client is None: - if allm_passthrough_route: + if _is_async: client = litellm.module_level_aclient else: client = litellm.module_level_client @@ -302,24 +302,40 @@ def llm_passthrough_route( # Update logging object with streaming status litellm_logging_obj.stream = is_streaming_request + ## LOGGING PRE-CALL + request_data = data if data else json + litellm_logging_obj.pre_call( + input=request_data, + api_key=provider_api_key, + additional_args={ + "complete_input_dict": request_data, + "api_base": str(updated_url), + "headers": headers, + }, + ) + try: - response = client.client.send(request=request, stream=is_streaming_request) - if asyncio.iscoroutine(response): - if is_streaming_request: - return _async_streaming(response, litellm_logging_obj, provider_config) - else: - return response - response.raise_for_status() - - if ( - hasattr(response, "iter_bytes") and is_streaming_request - ): # yield the chunk, so we can store it in the logging object - - return _sync_streaming(response, litellm_logging_obj, provider_config) + if _is_async: + # Return the coroutine to be awaited by the caller + return _async_passthrough_request( + client=client, + request=request, + is_streaming_request=is_streaming_request, + litellm_logging_obj=litellm_logging_obj, + provider_config=provider_config, + ) else: + # Sync path - client.client.send returns Response directly + response: httpx.Response = client.client.send(request=request, stream=is_streaming_request) # type: ignore + response.raise_for_status() - # For non-streaming responses, yield the entire response - return response + if ( + hasattr(response, "iter_bytes") and is_streaming_request + ): # yield the chunk, so we can store it in the logging object + return _sync_streaming(response, litellm_logging_obj, provider_config) + else: + # For non-streaming responses, yield the entire response + return response except Exception as e: if provider_config is None: raise e @@ -329,6 +345,39 @@ def llm_passthrough_route( ) +async def _async_passthrough_request( + client: Union[HTTPHandler, AsyncHTTPHandler], + request: httpx.Request, + is_streaming_request: bool, + litellm_logging_obj: "LiteLLMLoggingObj", + provider_config: "BasePassthroughConfig", +) -> Union[httpx.Response, AsyncGenerator[Any, Any]]: + """ + Handle async passthrough requests. + Uses async client to send request and properly handles streaming. + """ + # client.client.send returns a coroutine for async clients + response_result = client.client.send(request=request, stream=is_streaming_request) + + # Check if it's a coroutine and await it + if asyncio.iscoroutine(response_result): + if is_streaming_request: + # Pass the coroutine to _async_streaming which will await it + return _async_streaming( + response=response_result, + litellm_logging_obj=litellm_logging_obj, + provider_config=provider_config, + ) + else: + response = await response_result + await response.aread() + response.raise_for_status() + return response + else: + # Fallback for sync-like behavior (shouldn't happen in async path) + raise Exception("Expected coroutine from async client") + + def _sync_streaming( response: httpx.Response, litellm_logging_obj: "LiteLLMLoggingObj", diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c606ec048c2..4de258483b0 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -308,6 +308,7 @@ class ProxyBaseLLMRequestProcessing: "allm_passthrough_route", "avector_store_search", "avector_store_create", + "aocr", ], version: Optional[str] = None, user_model: Optional[str] = None, @@ -398,6 +399,7 @@ class ProxyBaseLLMRequestProcessing: "allm_passthrough_route", "avector_store_search", "avector_store_create", + "aocr", ], proxy_logging_obj: ProxyLogging, general_settings: dict, diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 33d432e8695..807b895bd03 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -1,6 +1,6 @@ import json import re -from typing import Any, Dict, List, Optional +from typing import Any, Collection, Dict, List, Optional import orjson from fastapi import Request, UploadFile, status @@ -149,7 +149,7 @@ def _safe_get_request_headers(request: Optional[Request]) -> dict: def check_file_size_under_limit( request_data: dict, file: UploadFile, - router_model_names: List[str], + router_model_names: Collection[str], ) -> bool: """ Check if any files passed in request are under max_file_size_mb diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py index 556c22b9495..29ede085ed6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/__init__.py @@ -23,6 +23,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" if not guardrail_name: raise ValueError("Pillar guardrail name is required") + optional_params = getattr(litellm_params, "optional_params", None) + _pillar_callback = PillarGuardrail( guardrail_name=guardrail_name, api_key=litellm_params.api_key, @@ -30,12 +32,34 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" on_flagged_action=getattr(litellm_params, "on_flagged_action", "monitor"), event_hook=litellm_params.mode, default_on=litellm_params.default_on, + async_mode=_get_config_value( + litellm_params, optional_params, "async_mode" + ), + persist_session=_get_config_value( + litellm_params, optional_params, "persist_session" + ), + include_scanners=_get_config_value( + litellm_params, optional_params, "include_scanners" + ), + include_evidence=_get_config_value( + litellm_params, optional_params, "include_evidence" + ), ) litellm.logging_callback_manager.add_litellm_callback(_pillar_callback) return _pillar_callback +def _get_config_value(litellm_params, optional_params, attribute_name): + """Return guardrail configuration value prioritising optional params when present.""" + + if optional_params is not None: + value = getattr(optional_params, attribute_name, None) + if value is not None: + return value + return getattr(litellm_params, attribute_name, None) + + guardrail_initializer_registry = { SupportedGuardrailIntegrations.PILLAR.value: initialize_guardrail, } diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py index f4741aa8e00..e19125ed031 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py @@ -69,6 +69,10 @@ class PillarGuardrail(CustomGuardrail): api_key: Optional[str] = None, api_base: Optional[str] = None, on_flagged_action: Optional[str] = None, + async_mode: Optional[bool] = None, + persist_session: Optional[bool] = None, + include_scanners: Optional[bool] = None, + include_evidence: Optional[bool] = None, **kwargs, ) -> None: """ @@ -110,6 +114,31 @@ class PillarGuardrail(CustomGuardrail): f"Pillar Guardrail: Initialized with on_flagged_action: {self.on_flagged_action}" ) + self.async_mode = self._resolve_bool_config( + provided_value=async_mode, + env_var="PILLAR_ASYNC", + default=None, + setting_name="async_mode", + ) + self.persist_session = self._resolve_bool_config( + provided_value=persist_session, + env_var="PILLAR_PERSIST", + default=None, + setting_name="persist_session", + ) + self.include_scanners = self._resolve_bool_config( + provided_value=include_scanners, + env_var="PILLAR_INCLUDE_SCANNERS", + default=True, + setting_name="include_scanners", + ) + self.include_evidence = self._resolve_bool_config( + provided_value=include_evidence, + env_var="PILLAR_INCLUDE_EVIDENCE", + default=True, + setting_name="include_evidence", + ) + # Define supported event hooks supported_event_hooks = [ GuardrailEventHooks.pre_call, @@ -347,12 +376,74 @@ class PillarGuardrail(CustomGuardrail): "Content-Type": "application/json", } - # Add Pillar-specific headers for enhanced response data - headers["plr_evidence"] = "true" - headers["plr_scanners"] = "true" + # Add Pillar-specific headers based on configuration + self._set_bool_header(headers, "plr_scanners", self.include_scanners) + self._set_bool_header(headers, "plr_evidence", self.include_evidence) + self._set_bool_header(headers, "plr_async", self.async_mode) + self._set_bool_header(headers, "plr_persist", self.persist_session) return headers + def _set_bool_header( + self, headers: Dict[str, str], header_name: str, value: Optional[bool] + ) -> None: + """Apply a boolean value as a lowercase string HTTP header when provided.""" + + if value is None: + return + headers[header_name] = "true" if value else "false" + + def _resolve_bool_config( + self, + provided_value: Optional[Union[bool, str, int]], + env_var: Optional[str], + default: Optional[bool], + setting_name: str, + ) -> Optional[bool]: + """Resolve configuration precedence: explicit value -> environment -> default.""" + + if provided_value is not None: + try: + return self._parse_bool_value(provided_value) + except ValueError: + verbose_proxy_logger.warning( + "Pillar Guardrail: Invalid boolean value '%s' for %s, falling back to default.", + provided_value, + setting_name, + ) + return default + + if env_var: + env_value = os.getenv(env_var) + if env_value is not None: + try: + return self._parse_bool_value(env_value) + except ValueError: + verbose_proxy_logger.warning( + "Pillar Guardrail: Invalid boolean env value '%s' for %s, falling back to default.", + env_value, + env_var, + ) + return default + + return default + + @staticmethod + def _parse_bool_value(value: Union[bool, str, int]) -> bool: + """Normalise various truthy/falsey inputs to a strict boolean.""" + + if isinstance(value, bool): + return value + if isinstance(value, int): + return bool(value) + + value_str = str(value).strip().lower() + if value_str in {"true", "1", "yes", "y", "on"}: + return True + if value_str in {"false", "0", "no", "n", "off"}: + return False + raise ValueError(f"Unrecognised boolean value: {value}") + def _extract_model_and_provider(self, data: dict) -> Tuple[str, str]: """ Extract the model and provider from the request data. diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index d8f426e1ae6..dec399c7f74 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -9,7 +9,10 @@ Has all /sso/* routes """ import asyncio +import base64 +import hashlib import os +import secrets from copy import deepcopy from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast @@ -381,7 +384,10 @@ async def get_generic_sso_response( try: result = await generic_sso.verify_and_process( request, - params={"include_client_id": generic_include_client_id}, + params=SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=request, + generic_include_client_id=generic_include_client_id, + ), headers=additional_generic_sso_headers_dict, ) @@ -1067,30 +1073,97 @@ class SSOAuthenticationHandler: allow_insecure_http=True, scope=generic_scope, ) - with generic_sso: - # TODO: state should be a random string and added to the user session with cookie - # or a cryptographicly signed state that we can verify stateless - # For simplification we are using a static state, this is not perfect but some - # SSO providers do not allow stateless verification - redirect_params = ( - SSOAuthenticationHandler._get_generic_sso_redirect_params( - state=state, - generic_authorization_endpoint=generic_authorization_endpoint, - ) - ) - - return await generic_sso.get_login_redirect(**redirect_params) # type: ignore + return await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=generic_sso, + state=state, + generic_authorization_endpoint=generic_authorization_endpoint, + ) raise ValueError( "Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso" ) + + @staticmethod + async def get_generic_sso_redirect_response( + generic_sso: Any, + state: Optional[str] = None, + generic_authorization_endpoint: Optional[str] = None, + ) -> Optional[RedirectResponse]: + """ + Get the redirect response for Generic SSO + """ + from urllib.parse import parse_qs, urlencode, urlparse, urlunparse + + from litellm.proxy.proxy_server import user_api_key_cache + with generic_sso: + # TODO: state should be a random string and added to the user session with cookie + # or a cryptographicly signed state that we can verify stateless + # For simplification we are using a static state, this is not perfect but some + # SSO providers do not allow stateless verification + redirect_params, code_verifier = ( + SSOAuthenticationHandler._get_generic_sso_redirect_params( + state=state, + generic_authorization_endpoint=generic_authorization_endpoint, + ) + ) + + # Separate PKCE params from state params (fastapi-sso doesn't accept code_challenge) + pkce_params = {} + state_only_params = {} + for key, value in redirect_params.items(): + if key in ("code_challenge", "code_challenge_method"): + pkce_params[key] = value + else: + state_only_params[key] = value + + # Get the redirect response from fastapi-sso with only state param + redirect_response = await generic_sso.get_login_redirect(**state_only_params) # type: ignore + + # If PKCE is enabled, add PKCE parameters to the redirect URL + if code_verifier and "state" in redirect_params: + + # Store code_verifier in cache (10 min TTL) + cache_key = f"pkce_verifier:{redirect_params['state']}" + user_api_key_cache.set_cache( + key=cache_key, + value=code_verifier, + ttl=600, + ) + + # Add PKCE parameters to the authorization URL + if pkce_params: + parsed_url = urlparse(str(redirect_response.headers["location"])) + query_params = parse_qs(parsed_url.query) + + # Add PKCE parameters + for key, value in pkce_params.items(): + query_params[key] = [value] + + # Reconstruct the URL with PKCE parameters + new_query = urlencode(query_params, doseq=True) + new_url = urlunparse(( + parsed_url.scheme, + parsed_url.netloc, + parsed_url.path, + parsed_url.params, + new_query, + parsed_url.fragment + )) + + # Update the redirect response + redirect_response.headers["location"] = new_url + verbose_proxy_logger.debug( + "PKCE parameters added to authorization URL" + ) + return redirect_response @staticmethod def _get_generic_sso_redirect_params( state: Optional[str] = None, generic_authorization_endpoint: Optional[str] = None, - ) -> dict: + ) -> Tuple[dict, Optional[str]]: """ Get redirect parameters for Generic SSO with proper state priority handling. + Optionally generates PKCE parameters if GENERIC_CLIENT_USE_PKCE is enabled. Priority order: 1. CLI state (if provided) @@ -1102,9 +1175,12 @@ class SSOAuthenticationHandler: generic_authorization_endpoint: Authorization endpoint URL Returns: - dict: Redirect parameters for SSO login + Tuple[dict, Optional[str]]: + - Redirect parameters for SSO login (may include PKCE params) + - code_verifier (if PKCE is enabled, None otherwise) """ redirect_params = {} + code_verifier: Optional[str] = None if state: # CLI state takes priority @@ -1122,7 +1198,18 @@ class SSOAuthenticationHandler: uuid.uuid4().hex ) # set state param for okta - required - return redirect_params + # Handle PKCE (Proof Key for Code Exchange) if enabled + # Set GENERIC_CLIENT_USE_PKCE=true to enable PKCE for enhanced OAuth security + use_pkce = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" + if use_pkce: + code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params() + redirect_params["code_challenge"] = code_challenge + redirect_params["code_challenge_method"] = "S256" + verbose_proxy_logger.debug( + "PKCE enabled - code_challenge added to authorization request" + ) + + return redirect_params, code_verifier @staticmethod def should_use_sso_handler( @@ -1606,6 +1693,69 @@ class SSOAuthenticationHandler: redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303) redirect_response.set_cookie(key="token", value=jwt_token) return redirect_response + + + @staticmethod + def prepare_token_exchange_parameters( + request: Request, + generic_include_client_id: bool, + ) -> dict: + """ + Prepare token exchange parameters for Generic SSO. + + Args: + request: Request object + generic_include_client_id: Generic OAuth Client ID + + Returns: + dict: Token exchange parameters + """ + # Prepare token exchange parameters + token_params = {"include_client_id": generic_include_client_id} + + # Retrieve PKCE code_verifier if PKCE was used in authorization + query_params = dict(request.query_params) + state = query_params.get("state") + if state: + from litellm.proxy.proxy_server import user_api_key_cache + + cache_key = f"pkce_verifier:{state}" + code_verifier = user_api_key_cache.get_cache(key=cache_key) + + if code_verifier: + # Add code_verifier to token exchange parameters + token_params["code_verifier"] = code_verifier + verbose_proxy_logger.debug( + "PKCE code_verifier retrieved and will be included in token exchange" + ) + + # Clean up the cache entry (single-use verifier) + user_api_key_cache.delete_cache(key=cache_key) + return token_params + + + @staticmethod + def generate_pkce_params() -> Tuple[str, str]: + """ + Generate PKCE (Proof Key for Code Exchange) parameters for OAuth 2.0. + + Returns: + Tuple[str, str]: (code_verifier, code_challenge) + - code_verifier: Random 43-128 character string (we use 43 for efficiency) + - code_challenge: Base64-URL-encoded SHA256 hash of the code_verifier + + Reference: https://datatracker.ietf.org/doc/html/rfc7636 + """ + # Generate a cryptographically random code_verifier (43 characters) + # Using 32 random bytes which becomes 43 characters when base64-url-encoded + code_verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).decode('utf-8').rstrip('=') + + # Generate code_challenge using S256 method (SHA256) + code_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest() + code_challenge = base64.urlsafe_b64encode(code_challenge_bytes).decode('utf-8').rstrip('=') + + return code_verifier, code_challenge + class MicrosoftSSOHandler: @@ -1739,7 +1889,7 @@ class MicrosoftSSOHandler: Extract app roles from the Microsoft Entra ID (Azure AD) id_token JWT. App roles are assigned in the Azure AD Enterprise Application and appear - in the 'roles' claim of the id_token. + in the 'app_roles' claim of the id_token. Args: id_token (Optional[str]): The JWT id_token from Microsoft SSO @@ -1758,8 +1908,9 @@ class MicrosoftSSOHandler: # (signature is already verified by fastapi_sso) decoded_token = jwt.decode(id_token, options={"verify_signature": False}) - # Extract roles claim from the token - roles = decoded_token.get("roles", []) + # Extract app_roles claim from the token + ## check for both 'roles' and 'app_roles' claims + roles = decoded_token.get("app_roles", []) or decoded_token.get("roles", []) if roles and isinstance(roles, list): verbose_proxy_logger.debug( diff --git a/litellm/proxy/ocr_endpoints/__init__.py b/litellm/proxy/ocr_endpoints/__init__.py new file mode 100644 index 00000000000..3488912f661 --- /dev/null +++ b/litellm/proxy/ocr_endpoints/__init__.py @@ -0,0 +1,2 @@ +# OCR Endpoints + diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py new file mode 100644 index 00000000000..c1092a06b48 --- /dev/null +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -0,0 +1,97 @@ +#### OCR Endpoints ##### + +import orjson +from fastapi import APIRouter, Depends, Request, Response +from fastapi.responses import ORJSONResponse + +from litellm.proxy._types import * +from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + +router = APIRouter() + + +@router.post( + "/v1/ocr", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, + tags=["ocr"], +) +@router.post( + "/ocr", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, + tags=["ocr"], +) +async def ocr( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + OCR endpoint for extracting text from documents and images. + + Follows the Mistral OCR API spec: + https://docs.mistral.ai/capabilities/vision/#optical-character-recognition-ocr + + Example: + ```bash + curl -X POST "http://localhost:4000/v1/ocr" \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "mistral/mistral-ocr-latest", + "document": { + "type": "document_url", + "document_url": "https://arxiv.org/pdf/2201.04234" + } + }' + ``` + """ + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + # Read request body + body = await request.body() + data = orjson.loads(body) + + # Process request using ProxyBaseLLMRequestProcessing + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="aocr", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=None, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) + diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index d07bfbb11ae..6f9f04e5cc2 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -8,7 +8,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc. import json import os -from typing import Optional, cast +from typing import Any, Optional, Union, cast import httpx from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket @@ -482,6 +482,172 @@ async def anthropic_proxy_route( return received_value +# Bedrock endpoint actions - consolidated list used for model extraction and streaming detection +BEDROCK_ENDPOINT_ACTIONS = { + "invoke", + "invoke-with-response-stream", + "converse", + "converse-stream", + "count_tokens", + "count-tokens", +} + +BEDROCK_STREAMING_ACTIONS = {"invoke-with-response-stream", "converse-stream"} + + +def _extract_model_from_bedrock_endpoint(endpoint: str) -> str: + """ + Extract model name from Bedrock endpoint path. + + Handles model names with slashes (e.g., aws/anthropic/bedrock-claude-3-5-sonnet-v1) + by finding the action in the endpoint and extracting everything between "model" and the action. + + Args: + endpoint: The endpoint path (e.g., "/model/aws/anthropic/model-name/invoke") + + Returns: + The extracted model name (e.g., "aws/anthropic/model-name") + + Raises: + ValueError: If model cannot be extracted from endpoint + """ + try: + endpoint_parts = endpoint.split("/") + + if "application-inference-profile" in endpoint: + # Format: model/application-inference-profile/{profile-id}/{action} + return "/".join(endpoint_parts[1:3]) + + # Format: model/{modelId}/{action} + # Find the index of the action in the endpoint parts + action_index = None + for idx, part in enumerate(endpoint_parts): + if part in BEDROCK_ENDPOINT_ACTIONS: + action_index = idx + break + + if action_index is not None and action_index > 1: + # Join all parts between "model" and the action + return "/".join(endpoint_parts[1:action_index]) + + # Fallback to taking everything after "model" if no action found + return "/".join(endpoint_parts[1:]) + + except Exception as e: + raise ValueError( + f"Model missing from endpoint. Expected format: /model/{{modelId}}/{{action}}. Got: {endpoint}" + ) from e + + +async def handle_bedrock_passthrough_router_model( + model: str, + endpoint: str, + request: Request, + request_body: dict, + llm_router: litellm.Router, +) -> Union[Response, StreamingResponse]: + """ + Handle Bedrock passthrough for router models (models defined in config.yaml). + + This helper delegates to llm_router.allm_passthrough_route for proper credential + and configuration management from the router. + + Args: + model: The router model name (e.g., "aws/anthropic/bedrock-claude-3-5-sonnet-v1") + endpoint: The Bedrock endpoint path (e.g., "/model/{modelId}/invoke") + request: The FastAPI request object + request_body: The parsed request body + llm_router: The LiteLLM router instance + + Returns: + Response or StreamingResponse depending on endpoint type + """ + # Detect streaming based on endpoint + is_streaming = any(action in endpoint for action in BEDROCK_STREAMING_ACTIONS) + + verbose_proxy_logger.debug( + f"Bedrock router passthrough: model='{model}', endpoint='{endpoint}', streaming={is_streaming}" + ) + + # Call router passthrough + try: + result = await llm_router.allm_passthrough_route( + model=model, + method=request.method, + endpoint=endpoint, + request_query_params=request.query_params, + request_headers=dict(request.headers), + stream=is_streaming, + content=None, + data=None, + files=None, + json=( + request_body + if request.headers.get("content-type") == "application/json" + else None + ), + params=None, + headers=None, + cookies=None, + ) + except httpx.HTTPStatusError as e: + # Handle HTTP errors from the provider by converting to HTTPException + error_body = await e.response.aread() + error_text = error_body.decode("utf-8") + + raise HTTPException( + status_code=e.response.status_code, + detail={"error": error_text}, + ) + except Exception as e: + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + # If it's a BaseLLMException (from non-HTTP errors), convert to HTTPException + if isinstance(e, BaseLLMException): + raise HTTPException( + status_code=e.status_code, + detail={"error": e.message}, + ) + # Re-raise any other exceptions + raise e + + # Handle streaming response + if is_streaming: + import inspect + + if inspect.isasyncgen(result): + # AsyncGenerator case + return StreamingResponse( + content=result, + status_code=200, + headers={"content-type": "application/vnd.amazon.eventstream"}, + ) + else: + # httpx.Response case + result = cast(httpx.Response, result) + return StreamingResponse( + content=result.aiter_bytes(), + status_code=result.status_code, + headers=HttpPassThroughEndpointHelpers.get_response_headers( + headers=result.headers, + custom_headers=None, + ), + ) + + # Handle non-streaming response + result = cast(httpx.Response, result) + content = await result.aread() + + return Response( + content=content, + status_code=result.status_code, + headers=HttpPassThroughEndpointHelpers.get_response_headers( + headers=result.headers, + custom_headers=None, + ), + ) + + async def handle_bedrock_count_tokens( endpoint: str, request: Request, @@ -560,6 +726,15 @@ async def bedrock_llm_proxy_route( ): """ Handles Bedrock LLM API calls. + + Supports both direct Bedrock models and router models from config.yaml. + + Endpoints: + - /model/{modelId}/invoke + - /model/{modelId}/invoke-with-response-stream + - /model/{modelId}/converse + - /model/{modelId}/converse-stream + - /model/application-inference-profile/{profileId}/{action} """ from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.proxy_server import ( @@ -588,24 +763,38 @@ async def bedrock_llm_proxy_route( request_body=request_body, ) - data: Dict[str, Any] = {} - base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) + # Extract model from endpoint path using helper try: - endpoint_parts = endpoint.split("/") - if "application-inference-profile" in endpoint: - # For application-inference-profile, include the profile ID part as well - model = "/".join(endpoint_parts[1:3]) - else: - model = endpoint_parts[1] - except Exception: + model = _extract_model_from_bedrock_endpoint(endpoint=endpoint) + except ValueError as e: raise HTTPException( status_code=400, - detail={ - "error": "Model missing from endpoint. Expected format: /model//. Got: " - + endpoint, - }, + detail={"error": str(e)}, ) + # Check if this is a router model (from config.yaml) + is_router_model = is_passthrough_request_using_router_model( + request_body={"model": model}, llm_router=llm_router + ) + + # If router model, use dedicated router passthrough handler + if is_router_model and llm_router: + return await handle_bedrock_passthrough_router_model( + model=model, + endpoint=endpoint, + request=request, + request_body=request_body, + llm_router=llm_router, + ) + + # Fall back to existing implementation for direct Bedrock models + verbose_proxy_logger.debug( + f"Bedrock passthrough: Using direct Bedrock model '{model}' for endpoint '{endpoint}'" + ) + + data: Dict[str, Any] = {} + base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) + data["method"] = request.method data["endpoint"] = endpoint data["data"] = request_body diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 338fa98118c..cca136811f7 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -1,16 +1,25 @@ model_list: - - model_name: db-openai-endpoint + - model_name: mistral/* litellm_params: - model: openai/gm - api_key: hi - api_base: https://exampleopenaiendpoint-production.up.railway.app/ - - -litellm_settings: - callbacks: ["dynamic_rate_limiter_v3"] - priority_reservation: - "prod": 0.9 # 90% reserved for production (9 RPM) - "dev": 0.1 # 10% reserved for development (1 RPM) - priority_reservation_settings: - default_priority: 0.2 # Weight (0%) assigned to keys without explicit priority metadata - saturation_threshold: 0.50 # A model is saturated if it has hit 50% of its RPM limit \ No newline at end of file + model: mistral/* + - model_name: special-bedrock-model + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock + - model_name: aws/anthropic/bedrock-claude-3-5-sonnet-v1 + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock + # Load balancing test - multiple deployments with same model_name + - model_name: load-balanced-claude + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-west-2 + custom_llm_provider: bedrock + - model_name: load-balanced-claude + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + aws_region_name: us-east-1 + custom_llm_provider: bedrock \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1106e0ed12f..5c830a9a1a1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -308,6 +308,7 @@ from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import ( ) from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware +from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, ) @@ -9774,6 +9775,7 @@ app.include_router(response_router) app.include_router(batches_router) app.include_router(public_endpoints_router) app.include_router(rerank_router) +app.include_router(ocr_router) app.include_router(image_router) app.include_router(fine_tuning_router) app.include_router(vector_store_router) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 1ae87637be0..f7b4cf0cbe1 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -25,6 +25,7 @@ ROUTE_ENDPOINT_MAPPING = { "alist_input_items": "/responses/{response_id}/input_items", "aimage_edit": "/images/edits", "acancel_responses": "/responses/{response_id}/cancel", + "aocr": "/ocr", } @@ -98,6 +99,7 @@ async def route_request( "allm_passthrough_route", "avector_store_search", "avector_store_create", + "aocr", ], ): """ diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8aa2f407555..803c9c1a4ff 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3707,6 +3707,7 @@ def construct_database_url_from_env_vars() -> Optional[str]: database_username = os.getenv("DATABASE_USERNAME") database_password = os.getenv("DATABASE_PASSWORD") database_name = os.getenv("DATABASE_NAME") + database_schema = os.getenv("DATABASE_SCHEMA") if database_host and database_username and database_name: # Handle the problem of special character escaping in the database URL @@ -3722,6 +3723,9 @@ def construct_database_url_from_env_vars() -> Optional[str]: else: database_url = f"postgresql://{database_username_enc}@{database_host}/{database_name_enc}" + if database_schema: + database_url += f"?schema={database_schema}" + return database_url return None diff --git a/litellm/router.py b/litellm/router.py index 5972b06f01a..842f6522660 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -189,7 +189,7 @@ class RoutingArgs(enum.Enum): class Router: - model_names: List = [] + model_names: set = set() cache_responses: Optional[bool] = False default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour tenacity = None @@ -872,6 +872,14 @@ class Router: generate_content_stream, call_type="generate_content_stream" ) + ######################################################### + # OCR routes + ######################################################### + from litellm.ocr import aocr, ocr + + self.aocr = self.factory_function(aocr, call_type="aocr") + self.ocr = self.factory_function(ocr, call_type="ocr") + def validate_fallbacks(self, fallback_param: Optional[List]): """ Validate the fallbacks parameter. @@ -1057,7 +1065,7 @@ class Router: self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) request_priority = kwargs.get("priority") or self.default_priority - start_time = time.time() + start_time = time.perf_counter() _is_prompt_management_model = self._is_prompt_management_model(model) if _is_prompt_management_model: @@ -1070,7 +1078,7 @@ class Router: response = await self.schedule_acompletion(**kwargs) else: response = await self.async_function_with_fallbacks(**kwargs) - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( @@ -1245,7 +1253,7 @@ class Router: input_kwargs_for_streaming_fallback["model"] = model parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) - start_time = time.time() + start_time = time.perf_counter() deployment = await self.async_get_available_deployment( model=model, messages=messages, @@ -1254,7 +1262,7 @@ class Router: ) _timeout_debug_deployment_dict = deployment - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( @@ -1834,8 +1842,8 @@ class Router: await self.scheduler.add_request(request=item) ## POLL QUEUE - end_time = time.time() + self.timeout - curr_time = time.time() + end_time = time.monotonic() + self.timeout + curr_time = time.monotonic() poll_interval = self.scheduler.polling_interval # poll every 3ms make_request = False @@ -1852,7 +1860,7 @@ class Router: break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) - curr_time = time.time() + curr_time = time.monotonic() if make_request: try: @@ -1896,8 +1904,8 @@ class Router: await self.scheduler.add_request(request=item) ## POLL QUEUE - end_time = time.time() + self.timeout - curr_time = time.time() + end_time = time.monotonic() + self.timeout + curr_time = time.monotonic() poll_interval = self.scheduler.polling_interval # poll every 3ms make_request = False @@ -1914,7 +1922,7 @@ class Router: break else: ## ELSE -> loop till default_timeout await asyncio.sleep(poll_interval) - curr_time = time.time() + curr_time = time.monotonic() if make_request: try: @@ -2732,6 +2740,37 @@ class Router: ) ) raise e + + def _add_deployment_model_to_endpoint_for_llm_passthrough_route( + self, kwargs: Dict[str, Any], + model: str, + model_name: str + ) -> Dict[str, Any]: + """ + Add the deployment model to the endpoint for LLM passthrough route. + + e.g for bedrock invoke users can pass endpoint as /model/special-bedrock-model/invoke + it should be actually sent as /model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke + """ + if "endpoint" in kwargs and kwargs["endpoint"]: + # For provider-specific endpoints, strip the provider prefix from model_name + # e.g., "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0" -> "us.anthropic.claude-3-5-sonnet-20240620-v1:0" + from litellm import get_llm_provider + + try: + # get_llm_provider returns (model_without_prefix, provider, api_key, api_base) + stripped_model_name, _, _, _ = get_llm_provider( + model=model_name, + custom_llm_provider=kwargs.get("custom_llm_provider"), + api_base=kwargs.get("api_base"), + ) + replacement_model_name = stripped_model_name + except Exception: + # If get_llm_provider fails, fall back to using model_name as-is + replacement_model_name = model_name + + kwargs["endpoint"] = kwargs["endpoint"].replace(model, replacement_model_name) + return kwargs async def _ageneric_api_call_with_fallbacks_helper( self, model: str, original_generic_function: Callable, **kwargs @@ -2764,6 +2803,7 @@ class Router: model_name = data["model"] self.total_calls[model_name] += 1 + self._add_deployment_model_to_endpoint_for_llm_passthrough_route(kwargs=kwargs, model=model, model_name=model_name) ### get custom response = original_generic_function( **{ @@ -2842,6 +2882,12 @@ class Router: self.total_calls[model_name] += 1 + # For passthrough routes, use the actual model from deployment + # and swap model name in endpoint if present + if "endpoint" in kwargs and kwargs["endpoint"]: + kwargs["endpoint"] = kwargs["endpoint"].replace(model, model_name) + kwargs["model"] = model_name + # Perform pre-call checks for routing strategy self.routing_strategy_pre_call_checks(deployment=deployment) @@ -3537,6 +3583,9 @@ class Router: "avector_store_create", "vector_store_search", "vector_store_create", + "aocr", + "ocr", + "aadapter_generate_content" ] = "assistants", ): """ @@ -3553,6 +3602,7 @@ class Router: "generate_content_stream", "vector_store_search", "vector_store_create", + "ocr", ): def sync_wrapper( @@ -3595,6 +3645,8 @@ class Router: "aimage_edit", "agenerate_content", "agenerate_content_stream", + "aocr", + "ocr", ): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, @@ -4915,22 +4967,25 @@ class Router: - hash - use hash as id """ - concat_str = model_group + # Optimized: Use list and join instead of string concatenation in loop + # This avoids creating many temporary string objects (O(n) vs O(n²) complexity) + parts = [model_group] for k, v in litellm_params.items(): if isinstance(k, str): - concat_str += k + parts.append(k) elif isinstance(k, dict): - concat_str += json.dumps(k) + parts.append(json.dumps(k)) else: - concat_str += str(k) + parts.append(str(k)) if isinstance(v, str): - concat_str += v + parts.append(v) elif isinstance(v, dict): - concat_str += json.dumps(v) + parts.append(json.dumps(v)) else: - concat_str += str(v) + parts.append(str(v)) + concat_str = "".join(parts) hash_object = hashlib.sha256(concat_str.encode()) return hash_object.hexdigest() @@ -5154,7 +5209,7 @@ class Router: verbose_router_logger.debug( f"\nInitialized Model List {self.get_model_names()}" ) - self.model_names = [m["model_name"] for m in model_list] + self.model_names = {m["model_name"] for m in model_list} # Build model_name index for O(1) lookups self._build_model_name_index(self.model_list) @@ -5360,7 +5415,7 @@ class Router: self._add_model_to_list_and_index_map( model=_deployment, model_id=deployment.model_info.id ) - self.model_names.append(deployment.model_name) + self.model_names.add(deployment.model_name) return deployment def _update_deployment_indices_after_removal( @@ -5519,9 +5574,15 @@ class Router: Returns -> Deployment or None Raise Exception -> if model found in invalid format + + Optimized with O(1) index lookup instead of O(n) linear scan. """ - for model in self.model_list: - if model["model_name"] == model_group_name: + # O(1) lookup in model_name index + if model_group_name in self.model_name_to_deployment_indices: + indices = self.model_name_to_deployment_indices[model_group_name] + if indices: + # Return first deployment for this model_name + model = self.model_list[indices[0]] if isinstance(model, dict): return Deployment(**model) elif isinstance(model, Deployment): @@ -5631,11 +5692,13 @@ class Router: Returns - dict: the model in list with 'model_name', 'litellm_params', Optional['model_info'] - None: could not find deployment in list + + Optimized with O(1) index lookup instead of O(n) linear scan. """ - for model in self.model_list: - if "model_info" in model and "id" in model["model_info"]: - if id == model["model_info"]["id"]: - return model + # O(1) lookup via model_id_to_deployment_index_map + if id in self.model_id_to_deployment_index_map: + idx = self.model_id_to_deployment_index_map[id] + return self.model_list[idx] return None def get_model_group(self, id: str) -> Optional[List]: @@ -6169,17 +6232,33 @@ class Router: if 'model_name' is none, returns all. Returns list of model id's. + + Optimized with O(1) or O(k) index lookup when model_name provided, + instead of O(n) linear scan. """ ids = [] - for model in self.model_list: - if "model_info" in model and "id" in model["model_info"]: - id = model["model_info"]["id"] - if exclude_team_models and model["model_info"].get("team_id"): - continue - if model_name is not None and model["model_name"] == model_name: - ids.append(id) - elif model_name is None: - ids.append(id) + + if model_name is not None: + # O(1) lookup in model_name index, then O(k) iteration where k = deployments for this model_name + if model_name in self.model_name_to_deployment_indices: + indices = self.model_name_to_deployment_indices[model_name] + for idx in indices: + model = self.model_list[idx] + if "model_info" in model and "id" in model["model_info"]: + if exclude_team_models and model["model_info"].get("team_id"): + continue + ids.append(model["model_info"]["id"]) + else: + # When model_name is None, return all model IDs + # Use the index map keys for O(n) where n = total deployments + for model_id in self.model_id_to_deployment_index_map.keys(): + idx = self.model_id_to_deployment_index_map[model_id] + model = self.model_list[idx] + if "model_info" in model and "id" in model["model_info"]: + if exclude_team_models and model["model_info"].get("team_id"): + continue + ids.append(model_id) + return ids def has_model_id(self, candidate_id: str) -> bool: @@ -6257,7 +6336,9 @@ class Router: model_name=model_name, model=model, team_id=team_id ): if model_alias is not None: - alias_model = copy.deepcopy(model) + # Optimized: Use shallow copy since we only modify top-level model_name + # This is much faster than deepcopy for nested dict structures + alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: @@ -6271,7 +6352,8 @@ class Router: model_name=model_name, model=model, team_id=team_id ): if model_alias is not None: - alias_model = copy.deepcopy(model) + # Optimized: Use shallow copy since we only modify top-level model_name + alias_model = model.copy() alias_model["model_name"] = model_alias returned_models.append(alias_model) else: @@ -7070,7 +7152,7 @@ class Router: if isinstance(healthy_deployments, dict): return healthy_deployments - start_time = time.time() + start_time = time.perf_counter() if ( self.routing_strategy == "usage-based-routing-v2" and self.lowesttpm_logger_v2 is not None @@ -7137,7 +7219,7 @@ class Router: f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}" ) - end_time = time.time() + end_time = time.perf_counter() _duration = end_time - start_time asyncio.create_task( self.service_logger_obj.async_service_success_hook( diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index e75f969892e..115cce9fb94 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -360,6 +360,22 @@ class PillarGuardrailConfigModel(BaseModel): default="monitor", description="Action to take when content is flagged: 'block' (raise exception) or 'monitor' (log only)", ) + async_mode: Optional[bool] = Field( + default=None, + description="Set to True to request asynchronous analysis (sets `plr_async` header). Defaults to provider behaviour when omitted.", + ) + persist_session: Optional[bool] = Field( + default=None, + description="Controls Pillar session persistence (sets `plr_persist` header). Set to False to disable persistence.", + ) + include_scanners: Optional[bool] = Field( + default=True, + description="Include scanner category summaries in responses (sets `plr_scanners` header).", + ) + include_evidence: Optional[bool] = Field( + default=True, + description="Include detailed evidence payloads in responses (sets `plr_evidence` header).", + ) class NomaGuardrailConfigModel(BaseModel): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py b/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py index e18f8dfb20e..4d0c9ed1cc5 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/pillar.py @@ -15,6 +15,22 @@ class PillarGuardrailConfigModelOptionalParams(BaseModel): default="monitor", description="Action to take when content is flagged: 'block' (raise exception) or 'monitor' (log only). If not provided, the `PILLAR_ON_FLAGGED_ACTION` environment variable is checked, defaults to 'monitor'.", ) + async_mode: Optional[bool] = Field( + default=None, + description="Set to True to request asynchronous analysis (sets `plr_async` header).", + ) + persist_session: Optional[bool] = Field( + default=None, + description="Set to False to disable session persistence (sets `plr_persist` header).", + ) + include_scanners: Optional[bool] = Field( + default=True, + description="Include scanner summaries in response payloads (sets `plr_scanners` header).", + ) + include_evidence: Optional[bool] = Field( + default=True, + description="Include detailed evidence objects in response payloads (sets `plr_evidence` header).", + ) class PillarGuardrailConfigModel( diff --git a/litellm/utils.py b/litellm/utils.py index 5861d703a34..f017543ceaa 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -143,8 +143,10 @@ from litellm.litellm_core_utils.token_counter import get_modified_max_tokens from litellm.llms.base_llm.google_genai.transformation import ( BaseGoogleGenAIGenerateContentConfig, ) +from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig from litellm.llms.bedrock.common_utils import BedrockModelInfo from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.mistral.ocr.transformation import MistralOCRConfig from litellm.router_utils.get_retry_from_policy import ( get_num_retries_from_retry_policy, reset_retry_policy, @@ -5119,9 +5121,9 @@ def json_schema_type(python_type_name: str): return python_to_json_schema_types.get(python_type_name, "string") -def function_to_dict(input_function): # noqa: C901 +def function_to_dict(input_function) -> dict: # noqa: C901 """Using type hints and numpy-styled docstring, - produce a dictionnary usable for OpenAI function calling + produce a dictionary usable for OpenAI function calling Parameters ---------- @@ -7211,6 +7213,8 @@ class ProviderConfigManager: return VolcEngineEmbeddingConfig() elif litellm.LlmProviders.OVHCLOUD == provider: return litellm.OVHCloudEmbeddingConfig() + elif litellm.LlmProviders.COMETAPI == provider: + return litellm.CometAPIEmbeddingConfig() return None @staticmethod @@ -7518,6 +7522,12 @@ class ProviderConfigManager: ) return get_aiml_image_generation_config(model) + elif LlmProviders.COMETAPI == provider: + from litellm.llms.cometapi.image_generation import ( + get_cometapi_image_generation_config, + ) + + return get_cometapi_image_generation_config(model) elif LlmProviders.GEMINI == provider: from litellm.llms.gemini.image_generation import ( get_gemini_image_generation_config, @@ -7549,11 +7559,9 @@ class ProviderConfigManager: provider: LlmProviders, ) -> Optional[BaseImageEditConfig]: if LlmProviders.OPENAI == provider: - from litellm.llms.openai.image_edit.transformation import ( - OpenAIImageEditConfig, - ) + from litellm.llms.openai.image_edit import get_openai_image_edit_config - return OpenAIImageEditConfig() + return get_openai_image_edit_config(model=model) elif LlmProviders.AZURE == provider: from litellm.llms.azure.image_edit.transformation import ( AzureImageEditConfig, @@ -7578,6 +7586,25 @@ class ProviderConfigManager: return LiteLLMProxyImageEditConfig() return None + @staticmethod + def get_provider_ocr_config( + model: str, + provider: LlmProviders, + ) -> Optional["BaseOCRConfig"]: + """ + Get OCR configuration for a given provider. + """ + from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig + + PROVIDER_TO_CONFIG_MAP = { + litellm.LlmProviders.MISTRAL: MistralOCRConfig, + litellm.LlmProviders.AZURE_AI: AzureAIOCRConfig, + } + config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None) + if config_class is None: + return None + return config_class() + @staticmethod def get_provider_google_genai_generate_content_config( model: str, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 91a23b7f00f..8cab321c4dd 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -400,6 +400,44 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "anthropic.claude-haiku-4-5@20251001": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -810,6 +848,25 @@ "supports_tool_choice": true, "supports_vision": true }, + "apac.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "apac.anthropic.claude-3-sonnet-20240229-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -4612,6 +4669,48 @@ "supports_web_search": true, "tool_use_system_prompt_tokens": 264 }, + "claude-haiku-4-5-20251001": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_computer_use": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "claude-haiku-4-5": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_computer_use": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "claude-3-5-sonnet-20240620": { "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -7741,6 +7840,25 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "eu.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "eu.anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -9486,6 +9604,54 @@ "supports_vision": true, "supports_web_search": true }, + "gemini-2.5-flash-image": { + "cache_read_input_token_cost": 3e-08, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "max_pdf_size_mb": 30, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "image_generation", + "output_cost_per_image": 0.039, + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "rpm": 100000, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_vision": true, + "supports_web_search": true, + "tpm": 8000000 + }, "gemini-2.5-flash-image-preview": { "cache_read_input_token_cost": 7.5e-08, "input_cost_per_audio_token": 1e-06, @@ -10939,6 +11105,54 @@ "supports_web_search": true, "tpm": 8000000 }, + "gemini/gemini-2.5-flash-image": { + "cache_read_input_token_cost": 3e-08, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "max_pdf_size_mb": 30, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "image_generation", + "output_cost_per_image": 0.039, + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "rpm": 100000, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-flash-image", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_vision": true, + "supports_web_search": true, + "tpm": 8000000 + }, "gemini/gemini-2.5-flash-image-preview": { "cache_read_input_token_cost": 7.5e-08, "input_cost_per_audio_token": 1e-06, @@ -13197,11 +13411,11 @@ "text" ], "supports_function_calling": true, - "supports_native_streaming": false, + "supports_native_streaming": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, - "supports_reasoning": false, + "supports_reasoning": true, "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, @@ -13650,8 +13864,56 @@ "lemonade/Qwen3-Coder-30B-A3B-Instruct-GGUF": { "input_cost_per_token": 0, "litellm_provider": "lemonade", - "max_tokens": 32768, - "max_input_tokens": 32768, + "max_tokens": 262144, + "max_input_tokens": 262144, + "max_output_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/gpt-oss-20b-mxfp4-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 131072, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/gpt-oss-120b-mxfp-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 131072, + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/Gemma-3-4b-it-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 128000, + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 0, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "lemonade/Qwen3-4B-Instruct-2507-GGUF": { + "input_cost_per_token": 0, + "litellm_provider": "lemonade", + "max_tokens": 262144, + "max_input_tokens": 262144, "max_output_tokens": 32768, "mode": "chat", "output_cost_per_token": 0, @@ -14466,6 +14728,25 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 346 }, + "jp.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "lambda_ai/deepseek-llama3.3-70b": { "input_cost_per_token": 2e-07, "litellm_provider": "lambda_ai", @@ -17153,6 +17434,8 @@ }, "openrouter/anthropic/claude-opus-4": { "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 1.875e-05, + "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "openrouter", "max_input_tokens": 200000, @@ -17163,6 +17446,7 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true, @@ -17170,6 +17454,9 @@ }, "openrouter/anthropic/claude-opus-4.1": { "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, + "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "openrouter", "max_input_tokens": 200000, @@ -17180,6 +17467,7 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true, @@ -17187,6 +17475,10 @@ }, "openrouter/anthropic/claude-sonnet-4": { "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_above_200k_tokens": 6e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, @@ -17199,6 +17491,31 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 159 + }, + "openrouter/anthropic/claude-sonnet-4.5": { + "input_cost_per_image": 0.0048, + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "input_cost_per_token_above_200k_tokens": 6e-06, + "output_cost_per_token_above_200k_tokens": 2.25e-05, + "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true, @@ -20097,6 +20414,25 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "us.anthropic.claude-haiku-4-5-20251001-v1:0": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "bedrock", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://aws.amazon.com/about-aws/whats-new/2025/10/claude-4-5-haiku-anthropic-amazon-bedrock", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "us.anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, "litellm_provider": "bedrock", @@ -21368,6 +21704,25 @@ "supports_pdf_input": true, "supports_tool_choice": true }, + "vertex_ai/claude-haiku-4-5@20251001": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 200000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "vertex_ai/claude-3-5-sonnet": { "input_cost_per_token": 3e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -21560,8 +21915,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, @@ -21577,8 +21932,8 @@ "input_cost_per_token_batches": 7.5e-06, "litellm_provider": "vertex_ai-anthropic_models", "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, + "max_output_tokens": 32000, + "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 7.5e-05, "output_cost_per_token_batches": 3.75e-05, diff --git a/pyproject.toml b/pyproject.toml index 0a379ec752b..12ac25b30a1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "1.78.1" +version = "1.78.3" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT" @@ -157,7 +157,7 @@ requires = ["poetry-core", "wheel"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "1.78.1" +version = "1.78.3" version_files = [ "pyproject.toml:^version" ] diff --git a/test_image_edit.png b/test_image_edit.png index 0380114f241..0386d2af106 100644 Binary files a/test_image_edit.png and b/test_image_edit.png differ diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index 74562c4648f..90544f747bb 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -9,6 +9,7 @@ import base64 from io import BytesIO from unittest.mock import patch, AsyncMock import json +from abc import ABC, abstractmethod sys.path.insert( 0, os.path.abspath("../..") @@ -30,6 +31,72 @@ class TestCustomLogger(CustomLogger): self.standard_logging_payload = kwargs.get("standard_logging_object", None) pass + +class BaseLLMImageEditTest(ABC): + """ + Abstract base test class that enforces a common test across all image edit test classes. + """ + + @property + def image_edit_function(self): + return litellm.image_edit + + @property + def async_image_edit_function(self): + return litellm.aimage_edit + + @abstractmethod + def get_base_image_edit_call_args(self) -> dict: + """Must return the base image edit call args""" + pass + + @pytest.fixture(autouse=True) + def _handle_rate_limits(self): + """Fixture to handle rate limit errors for all test methods""" + try: + yield + except litellm.RateLimitError: + pytest.skip("Rate limit exceeded") + except litellm.InternalServerError: + pytest.skip("Model is overloaded") + + @pytest.mark.parametrize("sync_mode", [True, False]) + @pytest.mark.flaky(retries=3, delay=2) + @pytest.mark.asyncio + async def test_openai_image_edit_litellm_sdk(self, sync_mode): + """ + Test image edit functionality with both sync and async modes. + """ + litellm._turn_on_debug() + try: + prompt = """ + Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. + """ + + call_args = self.get_base_image_edit_call_args() + call_args["prompt"] = prompt + + if sync_mode: + result = self.image_edit_function(**call_args) + else: + result = await self.async_image_edit_function(**call_args) + + print("result from image edit", result) + + # Validate the response meets expected schema + ImageResponse.model_validate(result) + + if isinstance(result, ImageResponse) and result.data: + image_base64 = result.data[0].b64_json + if image_base64: + image_bytes = base64.b64decode(image_base64) + + # Save the image to a file + with open("test_image_edit.png", "wb") as f: + f.write(image_bytes) + except litellm.ContentPolicyViolationError as e: + pass + # Get the current directory of the file being run pwd = os.path.dirname(os.path.realpath(__file__)) @@ -49,45 +116,31 @@ def get_test_images_as_bytesio(): bytesio_images.append(BytesIO(image_bytes)) return bytesio_images -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_openai_image_edit_litellm_sdk(sync_mode): - from litellm import image_edit, aimage_edit - litellm._turn_on_debug() - try: - prompt = """ - Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. - """ - if sync_mode: - result = image_edit( - prompt=prompt, - model="gpt-image-1", - image=TEST_IMAGES, - ) - else: - result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=TEST_IMAGES, - ) - print("result from image edit", result) +class TestOpenAIImageEditGPTImage1(BaseLLMImageEditTest): + """ + Concrete implementation of BaseLLMImageEditTest for OpenAI image edits. + """ - # Validate the response meets expected schema - ImageResponse.model_validate(result) - - if isinstance(result, ImageResponse) and result.data: - image_base64 = result.data[0].b64_json - if image_base64: - image_bytes = base64.b64decode(image_base64) + def get_base_image_edit_call_args(self) -> dict: + """Return base call args for OpenAI image edit""" + return { + "model": "gpt-image-1", + "image": TEST_IMAGES, + } - # Save the image to a file - with open("test_image_edit.png", "wb") as f: - f.write(image_bytes) - except litellm.ContentPolicyViolationError as e: - pass +class TestOpenAIImageEditDallE2(BaseLLMImageEditTest): + """ + Concrete implementation of BaseLLMImageEditTest for OpenAI DALL-E-2 image edits. + DALL-E-2 only supports a single image (not an array). + """ + def get_base_image_edit_call_args(self) -> dict: + """Return base call args for OpenAI DALL-E-2 image edit (single image only)""" + return { + "model": "dall-e-2", + "image": SINGLE_TEST_IMAGE, + } @pytest.mark.flaky(retries=3, delay=2) diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index d41448727d5..81593fb3f41 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3232,6 +3232,60 @@ async def test_bedrock_passthrough(sync_mode: bool): assert response.status_code == 200 +@pytest.mark.asyncio +async def test_bedrock_passthrough_router(): + """ + Test bedrock passthrough using litellm.Router with async mode. + Tests that the router: + 1. Resolves the router model name to the actual deployment + 2. Replaces the router model name in the endpoint with the actual deployment model + """ + import litellm + from litellm import Router + + litellm._turn_on_debug() + + router = Router( + model_list=[ + { + "model_name": "special-bedrock-model", + "litellm_params": { + "model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + }, + } + ] + ) + + data = { + "max_tokens": 512, + "messages": [{"role": "user", "content": "Hey"}], + "system": [ + { + "type": "text", + "text": "Analyze if this message indicates a new conversation topic. If it does, extract a 2-3 word title that captures the new topic. Format your response as a JSON object with two fields: 'isNewTopic' (boolean) and 'title' (string, or null if isNewTopic is false). Only include these fields, no other text.", + } + ], + "temperature": 0, + "metadata": { + "user_id": "5dd07c33da27e6d2968d94ea20bf47a7b090b6b158b82328d54da2909a108e84" + }, + "anthropic_version": "bedrock-2023-05-31", + "anthropic_beta": ["claude-code-20250219"], + } + + # Endpoint uses the router model name which should be replaced with actual deployment + response = await router.allm_passthrough_route( + model="special-bedrock-model", + method="POST", + endpoint="/model/special-bedrock-model/invoke", + data=data, + ) + + print(response.text) + + assert response.status_code == 200 + + @pytest.mark.asyncio async def test_bedrock_converse__streaming_passthrough(monkeypatch): import litellm diff --git a/tests/ocr_tests/base_ocr_unit_tests.py b/tests/ocr_tests/base_ocr_unit_tests.py new file mode 100644 index 00000000000..12264844792 --- /dev/null +++ b/tests/ocr_tests/base_ocr_unit_tests.py @@ -0,0 +1,140 @@ +""" +Base test class for OCR functionality across different providers. + +This follows the same pattern as BaseLLMChatTest in tests/llm_translation/base_llm_unit_tests.py +""" +import pytest +import litellm +from abc import ABC, abstractmethod + + +# Test resources +TEST_IMAGE_PATH = "test_image_edit.png" +TEST_PDF_URL = "https://arxiv.org/pdf/2201.04234" + + +class BaseOCRTest(ABC): + """ + Abstract base test class that enforces common OCR tests across all providers. + + Each provider-specific test class should inherit from this and implement + get_base_ocr_call_args() to return provider-specific configuration. + """ + + @abstractmethod + def get_base_ocr_call_args(self) -> dict: + """Must return the base OCR call args for the specific provider""" + pass + + @pytest.fixture(autouse=True) + def _handle_rate_limits(self): + """Fixture to handle rate limit errors for all test methods""" + try: + yield + except litellm.RateLimitError: + pytest.skip("Rate limit exceeded") + except litellm.InternalServerError: + pytest.skip("Model is overloaded") + + @pytest.mark.parametrize("sync_mode", [True, False]) + @pytest.mark.asyncio + async def test_basic_ocr_with_url(self, sync_mode): + """ + Test basic OCR with a public URL. + """ + litellm._turn_on_debug() + base_ocr_call_args = self.get_base_ocr_call_args() + print("BASE OCR Call args=", base_ocr_call_args) + + try: + if sync_mode: + response = litellm.ocr( + document={ + "type": "document_url", + "document_url": TEST_PDF_URL + }, + **base_ocr_call_args, + ) + else: + response = await litellm.aocr( + document={ + "type": "document_url", + "document_url": TEST_PDF_URL + }, + **base_ocr_call_args, + ) + + print(f"\n{'='*80}") + print(f"Sync Mode: {sync_mode}") + print(f"Response type: {type(response)}") + print(f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}") + + # Check if response has expected OCR format + assert hasattr(response, "pages"), "Response should have 'pages' attribute" + assert hasattr(response, "model"), "Response should have 'model' attribute" + assert hasattr(response, "object"), "Response should have 'object' attribute" + assert response.object == "ocr", f"Expected object='ocr', got '{response.object}'" + + # Validate pages structure + assert isinstance(response.pages, list), "pages should be a list" + assert len(response.pages) > 0, "Should have at least one page" + + # Check first page structure + first_page = response.pages[0] + assert hasattr(first_page, "index"), "Page should have 'index' attribute" + assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute" + + # Extract text from all pages for validation + total_text = "\n\n".join(page.markdown for page in response.pages if page.markdown) + print(f"Total pages: {len(response.pages)}") + print(f"Total extracted text length: {len(total_text)} characters") + print(f"First 200 chars: {total_text[:200]}") + print(f"Model: {response.model}") + if response.usage_info: + print(f"Pages processed: {response.usage_info.pages_processed}") + print(f"{'='*80}\n") + + assert len(total_text) > 0, "Should extract some text from the document" + + except Exception as e: + pytest.fail(f"OCR call failed: {str(e)}") + + def test_ocr_response_structure(self): + """ + Test that the OCR response has the correct structure. + """ + litellm.set_verbose = True + base_ocr_call_args = self.get_base_ocr_call_args() + + response = litellm.ocr( + document={ + "type": "document_url", + "document_url": TEST_PDF_URL + }, + **base_ocr_call_args, + ) + + # Validate response structure + assert hasattr(response, "pages"), "Response should have 'pages' attribute" + assert hasattr(response, "model"), "Response should have 'model' attribute" + assert hasattr(response, "object"), "Response should have 'object' attribute" + assert hasattr(response, "usage_info"), "Response should have 'usage_info' attribute" + + assert isinstance(response.pages, list), "pages should be a list" + assert len(response.pages) > 0, "Should have at least one page" + assert response.object == "ocr", "object should be 'ocr'" + + # Validate first page structure + first_page = response.pages[0] + assert hasattr(first_page, "index"), "Page should have 'index' attribute" + assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute" + assert isinstance(first_page.markdown, str), "markdown should be a string" + + print(f"\nResponse structure validated:") + print(f" - object: {response.object}") + print(f" - model: {response.model}") + print(f" - pages: {len(response.pages)}") + if response.usage_info: + print(f" - pages_processed: {response.usage_info.pages_processed}") + print(f" - doc_size_bytes: {response.usage_info.doc_size_bytes}") + diff --git a/tests/ocr_tests/test_ocr_azure_ai.py b/tests/ocr_tests/test_ocr_azure_ai.py new file mode 100644 index 00000000000..34fd9b37ba2 --- /dev/null +++ b/tests/ocr_tests/test_ocr_azure_ai.py @@ -0,0 +1,27 @@ +""" +Test OCR functionality with Azure AI API. + +Note: Azure AI OCR automatically converts URLs to base64 data URIs since +the Azure AI endpoint doesn't have internet access. +""" +import os +from base_ocr_unit_tests import BaseOCRTest + +class TestAzureAIOCR(BaseOCRTest): + """ + Test class for Azure AI OCR functionality. + Inherits from BaseOCRTest and provides Azure AI-specific configuration. + + Note: For Azure AI, LiteLLM will automatically convert URLs to base64 data URIs before + sending to the API, since Azure AI OCR endpoint doesn't have internet access. + """ + + def get_base_ocr_call_args(self) -> dict: + """ + Return the base OCR call args for Azure AI. + """ + return { + "model": "azure_ai/mistral-document-ai-2505", + "api_key": os.getenv("AZURE_AI_API_KEY_MISTRAL"), + "api_base": os.getenv("AZURE_AI_API_BASE_MISTRAL"), + } diff --git a/tests/ocr_tests/test_ocr_mistral.py b/tests/ocr_tests/test_ocr_mistral.py new file mode 100644 index 00000000000..ea647459271 --- /dev/null +++ b/tests/ocr_tests/test_ocr_mistral.py @@ -0,0 +1,88 @@ +""" +Test OCR functionality with Mistral API. +""" +import os +import sys +import pytest +import litellm +from litellm import Router +from base_ocr_unit_tests import BaseOCRTest, TEST_PDF_URL + + +class TestMistralOCR(BaseOCRTest): + """ + Test class for Mistral OCR functionality. + """ + + def get_base_ocr_call_args(self) -> dict: + """Return the base OCR call args for Mistral""" + return { + "model": "mistral/mistral-ocr-latest", + "api_key": os.getenv("MISTRAL_API_KEY"), + } + +@pytest.mark.asyncio +async def test_router_aocr_with_mistral(): + """ + Test OCR with Router using Mistral OCR deployment. + """ + litellm.set_verbose = True + + # Create router with Mistral OCR deployment + router = Router( + model_list=[ + { + "model_name": "mistral-ocr", + "litellm_params": { + "model": "mistral/mistral-ocr-latest", + "api_key": os.getenv("MISTRAL_API_KEY"), + }, + } + ] + ) + + try: + # Call OCR through router + response = await router.aocr( + model="mistral-ocr", + document={ + "type": "document_url", + "document_url": TEST_PDF_URL + }, + ) + + print(f"\n{'='*80}") + print("Router OCR Test") + print(f"Response type: {type(response)}") + print(f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}") + + # Check if response has expected Mistral OCR format + assert hasattr(response, "pages"), "Response should have 'pages' attribute" + assert hasattr(response, "model"), "Response should have 'model' attribute" + assert hasattr(response, "object"), "Response should have 'object' attribute" + assert response.object == "ocr", f"Expected object='ocr', got '{response.object}'" + + # Validate pages structure + assert isinstance(response.pages, list), "pages should be a list" + assert len(response.pages) > 0, "Should have at least one page" + + # Check first page structure + first_page = response.pages[0] + assert hasattr(first_page, "index"), "Page should have 'index' attribute" + assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute" + + # Extract text from all pages for validation + total_text = "\n\n".join(page.markdown for page in response.pages if page.markdown) + print(f"Total pages: {len(response.pages)}") + print(f"Total extracted text length: {len(total_text)} characters") + print(f"First 200 chars: {total_text[:200]}") + print(f"Model: {response.model}") + if response.usage_info: + print(f"Pages processed: {response.usage_info.pages_processed}") + print(f"{'='*80}\n") + + assert len(total_text) > 0, "Should extract some text from the document" + + except Exception as e: + pytest.fail(f"Router OCR call failed: {str(e)}") + diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index a31c4d8210f..c2339d9eec5 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -1411,7 +1411,8 @@ def test_generate_model_id_with_deployment_model_name(model_list): "Expected TypeError when model_group is None - this confirms our fix is needed" ) except TypeError as e: - assert "unsupported operand type(s) for +=" in str(e) + # After optimization, error message changed but still fails appropriately on None + assert "unsupported operand type(s) for +=" in str(e) or "expected str instance, NoneType found" in str(e) print(f"✓ Correctly failed with None model_group (as expected): {e}") except Exception as e: pytest.fail(f"Unexpected error with None model_group: {e}") diff --git a/tests/router_unit_tests/test_router_index_management.py b/tests/router_unit_tests/test_router_index_management.py index 04ea9214991..0c313e9ef4e 100644 --- a/tests/router_unit_tests/test_router_index_management.py +++ b/tests/router_unit_tests/test_router_index_management.py @@ -1,6 +1,8 @@ import sys import os import pytest +import ast +import ast sys.path.insert( 0, os.path.abspath("../..") @@ -177,3 +179,97 @@ class TestRouterIndexManagement: # Verify: New entry is added assert "claude-3" in router.model_name_to_deployment_indices assert router.model_name_to_deployment_indices["claude-3"] == [0] + + def test_no_linear_scans_in_router(self): + """ + Static analysis test to ensure Router doesn't use O(n) linear scans. + + Scans router.py for 'in self.model_list' pattern which indicates + inefficient O(n) iteration instead of using index-based O(1) lookups. + + Methods should use: + - model_id_to_deployment_index_map for O(1) model_id lookups + - model_name_to_deployment_indices for O(1) + O(k) model_name lookups + """ + # Methods that are allowed to iterate through self.model_list + ALLOWED_METHODS = [ + "_get_deployment_by_litellm_model", # Edge case: lookup by litellm_params.model (not indexed) + ] + + # Get path to router.py + router_file = os.path.join( + os.path.dirname(os.path.dirname(os.path.dirname(__file__))), + "litellm", + "router.py" + ) + + # Read the file + with open(router_file, 'r') as f: + content = f.read() + + # Parse with AST + tree = ast.parse(content) + + # Find violations + violations = [] + ignore_methods = set(ALLOWED_METHODS) + + for node in ast.walk(tree): + if isinstance(node, ast.FunctionDef): + method_name = node.name + + # Skip ignored methods + if method_name in ignore_methods: + continue + + # Get source for this method + try: + method_source = ast.get_source_segment(content, node) + if not method_source: + continue + + # Check for the anti-pattern: "in self.model_list" + # This catches: for x in self.model_list, if x in self.model_list, etc. + if "in self.model_list" in method_source: + # Extract the specific line for better error reporting + lines = method_source.split('\n') + pattern_line = None + for line in lines: + if "in self.model_list" in line: + pattern_line = line.strip() + break + + violations.append({ + "method": method_name, + "line": node.lineno, + "pattern": pattern_line or "in self.model_list" + }) + except Exception: + # Skip if we can't get source segment + pass + + # Assert no violations + if violations: + error_msg = "\n".join([ + f" - {v['method']}() at line {v['line']}: {v['pattern']}" + for v in violations + ]) + + pytest.fail( + f"\n{'='*70}\n" + f"Found O(n) linear scan pattern in router.py:\n\n" + f"{error_msg}\n\n" + f"These methods should use index maps instead:\n" + f" - model_id_to_deployment_index_map (for model_id lookups)\n" + f" - model_name_to_deployment_indices (for model_name lookups)\n\n" + f"If a method legitimately needs O(n) iteration, add it to\n" + f"ALLOWED_METHODS in this test method.\n" + f"{'='*70}\n" + ) + def test_model_names_is_set(self): + """Verify that model_names uses a set for O(1) lookups, not a list (O(n))""" + router = Router(model_list=[]) + + assert isinstance(router.model_names, set), ( + f"model_names should be a set for O(1) lookups, but got {type(router.model_names)}" + ) diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 1541eea4939..664ecb896b6 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -649,5 +649,5 @@ def test_bedrock_anthropic_prompt_caching(): assert prompt_cost >= 0 assert completion_cost >= 0 - assert round(prompt_cost, 3) == 0.845 + assert round(prompt_cost, 3) == 0.111 assert round(completion_cost, 5) == 0.00820 diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 0e31699fd83..09fc31d18b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -399,3 +399,41 @@ async def test_session_validation(): mock_valid_session = MockClientSession() transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore assert transport3.client is mock_valid_session # Should reuse session + + +@pytest.mark.parametrize( + "env_curve,litellm_curve,expected_curve,should_call", + [ + # env_curve: SSL_ECDH_CURVE env var | litellm_curve: litellm.ssl_ecdh_curve variable + # expected_curve: curve that should be set | should_call: whether set_ecdh_curve() should be called + + # Valid configurations + ("X25519", None, "X25519", True), # Env var only + ("prime256v1", None, "prime256v1", True), # Different valid curve + (None, "secp384r1", "secp384r1", True), # litellm variable only + ("X25519", "secp521r1", "X25519", True), # Env var takes precedence + # Empty/None configurations - should skip + ("", None, None, False), # Empty string - skip configuration + (None, None, None, False), # None value - skip configuration + ] +) +def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, monkeypatch): + """Test SSL ECDH curve configuration with valid curves and precedence""" + with patch.dict(os.environ, clear=True): + if env_curve: + monkeypatch.setenv("SSL_ECDH_CURVE", env_curve) + + original_value = litellm.ssl_ecdh_curve + try: + litellm.ssl_ecdh_curve = litellm_curve + + with patch.object(ssl.SSLContext, 'set_ecdh_curve') as mock_set_curve: + ssl_context = get_ssl_configuration() + + if should_call: + mock_set_curve.assert_called_once_with(expected_curve) + else: + mock_set_curve.assert_not_called() + assert isinstance(ssl_context, ssl.SSLContext) + finally: + litellm.ssl_ecdh_curve = original_value diff --git a/tests/guardrails_tests/test_pillar_guardrails.py b/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py similarity index 90% rename from tests/guardrails_tests/test_pillar_guardrails.py rename to tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py index aeb2227f9b9..67030a8161f 100644 --- a/tests/guardrails_tests/test_pillar_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py @@ -8,8 +8,12 @@ and following LiteLLM testing patterns and best practices. # Standard library imports import os import sys +from typing import Dict from unittest.mock import Mock, patch +# Add parent directory to path for imports +sys.path.insert(0, os.path.abspath("../../..")) + # Third-party imports import pytest from fastapi.exceptions import HTTPException @@ -26,9 +30,6 @@ from litellm.proxy.guardrails.guardrail_hooks.pillar import ( ) from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 -# Add parent directory to path for imports -sys.path.insert(0, os.path.abspath("../..")) - # ============================================================================ # FIXTURES @@ -221,6 +222,18 @@ def mock_llm_response(): return mock_response +@pytest.fixture +def pillar_async_response(): + """Fixture providing an asynchronous Pillar API queue response.""" + return Response( + json={"status": "queued", "session_id": "async-session", "position": 1}, + status_code=202, + request=Request( + method="POST", url="https://api.pillar.security/api/v1/protect" + ), + ) + + @pytest.fixture def mock_llm_response_with_tools(): """Fixture providing a mock LLM response with tool calls.""" @@ -440,6 +453,55 @@ async def test_post_call_hook_with_tool_calls( assert result == mock_llm_response_with_tools +# ========================================================================= +# HEADER CONFIGURATION TESTS +# ========================================================================= + + +@pytest.mark.asyncio +async def test_pre_call_hook_custom_header_overrides( + sample_request_data, + user_api_key_dict, + dual_cache, + pillar_async_response, +): + """Ensure configuration values translate into correct Protect headers.""" + + guardrail = PillarGuardrail( + guardrail_name="pillar-header-test", + api_key="test-pillar-key", + api_base="https://api.pillar.security", + on_flagged_action="monitor", + persist_session=False, + async_mode=True, + include_scanners=False, + include_evidence=False, + ) + + captured_headers: Dict[str, str] = {} + + async def _mock_post(*args, **kwargs): + captured_headers.update(kwargs.get("headers", {})) + return pillar_async_response + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new=_mock_post, + ): + result = await guardrail.async_pre_call_hook( + data=sample_request_data, + cache=dual_cache, + user_api_key_dict=user_api_key_dict, + call_type="completion", + ) + + assert result == sample_request_data + assert captured_headers.get("plr_persist") == "false" + assert captured_headers.get("plr_async") == "true" + assert captured_headers.get("plr_scanners") == "false" + assert captured_headers.get("plr_evidence") == "false" + + # ============================================================================ # EDGE CASE TESTS # ============================================================================ diff --git a/tests/test_litellm/proxy/management_endpoints/test_entraid_app_roles.py b/tests/test_litellm/proxy/management_endpoints/test_entraid_app_roles.py new file mode 100644 index 00000000000..f6248d36628 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_entraid_app_roles.py @@ -0,0 +1,55 @@ +""" +Unit tests for EntraID app roles JWT claim extraction. + +This module tests the get_app_roles_from_id_token method to ensure it correctly +extracts app roles from Microsoft EntraID JWT tokens and prevents regressions. +""" + +import pytest +import jwt + +from litellm.proxy.management_endpoints.ui_sso import MicrosoftSSOHandler + + +class TestEntraIDAppRoles: + """Test EntraID app roles extraction from JWT tokens""" + + def test_get_app_roles_from_id_token_works_without_roles(self): + """Test that JWT token works fine without app_roles claim""" + # Arrange - Token without app_roles (normal user) + payload = { + "sub": "user123", + "email": "user@company.com", + "aud": "litellm-app", + "iss": "https://login.microsoftonline.com/tenant-id/v2.0", + "exp": 9999999999, + } + no_roles_token = jwt.encode(payload, "secret", algorithm="HS256") + + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(no_roles_token) + + # Assert - Should return empty list, not error + assert result == [] + assert len(result) == 0 + + def test_get_app_roles_from_id_token_assigns_roles_when_present(self): + """Test that valid app roles are properly assigned when present""" + # Arrange - Token with valid roles + payload = { + "sub": "user123", + "email": "admin@company.com", + "app_roles": ["proxy_admin"], + "aud": "litellm-app", + "iss": "https://login.microsoftonline.com/tenant-id/v2.0", + "exp": 9999999999, + } + valid_roles_token = jwt.encode(payload, "secret", algorithm="HS256") + + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(valid_roles_token) + + # Assert - Should extract the role + assert result == ["proxy_admin"] + assert len(result) == 1 + assert "proxy_admin" in result 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 941517db120..0f403ae5d65 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -99,7 +99,9 @@ def test_microsoft_sso_handler_with_empty_response(): # Test with None response # Act - result = MicrosoftSSOHandler.openid_from_response(response=None, team_ids=[], user_role=None) + result = MicrosoftSSOHandler.openid_from_response( + response=None, team_ids=[], user_role=None + ) # Assert assert isinstance(result, CustomOpenID) @@ -789,11 +791,11 @@ class TestCLISSOCallbackFunction: "not-sk-key", "sk", # too short ] - + for invalid_key in invalid_keys: # This should fail validation before any database operations # We can test this by checking if the key starts with 'sk-' - if not invalid_key or not invalid_key.startswith('sk-'): + if not invalid_key or not invalid_key.startswith("sk-"): # This would trigger the validation error assert True # Validation works as expected @@ -806,14 +808,14 @@ class TestCLIPollingFunction: # Test key format validation logic invalid_keys = [ "invalid-key", - "not-sk-key", + "not-sk-key", "", "sk", # too short ] - + for invalid_key in invalid_keys: # Validation logic: key must start with 'sk-' - if not invalid_key.startswith('sk-'): + if not invalid_key.startswith("sk-"): # This would trigger the validation error in the actual function assert True # Validation works as expected @@ -827,7 +829,7 @@ class TestAuthCallbackRouting: # Test CLI state detection logic cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123" - + # This mimics the logic in auth_callback if cli_state and cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"): # Extract the key ID from the state @@ -839,18 +841,20 @@ class TestAuthCallbackRouting: def test_non_cli_state_routing(self): """Test that non-CLI states don't trigger CLI routing""" from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX - + non_cli_states = [ "regular_oauth_state", - "some_random_string", + "some_random_string", None, "", - "not_session_token:something" + "not_session_token:something", ] - + for state in non_cli_states: # This mimics the routing logic in auth_callback - should_route_to_cli = state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") + should_route_to_cli = state and state.startswith( + f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:" + ) assert not should_route_to_cli, f"State '{state}' should not route to CLI" @@ -864,9 +868,9 @@ class TestGoogleLoginCLIIntegration: # Test the CLI state generation logic used in google_login source = "litellm-cli" key = "sk-test123" - + cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key) - + assert cli_state is not None assert cli_state.startswith("litellm-session-token:") assert "sk-test123" in cli_state @@ -882,10 +886,12 @@ class TestGoogleLoginCLIIntegration: (None, "sk-test123"), ("wrong-source", "sk-test123"), ] - + for source, key in test_cases: cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key) - assert cli_state is None, f"CLI state should not be generated for source='{source}', key='{key}'" + assert ( + cli_state is None + ), f"CLI state should not be generated for source='{source}', key='{key}'" class TestSSOHandlerIntegration: @@ -896,13 +902,24 @@ class TestSSOHandlerIntegration: from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler # Test that SSO handler is used when client IDs are provided - assert SSOAuthenticationHandler.should_use_sso_handler(google_client_id="test") is True - assert SSOAuthenticationHandler.should_use_sso_handler(microsoft_client_id="test") is True - assert SSOAuthenticationHandler.should_use_sso_handler(generic_client_id="test") is True - + assert ( + SSOAuthenticationHandler.should_use_sso_handler(google_client_id="test") + is True + ) + assert ( + SSOAuthenticationHandler.should_use_sso_handler(microsoft_client_id="test") + is True + ) + assert ( + SSOAuthenticationHandler.should_use_sso_handler(generic_client_id="test") + is True + ) + # Test that SSO handler is not used when no client IDs are provided assert SSOAuthenticationHandler.should_use_sso_handler() is False - assert SSOAuthenticationHandler.should_use_sso_handler(None, None, None) is False + assert ( + SSOAuthenticationHandler.should_use_sso_handler(None, None, None) is False + ) def test_get_redirect_url_for_sso(self): """Test the redirect URL generation for SSO""" @@ -911,13 +928,12 @@ class TestSSOHandlerIntegration: # Mock request object mock_request = MagicMock() mock_request.base_url = "https://test.litellm.ai/" - + # Test redirect URL generation redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( - request=mock_request, - sso_callback_route="sso/callback" + request=mock_request, sso_callback_route="sso/callback" ) - + assert redirect_url.startswith("https://test.litellm.ai") assert "sso/callback" in redirect_url @@ -928,21 +944,25 @@ class TestUISSO_FunctionsExistence: def test_cli_sso_callback_exists(self): """Test that cli_sso_callback function exists""" from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback + assert callable(cli_sso_callback) def test_cli_poll_key_exists(self): """Test that cli_poll_key function exists""" from litellm.proxy.management_endpoints.ui_sso import cli_poll_key + assert callable(cli_poll_key) def test_auth_callback_exists(self): """Test that auth_callback function exists""" from litellm.proxy.management_endpoints.ui_sso import auth_callback + assert callable(auth_callback) def test_google_login_exists(self): """Test that google_login function exists""" from litellm.proxy.management_endpoints.ui_sso import google_login + assert callable(google_login) def test_sso_authentication_handler_exists(self): @@ -951,9 +971,9 @@ class TestUISSO_FunctionsExistence: # Check that the class exists assert SSOAuthenticationHandler is not None - + # Check that the new _get_cli_state method exists - assert hasattr(SSOAuthenticationHandler, '_get_cli_state') + assert hasattr(SSOAuthenticationHandler, "_get_cli_state") assert callable(SSOAuthenticationHandler._get_cli_state) @@ -963,9 +983,11 @@ class TestSSOStateHandling: def test_get_cli_state_valid(self): """Test generating CLI state with valid parameters""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - - state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key="sk-test123") - + + state = SSOAuthenticationHandler._get_cli_state( + source="litellm-cli", key="sk-test123" + ) + assert state is not None assert state.startswith("litellm-session-token:") assert "sk-test123" in state @@ -973,37 +995,39 @@ class TestSSOStateHandling: def test_get_cli_state_invalid_source(self): """Test generating CLI state with invalid source""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - - state = SSOAuthenticationHandler._get_cli_state(source="invalid_source", key="sk-test123") - + + state = SSOAuthenticationHandler._get_cli_state( + source="invalid_source", key="sk-test123" + ) + assert state is None def test_get_cli_state_no_key(self): """Test generating CLI state without key""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key=None) - + assert state is None def test_get_cli_state_no_source(self): """Test generating CLI state without source""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state(source=None, key="sk-test123") - + assert state is None def test_get_cli_state_with_existing_key(self): """Test generating CLI state with existing_key embedded in state parameter""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state( - source="litellm-cli", + source="litellm-cli", key="sk-new-key-123", - existing_key="sk-existing-key-456" + existing_key="sk-existing-key-456", ) - + assert state is not None assert state.startswith("litellm-session-token:") assert "sk-new-key-123" in state @@ -1014,13 +1038,11 @@ class TestSSOStateHandling: def test_get_cli_state_without_existing_key(self): """Test generating CLI state without existing_key""" from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + state = SSOAuthenticationHandler._get_cli_state( - source="litellm-cli", - key="sk-new-key-789", - existing_key=None + source="litellm-cli", key="sk-new-key-789", existing_key=None ) - + assert state is not None assert state.startswith("litellm-session-token:") assert "sk-new-key-789" in state @@ -1039,7 +1061,7 @@ class TestStateRouting: # Test CLI state format cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123" assert cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") - + # Test extraction of key from state key_id = cli_state.split(":", 1)[1] assert key_id == "sk-test123" @@ -1049,13 +1071,15 @@ class TestStateRouting: from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX # State format: {PREFIX}:{key}:{existing_key} - cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-456:sk-existing-key-789" - + cli_state = ( + f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-456:sk-existing-key-789" + ) + # Parse as done in auth_callback state_parts = cli_state.split(":", 2) # Split into max 3 parts key_id = state_parts[1] if len(state_parts) > 1 else None existing_key = state_parts[2] if len(state_parts) > 2 else None - + assert key_id == "sk-new-key-456" assert existing_key == "sk-existing-key-789" @@ -1065,12 +1089,12 @@ class TestStateRouting: # State format: {PREFIX}:{key} cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-999" - + # Parse as done in auth_callback state_parts = cli_state.split(":", 2) # Split into max 3 parts key_id = state_parts[1] if len(state_parts) > 1 else None existing_key = state_parts[2] if len(state_parts) > 2 else None - + assert key_id == "sk-new-key-999" assert existing_key is None @@ -1084,9 +1108,9 @@ class TestStateRouting: "some_random_string", None, "", - "not_session_token:something" + "not_session_token:something", ] - + for state in test_states: if state: assert not state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") @@ -1105,10 +1129,10 @@ class TestHTMLIntegration: # Test that function exists and is callable assert callable(render_cli_sso_success_page) - + # Test that it returns expected type html = render_cli_sso_success_page() - + assert isinstance(html, str) assert len(html) > 0 @@ -1125,11 +1149,19 @@ class TestCustomUISSO: # Mock request mock_request = MagicMock() mock_request.base_url = "https://test.example.com/" - + # Mock user_custom_ui_sso_sign_in_handler to exist but make enterprise import fail with patch("litellm.proxy.proxy_server.premium_user", True): - with patch("litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", MagicMock()): - with patch.dict('sys.modules', {'enterprise.litellm_enterprise.proxy.auth.custom_sso_handler': None}): + with patch( + "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", + MagicMock(), + ): + with patch.dict( + "sys.modules", + { + "enterprise.litellm_enterprise.proxy.auth.custom_sso_handler": None + }, + ): # Temporarily mock the google_login function call to test the import error path async def mock_google_login(): # This mimics the relevant part of google_login that would trigger the import error @@ -1137,13 +1169,19 @@ class TestCustomUISSO: from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import ( EnterpriseCustomSSOHandler, ) + return "success" except ImportError: - raise ValueError("Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise.") - + raise ValueError( + "Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise." + ) + # Test that the ValueError is raised with the correct message import pytest - with pytest.raises(ValueError, match="Enterprise features are not available"): + + with pytest.raises( + ValueError, match="Enterprise features are not available" + ): asyncio.run(mock_google_login()) @pytest.mark.asyncio @@ -1195,8 +1233,10 @@ class TestCustomUISSO: return_value=mock_redirect_response, ) as mock_get_redirect: # Act - result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( - request=mock_request + result = ( + await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( + request=mock_request + ) ) # Assert @@ -1241,12 +1281,14 @@ class TestCustomUISSO: async def handle_custom_ui_sso_sign_in(self, request: Request) -> OpenID: self.method_called = True self.received_request = request - + # Parse headers like the actual implementation would request_headers_dict = dict(request.headers) return OpenID( id=request_headers_dict.get("x-litellm-user-id", "default_user"), - email=request_headers_dict.get("x-litellm-user-email", "default@test.com"), + email=request_headers_dict.get( + "x-litellm-user-email", "default@test.com" + ), first_name="Custom", last_name="Handler", display_name="Custom Handler Test", @@ -1260,7 +1302,7 @@ class TestCustomUISSO: # Mock request with custom headers mock_request = MagicMock(spec=Request) mock_request.headers = { - "x-litellm-user-id": "custom_test_user_456", + "x-litellm-user-id": "custom_test_user_456", "x-litellm-user-email": "custom@example.com", "x-forwarded-for": "10.0.0.1", } @@ -1281,8 +1323,10 @@ class TestCustomUISSO: return_value=mock_redirect_response, ) as mock_get_redirect: # Act - result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( - request=mock_request + result = ( + await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( + request=mock_request + ) ) # Assert that our custom handler was executed @@ -1292,7 +1336,7 @@ class TestCustomUISSO: # Verify the redirect response was called with the OpenID from our custom handler mock_get_redirect.assert_called_once() call_args = mock_get_redirect.call_args.kwargs - + # Verify the OpenID object has the expected values from our custom handler openid_result = call_args["result"] assert openid_result.id == "custom_test_user_456" @@ -1323,25 +1367,30 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock(spec=Request) - + # Test data existing_key = "sk-existing-key-123" new_key = "sk-new-key-456" - + # Mock the regenerate helper function - with patch("litellm.proxy.management_endpoints.ui_sso._regenerate_cli_key") as mock_regenerate, \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success"): - + with patch( + "litellm.proxy.management_endpoints.ui_sso._regenerate_cli_key" + ) as mock_regenerate, patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ): + # Act result = await cli_sso_callback( - request=mock_request, - key=new_key, - existing_key=existing_key + request=mock_request, key=new_key, existing_key=existing_key ) - + # Assert - mock_regenerate.assert_called_once_with(existing_key=existing_key, new_key=new_key, user_id=None) + mock_regenerate.assert_called_once_with( + existing_key=existing_key, new_key=new_key, user_id=None + ) assert result.status_code == 200 assert "Success" in result.body.decode() @@ -1352,22 +1401,25 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock(spec=Request) - + # Test data new_key = "sk-new-key-789" - + # Mock the create helper function - with patch("litellm.proxy.management_endpoints.ui_sso._create_new_cli_key") as mock_create, \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success"): - + with patch( + "litellm.proxy.management_endpoints.ui_sso._create_new_cli_key" + ) as mock_create, patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ): + # Act result = await cli_sso_callback( - request=mock_request, - key=new_key, - existing_key=None + request=mock_request, key=new_key, existing_key=None ) - + # Assert mock_create.assert_called_once_with(key=new_key, user_id=None) assert result.status_code == 200 @@ -1381,32 +1433,42 @@ class TestCLIKeyRegenerationFlow: # Mock request (no query params needed - existing_key is in state) mock_request = MagicMock(spec=Request) - + # CLI state with existing_key embedded: {PREFIX}:{key}:{existing_key} cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-session-key-456:sk-existing-cli-key-123" - + # Mock the CLI callback and required proxy server components mock_result = {"user_id": "test-user", "email": "test@example.com"} - - with patch("litellm.proxy.management_endpoints.ui_sso.cli_sso_callback") as mock_cli_callback, \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.proxy_server.master_key", "test-master-key"), \ - patch("litellm.proxy.proxy_server.general_settings", {}), \ - patch("litellm.proxy.proxy_server.jwt_handler", MagicMock()), \ - patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), \ - patch.dict(os.environ, {"GOOGLE_CLIENT_ID": "test-google-id"}, clear=True), \ - patch("litellm.proxy.management_endpoints.ui_sso.GoogleSSOHandler.get_google_callback_response", return_value=mock_result): + + with patch( + "litellm.proxy.management_endpoints.ui_sso.cli_sso_callback" + ) as mock_cli_callback, patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.proxy_server.master_key", "test-master-key" + ), patch( + "litellm.proxy.proxy_server.general_settings", {} + ), patch( + "litellm.proxy.proxy_server.jwt_handler", MagicMock() + ), patch( + "litellm.proxy.proxy_server.user_api_key_cache", MagicMock() + ), patch.dict( + os.environ, {"GOOGLE_CLIENT_ID": "test-google-id"}, clear=True + ), patch( + "litellm.proxy.management_endpoints.ui_sso.GoogleSSOHandler.get_google_callback_response", + return_value=mock_result, + ): mock_cli_callback.return_value = MagicMock() - + # Act await auth_callback(request=mock_request, state=cli_state) - + # Assert - existing_key should be extracted from state parameter mock_cli_callback.assert_called_once_with( request=mock_request, key="sk-new-session-key-456", existing_key="sk-existing-cli-key-123", - result=mock_result + result=mock_result, ) def test_get_redirect_url_does_not_include_existing_key_in_url(self): @@ -1416,15 +1478,17 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock() mock_request.base_url = "https://test.litellm.ai/" - - with patch("litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"): + + with patch( + "litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai" + ): # Test with existing_key - should NOT be in URL redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( request=mock_request, sso_callback_route="sso/callback", - existing_key="sk-existing-123" + existing_key="sk-existing-123", ) - + # existing_key should NOT be in the URL assert "https://test.litellm.ai/sso/callback" == redirect_url assert "existing_key" not in redirect_url @@ -1436,44 +1500,181 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock() mock_request.base_url = "https://test.litellm.ai/" - - with patch("litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai"): + + with patch( + "litellm.proxy.utils.get_custom_url", return_value="https://test.litellm.ai" + ): # Test without existing_key redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( - request=mock_request, - sso_callback_route="sso/callback" + request=mock_request, sso_callback_route="sso/callback" ) - + assert "https://test.litellm.ai/sso/callback" == redirect_url @pytest.mark.asyncio async def test_cli_sso_callback_regenerate_vs_create_flow(self): """Test CLI SSO callback calls regenerate_key_fn when existing_key provided, generate_key_helper_fn when not""" from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback - + mock_request = MagicMock(spec=Request) - - with patch("litellm.proxy.management_endpoints.key_management_endpoints.regenerate_key_fn") as mock_regenerate, \ - patch("litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn") as mock_generate, \ - patch("litellm.proxy._types.UserAPIKeyAuth.get_litellm_cli_user_api_key_auth"), \ - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), \ - patch("litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success"): - + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.regenerate_key_fn" + ) as mock_regenerate, patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn" + ) as mock_generate, patch( + "litellm.proxy._types.UserAPIKeyAuth.get_litellm_cli_user_api_key_auth" + ), patch( + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), patch( + "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", + return_value="Success", + ): + # Test regeneration path - await cli_sso_callback(mock_request, key="sk-new-123", existing_key="sk-existing-456") + await cli_sso_callback( + mock_request, key="sk-new-123", existing_key="sk-existing-456" + ) mock_regenerate.assert_called_once() mock_generate.assert_not_called() - + # Reset mocks mock_regenerate.reset_mock() mock_generate.reset_mock() - + # Test creation path await cli_sso_callback(mock_request, key="sk-new-789", existing_key=None) mock_regenerate.assert_not_called() mock_generate.assert_called_once() +class TestGetAppRolesFromIdToken: + """Test the get_app_roles_from_id_token method""" + + def test_roles_picked_when_app_roles_not_exists(self): + """Test that 'roles' is picked when 'app_roles' doesn't exist""" + import jwt + + # Create a token with only 'roles' claim + token_payload = { + "sub": "user123", + "email": "test@example.com", + "roles": ["Admin", "User", "Developer"], + } + + # Create a mock JWT token + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload) as mock_jwt_decode: + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert + assert result == ["Admin", "User", "Developer"] + mock_jwt_decode.assert_called_once_with( + mock_token, options={"verify_signature": False} + ) + + def test_app_roles_picked_when_both_exist(self): + """Test that 'app_roles' takes precedence when both 'app_roles' and 'roles' exist""" + import jwt + + # Create a token with both 'app_roles' and 'roles' claims + token_payload = { + "sub": "user123", + "email": "test@example.com", + "app_roles": ["AppAdmin", "AppUser"], + "roles": ["RoleAdmin", "RoleUser"], + } + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert - app_roles should be picked, not roles + assert result == ["AppAdmin", "AppUser"] + + def test_roles_picked_when_app_roles_is_empty(self): + """Test that 'roles' is picked when 'app_roles' exists but is empty""" + import jwt + + # Create a token with empty 'app_roles' and populated 'roles' + token_payload = { + "sub": "user123", + "email": "test@example.com", + "app_roles": [], + "roles": ["Admin", "User"], + } + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert - roles should be picked since app_roles is empty + assert result == ["Admin", "User"] + + def test_empty_list_when_neither_exists(self): + """Test that empty list is returned when neither 'app_roles' nor 'roles' exist""" + import jwt + + # Create a token without roles claims + token_payload = {"sub": "user123", "email": "test@example.com"} + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert + assert result == [] + + def test_empty_list_when_no_token_provided(self): + """Test that empty list is returned when no token is provided""" + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(None) + + # Assert + assert result == [] + + def test_empty_list_when_roles_not_a_list(self): + """Test that empty list is returned when roles is not a list""" + import jwt + + # Create a token with non-list roles + token_payload = { + "sub": "user123", + "email": "test@example.com", + "roles": "Admin", # String instead of list + } + + mock_token = "mock.jwt.token" + + with patch("jwt.decode", return_value=token_payload): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert + assert result == [] + + def test_error_handling_on_jwt_decode_exception(self): + """Test that exceptions during JWT decode are handled gracefully""" + import jwt + + mock_token = "invalid.jwt.token" + + with patch("jwt.decode", side_effect=Exception("Invalid token")): + # Act + result = MicrosoftSSOHandler.get_app_roles_from_id_token(mock_token) + + # Assert - should return empty list on error + assert result == [] + + class TestProcessSSOJWTAccessToken: """Test the process_sso_jwt_access_token helper function""" @@ -1496,10 +1697,12 @@ class TestProcessSSOJWTAccessToken: "sub": "1234567890", "name": "John Doe", "iat": 1516239022, - "groups": ["team1", "team2", "team3"] + "groups": ["team1", "team2", "team3"], } - def test_process_sso_jwt_access_token_with_valid_token(self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload): + def test_process_sso_jwt_access_token_with_valid_token( + self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload + ): """Test processing a valid JWT access token with team extraction""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1513,7 +1716,7 @@ class TestProcessSSOJWTAccessToken: last_name="User", display_name="Test User", provider="generic", - team_ids=[] + team_ids=[], ) with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode: @@ -1521,7 +1724,7 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert @@ -1529,14 +1732,18 @@ class TestProcessSSOJWTAccessToken: mock_jwt_decode.assert_called_once_with( sample_jwt_token, options={"verify_signature": False} ) - + # Verify team IDs were extracted from JWT - mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(sample_jwt_payload) - + mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with( + sample_jwt_payload + ) + # Verify team IDs were set on the result object assert result.team_ids == ["team1", "team2", "team3"] - def test_process_sso_jwt_access_token_with_existing_team_ids(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_with_existing_team_ids( + self, mock_jwt_handler, sample_jwt_token + ): """Test that existing team IDs are not overwritten""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1551,7 +1758,7 @@ class TestProcessSSOJWTAccessToken: last_name="User", display_name="Test User", provider="generic", - team_ids=existing_team_ids + team_ids=existing_team_ids, ) with patch("jwt.decode") as mock_jwt_decode: @@ -1559,51 +1766,53 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert # JWT should still be decoded mock_jwt_decode.assert_called_once() - + # But team IDs should NOT be extracted since they already exist mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - + # Existing team IDs should remain unchanged assert result.team_ids == existing_team_ids - def test_process_sso_jwt_access_token_with_dict_result(self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload): + def test_process_sso_jwt_access_token_with_dict_result( + self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload + ): """Test processing with a dictionary result object""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, ) # Create a dictionary result without team_ids - result = { - "id": "test_user", - "email": "test@example.com", - "name": "Test User" - } + result = {"id": "test_user", "email": "test@example.com", "name": "Test User"} with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode: # Act process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert mock_jwt_decode.assert_called_once_with( sample_jwt_token, options={"verify_signature": False} ) - mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(sample_jwt_payload) - + mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with( + sample_jwt_payload + ) + # Verify team_ids was added to the dict as a key assert "team_ids" in result assert result["team_ids"] == ["team1", "team2", "team3"] - def test_process_sso_jwt_access_token_with_dict_existing_team_ids(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_with_dict_existing_team_ids( + self, mock_jwt_handler, sample_jwt_token + ): """Test that existing team IDs in dictionary are not overwritten""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1615,7 +1824,7 @@ class TestProcessSSOJWTAccessToken: "id": "test_user", "email": "test@example.com", "name": "Test User", - "team_ids": existing_team_ids + "team_ids": existing_team_ids, } with patch("jwt.decode") as mock_jwt_decode: @@ -1623,16 +1832,16 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert # JWT should still be decoded mock_jwt_decode.assert_called_once() - + # But team IDs should NOT be extracted since they already exist mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - + # Existing team IDs should remain unchanged assert result["team_ids"] == existing_team_ids @@ -1642,20 +1851,14 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token, ) - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) # Test with None access token with patch("jwt.decode") as mock_jwt_decode: process_sso_jwt_access_token( - access_token_str=None, - sso_jwt_handler=mock_jwt_handler, - result=result + access_token_str=None, sso_jwt_handler=mock_jwt_handler, result=result ) - + # Assert nothing was processed mock_jwt_decode.assert_not_called() mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() @@ -1664,11 +1867,9 @@ class TestProcessSSOJWTAccessToken: # Test with empty string access token with patch("jwt.decode") as mock_jwt_decode: process_sso_jwt_access_token( - access_token_str="", - sso_jwt_handler=mock_jwt_handler, - result=result + access_token_str="", sso_jwt_handler=mock_jwt_handler, result=result ) - + # Assert nothing was processed mock_jwt_decode.assert_not_called() mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() @@ -1680,25 +1881,21 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token, ) - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) with patch("jwt.decode") as mock_jwt_decode: # Act process_sso_jwt_access_token( - access_token_str=sample_jwt_token, - sso_jwt_handler=None, - result=result + access_token_str=sample_jwt_token, sso_jwt_handler=None, result=result ) # Assert nothing was processed mock_jwt_decode.assert_not_called() assert result.team_ids == [] - def test_process_sso_jwt_access_token_no_result(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_no_result( + self, mock_jwt_handler, sample_jwt_token + ): """Test that nothing happens when result is None""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1709,32 +1906,32 @@ class TestProcessSSOJWTAccessToken: process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=None + result=None, ) # Assert nothing was processed mock_jwt_decode.assert_not_called() mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - def test_process_sso_jwt_access_token_jwt_decode_exception(self, mock_jwt_handler, sample_jwt_token): + def test_process_sso_jwt_access_token_jwt_decode_exception( + self, mock_jwt_handler, sample_jwt_token + ): """Test that JWT decode exceptions are not caught (should propagate up)""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, ) - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) - with patch("jwt.decode", side_effect=Exception("JWT decode error")) as mock_jwt_decode: + with patch( + "jwt.decode", side_effect=Exception("JWT decode error") + ) as mock_jwt_decode: # Act & Assert with pytest.raises(Exception, match="JWT decode error"): process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Verify JWT decode was attempted @@ -1742,7 +1939,9 @@ class TestProcessSSOJWTAccessToken: # But team extraction should not have been called mock_jwt_handler.get_team_ids_from_jwt.assert_not_called() - def test_process_sso_jwt_access_token_empty_team_ids_from_jwt(self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload): + def test_process_sso_jwt_access_token_empty_team_ids_from_jwt( + self, mock_jwt_handler, sample_jwt_token, sample_jwt_payload + ): """Test processing when JWT handler returns empty team IDs""" from litellm.proxy.management_endpoints.ui_sso import ( process_sso_jwt_access_token, @@ -1751,24 +1950,126 @@ class TestProcessSSOJWTAccessToken: # Configure mock to return empty team IDs mock_jwt_handler.get_team_ids_from_jwt.return_value = [] - result = CustomOpenID( - id="test_user", - email="test@example.com", - team_ids=[] - ) + result = CustomOpenID(id="test_user", email="test@example.com", team_ids=[]) with patch("jwt.decode", return_value=sample_jwt_payload) as mock_jwt_decode: # Act process_sso_jwt_access_token( access_token_str=sample_jwt_token, sso_jwt_handler=mock_jwt_handler, - result=result + result=result, ) # Assert mock_jwt_decode.assert_called_once() - mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with(sample_jwt_payload) - + mock_jwt_handler.get_team_ids_from_jwt.assert_called_once_with( + sample_jwt_payload + ) + # Even empty team IDs should be set assert result.team_ids == [] + +class TestPKCEFunctionality: + """Test PKCE (Proof Key for Code Exchange) functionality""" + + def test_generate_pkce_params(self): + """ + Test that generate_pkce_params generates valid PKCE parameters + """ + import base64 + import hashlib + + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Act + code_verifier, code_challenge = SSOAuthenticationHandler.generate_pkce_params() + + # Assert + assert len(code_verifier) == 43 + assert isinstance(code_verifier, str) + + # Verify code_challenge is correctly generated from code_verifier + expected_challenge_bytes = hashlib.sha256(code_verifier.encode('utf-8')).digest() + expected_challenge = base64.urlsafe_b64encode(expected_challenge_bytes).decode('utf-8').rstrip('=') + assert code_challenge == expected_challenge + + # Verify both are base64url encoded (no padding) + assert '=' not in code_verifier + assert '=' not in code_challenge + + @pytest.mark.asyncio + async def test_prepare_token_exchange_parameters_with_pkce(self): + """ + Test prepare_token_exchange_parameters retrieves PKCE code_verifier from cache + """ + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Mock request with state parameter + mock_request = MagicMock(spec=Request) + test_state = "test_oauth_state_123" + mock_request.query_params = {"state": test_state} + + # Mock cache + mock_cache = MagicMock() + test_code_verifier = "test_code_verifier_abc123xyz" + mock_cache.get_cache.return_value = test_code_verifier + + with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + # Act + token_params = SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=mock_request, + generic_include_client_id=False + ) + + # Assert + assert token_params["include_client_id"] is False + assert token_params["code_verifier"] == test_code_verifier + + # Verify cache was accessed and deleted + mock_cache.get_cache.assert_called_once_with(key=f"pkce_verifier:{test_state}") + mock_cache.delete_cache.assert_called_once_with(key=f"pkce_verifier:{test_state}") + + @pytest.mark.asyncio + async def test_get_generic_sso_redirect_response_with_pkce(self): + """ + Test get_generic_sso_redirect_response with PKCE enabled stores verifier and adds challenge to URL + """ + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Mock SSO provider + mock_sso = MagicMock() + mock_redirect_response = MagicMock() + original_location = "https://auth.example.com/authorize?state=test456&client_id=abc" + mock_redirect_response.headers = {"location": original_location} + mock_sso.get_login_redirect = AsyncMock(return_value=mock_redirect_response) + mock_sso.__enter__ = MagicMock(return_value=mock_sso) + mock_sso.__exit__ = MagicMock(return_value=False) + + test_state = "test456" + mock_cache = MagicMock() + + with patch.dict(os.environ, {"GENERIC_CLIENT_USE_PKCE": "true"}): + with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + # Act + result = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=mock_sso, + state=test_state, + generic_authorization_endpoint="https://auth.example.com/authorize" + ) + + # Assert + # Verify cache was called to store code_verifier + mock_cache.set_cache.assert_called_once() + cache_call = mock_cache.set_cache.call_args + assert cache_call.kwargs["key"] == f"pkce_verifier:{test_state}" + assert cache_call.kwargs["ttl"] == 600 + assert len(cache_call.kwargs["value"]) == 43 + + # Verify PKCE parameters were added to the redirect URL + assert result is not None + updated_location = str(result.headers["location"]) + assert "code_challenge=" in updated_location + assert "code_challenge_method=S256" in updated_location + assert f"state={test_state}" in updated_location + diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 239f83b21ad..6eeca946190 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -18,12 +18,12 @@ import litellm from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( BaseOpenAIPassThroughHandler, RouteChecks, + bedrock_llm_proxy_route, create_pass_through_route, llm_passthrough_factory_proxy_route, - vllm_proxy_route, vertex_discovery_proxy_route, vertex_proxy_route, - bedrock_llm_proxy_route, + vllm_proxy_route, ) from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials @@ -996,6 +996,64 @@ class TestBedrockLLMProxyRoute: assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0" assert result == "success" + @pytest.mark.asyncio + async def test_bedrock_error_handling_returns_actual_error(self): + """ + Test that when Bedrock API returns an error, it is properly propagated to the user + instead of being returned as a generic "Internal Server Error". + """ + from fastapi import HTTPException + + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + handle_bedrock_passthrough_router_model, + ) + + mock_request = Mock() + mock_request.method = "POST" + mock_request.headers = {"content-type": "application/json"} + mock_request.query_params = {} + + mock_request_body = { + "messages": [ + { + "role": "user", + "content": [{"textaaa": "Hello"}] + } + ] + } + + bedrock_error_message = '{"message":"ContentBlock object at messages.0.content.0 must set one of the following keys: text, image, toolUse, toolResult, document, video."}' + + # Create a mock httpx.Response for the error + mock_error_response = Mock(spec=httpx.Response) + mock_error_response.status_code = 400 + mock_error_response.aread = AsyncMock(return_value=bedrock_error_message.encode('utf-8')) + + # Create the HTTPStatusError + mock_http_error = httpx.HTTPStatusError( + message="Bad Request", + request=Mock(spec=httpx.Request), + response=mock_error_response, + ) + + mock_llm_router = Mock() + mock_llm_router.allm_passthrough_route = AsyncMock(side_effect=mock_http_error) + + endpoint = "model/test-model/converse" + model = "test-model" + + with pytest.raises(HTTPException) as exc_info: + await handle_bedrock_passthrough_router_model( + model=model, + endpoint=endpoint, + request=mock_request, + request_body=mock_request_body, + llm_router=mock_llm_router, + ) + + assert exc_info.value.status_code == 400 + assert "ContentBlock object at messages.0.content.0 must set one of the following keys" in str(exc_info.value.detail) + class TestLLMPassthroughFactoryProxyRoute: @pytest.mark.asyncio diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 877a9d069e7..3e825e935d1 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1548,3 +1548,74 @@ def test_get_deployment_model_info_base_model_merge_priority(): assert result["key"] == "gpt-4" print("✓ Base model merge priority test passed!") + + +def test_add_deployment_model_to_endpoint_for_llm_passthrough_route(): + """ + Test that _add_deployment_model_to_endpoint_for_llm_passthrough_route correctly strips bedrock provider prefix + """ + router = litellm.Router( + model_list=[ + { + "model_name": "special-bedrock-model", + "litellm_params": { + "model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + }, + } + ], + ) + + # Test Case 1: Bedrock model with provider prefix - should strip "bedrock/" prefix + kwargs = { + "endpoint": "/model/special-bedrock-model/invoke", + "custom_llm_provider": "bedrock", + } + result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( + kwargs=kwargs, + model="special-bedrock-model", + model_name="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + ) + assert ( + result["endpoint"] == "/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke" + ), f"Expected '/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke', got '{result['endpoint']}'" + + # Test Case 2: Bedrock invoke-with-response-stream endpoint + kwargs = { + "endpoint": "/model/special-bedrock-model/invoke-with-response-stream", + "custom_llm_provider": "bedrock", + } + result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( + kwargs=kwargs, + model="special-bedrock-model", + model_name="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + ) + assert ( + result["endpoint"] == "/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke-with-response-stream" + ), f"Expected streaming endpoint with stripped prefix, got '{result['endpoint']}'" + + # Test Case 3: Bedrock converse endpoint + kwargs = { + "endpoint": "/model/bedrock-model/converse", + "custom_llm_provider": "bedrock", + } + result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( + kwargs=kwargs, + model="bedrock-model", + model_name="bedrock/us.meta.llama3-8b-instruct-v1:0", + ) + assert ( + result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/converse" + ), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/converse', got '{result['endpoint']}'" + + # Test Case 4: Bedrock provider prefix auto-detected from model_name + kwargs = { + "endpoint": "/model/router-model/invoke", + } + result = router._add_deployment_model_to_endpoint_for_llm_passthrough_route( + kwargs=kwargs, + model="router-model", + model_name="bedrock/us.meta.llama3-8b-instruct-v1:0", + ) + assert ( + result["endpoint"] == "/model/us.meta.llama3-8b-instruct-v1:0/invoke" + ), f"Expected '/model/us.meta.llama3-8b-instruct-v1:0/invoke', got '{result['endpoint']}'"