mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
LiteLLM SDK <-> Proxy improvement (don't transform message client-side) + Bedrock - handle qs:.. in base64 file data + Tag Management - support adding public model names (#11908)
* fix(factory.py): handle qs:.. in mime type Fixes https://github.com/BerriAI/litellm/issues/11839 * feat(litellm_proxy/): don't transform messages client-side leave litellm proxy messages untouched - allow proxy to handle transformation prevents double transformation * feat(tag_management_endpoints.py): support adding models to tag by adding model_name Closes https://github.com/BerriAI/litellm/issues/11884 * test(test_tag_management_endpoints.py): add unit tests for adding new model by public model name * test: update test
This commit is contained in:
parent
0e6f1c089d
commit
308e82d885
8 changed files with 265 additions and 33 deletions
|
|
@ -1053,10 +1053,10 @@ def convert_to_gemini_tool_call_invoke(
|
|||
if tool_calls is not None:
|
||||
for tool in tool_calls:
|
||||
if "function" in tool:
|
||||
gemini_function_call: Optional[
|
||||
VertexFunctionCall
|
||||
] = _gemini_tool_call_invoke_helper(
|
||||
function_call_params=tool["function"]
|
||||
gemini_function_call: Optional[VertexFunctionCall] = (
|
||||
_gemini_tool_call_invoke_helper(
|
||||
function_call_params=tool["function"]
|
||||
)
|
||||
)
|
||||
if gemini_function_call is not None:
|
||||
_parts_list.append(
|
||||
|
|
@ -1573,9 +1573,9 @@ def anthropic_messages_pt( # noqa: PLR0915
|
|||
)
|
||||
|
||||
if "cache_control" in _content_element:
|
||||
_anthropic_content_element[
|
||||
"cache_control"
|
||||
] = _content_element["cache_control"]
|
||||
_anthropic_content_element["cache_control"] = (
|
||||
_content_element["cache_control"]
|
||||
)
|
||||
user_content.append(_anthropic_content_element)
|
||||
elif m.get("type", "") == "text":
|
||||
m = cast(ChatCompletionTextObject, m)
|
||||
|
|
@ -1613,9 +1613,9 @@ def anthropic_messages_pt( # noqa: PLR0915
|
|||
)
|
||||
|
||||
if "cache_control" in _content_element:
|
||||
_anthropic_content_text_element[
|
||||
"cache_control"
|
||||
] = _content_element["cache_control"]
|
||||
_anthropic_content_text_element["cache_control"] = (
|
||||
_content_element["cache_control"]
|
||||
)
|
||||
|
||||
user_content.append(_anthropic_content_text_element)
|
||||
|
||||
|
|
@ -2433,8 +2433,10 @@ class BedrockImageProcessor:
|
|||
|
||||
# Extract MIME type using regular expression
|
||||
mime_type_match = re.match(r"data:(.*?);base64", image_metadata)
|
||||
|
||||
if mime_type_match:
|
||||
mime_type = mime_type_match.group(1)
|
||||
mime_type = mime_type.split(";")[0]
|
||||
image_format = mime_type.split("/")[1]
|
||||
else:
|
||||
mime_type = "image/jpeg"
|
||||
|
|
@ -2458,6 +2460,7 @@ class BedrockImageProcessor:
|
|||
|
||||
document_types = ["application", "text"]
|
||||
is_document = any(mime_type.startswith(doc_type) for doc_type in document_types)
|
||||
|
||||
supported_image_and_video_formats: List[str] = (
|
||||
supported_video_formats + supported_image_formats
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,13 +2,16 @@
|
|||
Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
|
||||
from litellm.secret_managers.main import get_secret_bool, get_secret_str
|
||||
from litellm.types.router import LiteLLM_Params
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
class LiteLLMProxyChatConfig(OpenAIGPTConfig):
|
||||
def get_supported_openai_params(self, model: str) -> List:
|
||||
|
|
@ -113,3 +116,33 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig):
|
|||
)
|
||||
|
||||
return model, custom_llm_provider, api_key, api_base
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List["AllMessageValues"],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
# don't transform the request
|
||||
return {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
async def async_transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List["AllMessageValues"],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
# don't transform the request
|
||||
return {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,12 @@ model_list:
|
|||
litellm_params:
|
||||
model: codex-mini-latest
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
- model_name: bedrock/*
|
||||
litellm_params:
|
||||
model: bedrock/*
|
||||
- model_name: eu.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
litellm_params:
|
||||
model: eu.anthropic.claude-3-5-sonnet-20240620-v1:0
|
||||
- model_name: "gpt-4o-mini-openai"
|
||||
litellm_params:
|
||||
model: gpt-4o-mini
|
||||
|
|
|
|||
|
|
@ -1,18 +1,19 @@
|
|||
"""
|
||||
TAG MANAGEMENT
|
||||
|
||||
All /tag management endpoints
|
||||
All /tag management endpoints
|
||||
|
||||
/tag/new
|
||||
/tag/new
|
||||
/tag/info
|
||||
/tag/update
|
||||
/tag/delete
|
||||
/tag/list
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
from typing import Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
|
|
@ -33,6 +34,10 @@ from litellm.types.tag_management import (
|
|||
TagUpdateRequest,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import Router
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -111,6 +116,33 @@ async def _save_tags_config(prisma_client, tags_config: Dict[str, TagConfig]):
|
|||
)
|
||||
|
||||
|
||||
async def get_deployments_by_model(
|
||||
model: str, llm_router: "Router"
|
||||
) -> List["Deployment"]:
|
||||
"""
|
||||
Get all deployments by model
|
||||
"""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
# Check if model id
|
||||
deployment = llm_router.get_deployment(model_id=model)
|
||||
if deployment is not None:
|
||||
return [deployment]
|
||||
|
||||
# Check if model name
|
||||
deployments = llm_router.get_model_list(model_name=model)
|
||||
if deployments is None:
|
||||
return []
|
||||
return [
|
||||
Deployment(
|
||||
model_name=deployment["model_name"],
|
||||
litellm_params=LiteLLM_Params(**deployment["litellm_params"]), # type: ignore
|
||||
model_info=ModelInfo(**deployment.get("model_info") or {}),
|
||||
)
|
||||
for deployment in deployments
|
||||
]
|
||||
|
||||
|
||||
@router.post(
|
||||
"/tag/new",
|
||||
tags=["tag management"],
|
||||
|
|
@ -126,12 +158,19 @@ async def new_tag(
|
|||
Parameters:
|
||||
- name: str - The name of the tag
|
||||
- description: Optional[str] - Description of what this tag represents
|
||||
- models: List[str] - List of LLM models allowed for this tag
|
||||
- models: List[str] - List of either 'model_id' or 'model_name' allowed for this tag
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.proxy_server import llm_router, prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.no_llm_router.value
|
||||
)
|
||||
try:
|
||||
# Get existing tags config
|
||||
tags_config = await _get_tags_config(prisma_client)
|
||||
|
|
@ -160,11 +199,19 @@ async def new_tag(
|
|||
|
||||
# Update models with new tag
|
||||
if tag.models:
|
||||
for model_id in tag.models:
|
||||
await _add_tag_to_deployment(
|
||||
model_id=model_id,
|
||||
tag=tag.name,
|
||||
tasks = []
|
||||
for model in tag.models:
|
||||
deployments = await get_deployments_by_model(model, llm_router)
|
||||
tasks.extend(
|
||||
[
|
||||
_add_tag_to_deployment(
|
||||
deployment=deployment,
|
||||
tag=tag.name,
|
||||
)
|
||||
for deployment in deployments
|
||||
]
|
||||
)
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Get model names for response
|
||||
model_info = await _get_model_names(prisma_client, tag.models or [])
|
||||
|
|
@ -179,27 +226,26 @@ async def new_tag(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
async def _add_tag_to_deployment(model_id: str, tag: str):
|
||||
async def _add_tag_to_deployment(deployment: "Deployment", tag: str):
|
||||
"""Helper function to add tag to deployment"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
deployment = await prisma_client.db.litellm_proxymodeltable.find_unique(
|
||||
where={"model_id": model_id}
|
||||
)
|
||||
if deployment is None:
|
||||
raise HTTPException(status_code=404, detail=f"Deployment {model_id} not found")
|
||||
|
||||
litellm_params = deployment.litellm_params
|
||||
if "tags" not in litellm_params:
|
||||
litellm_params["tags"] = []
|
||||
litellm_params["tags"].append(tag)
|
||||
await prisma_client.db.litellm_proxymodeltable.update(
|
||||
where={"model_id": model_id},
|
||||
data={"litellm_params": safe_dumps(litellm_params)},
|
||||
)
|
||||
|
||||
try:
|
||||
await prisma_client.db.litellm_proxymodeltable.update(
|
||||
where={"model_id": deployment.model_info.id},
|
||||
data={"litellm_params": safe_dumps(litellm_params)},
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error adding tag to deployment: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
|
||||
from litellm.types.integrations.langfuse_otel import LangfuseOtelConfig
|
||||
|
||||
|
|
@ -308,3 +308,20 @@ def test_vertex_ai_transform_empty_function_call_arguments():
|
|||
assert result["args"] == {
|
||||
"type": "object",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_process_image_async_factory():
|
||||
"""
|
||||
Test that the _process_image_async_factory method handles image input correctly
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
BedrockImageProcessor,
|
||||
)
|
||||
|
||||
image_url = "data:application/pdf; qs=0.001;base64,JVBERi0xLjQKJcOkw7zDtsOfCjIgMCBvYmoKPDwvTGVuZ3RoIDMgMCBSL0ZpbHRlci9GbGF0ZURlY29kZT4"
|
||||
|
||||
content_block = await BedrockImageProcessor.process_image_async(
|
||||
image_url=image_url, format=None
|
||||
)
|
||||
print(f"content_block: {content_block}")
|
||||
|
|
|
|||
|
|
@ -0,0 +1,32 @@
|
|||
from typing import Optional
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.litellm_proxy.chat.transformation import LiteLLMProxyChatConfig
|
||||
|
||||
|
||||
def test_litellm_proxy_chat_transformation():
|
||||
"""
|
||||
Assert messages are not transformed when calling litellm proxy
|
||||
"""
|
||||
config = LiteLLMProxyChatConfig()
|
||||
file_content = [
|
||||
{"type": "text", "text": "What is this document about?"},
|
||||
{
|
||||
"type": "file",
|
||||
"file": {
|
||||
"file_id": "https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf",
|
||||
"format": "application/pdf",
|
||||
},
|
||||
},
|
||||
]
|
||||
messages = [{"role": "user", "content": file_content}]
|
||||
assert config.transform_request(
|
||||
model="model",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
) == {"model": "model", "messages": messages}
|
||||
|
|
@ -26,6 +26,8 @@ async def test_create_and_get_tag():
|
|||
"""
|
||||
# Mock the prisma client and _get_tags_config and _save_tags_config
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
|
||||
"litellm.proxy.proxy_server.llm_router"
|
||||
) as mock_router, patch(
|
||||
"litellm.proxy.management_endpoints.tag_management_endpoints._get_tags_config"
|
||||
) as mock_get_tags, patch(
|
||||
"litellm.proxy.management_endpoints.tag_management_endpoints._save_tags_config"
|
||||
|
|
@ -50,6 +52,7 @@ async def test_create_and_get_tag():
|
|||
|
||||
# Test tag creation
|
||||
response = client.post("/tag/new", json=tag_data, headers=headers)
|
||||
print(f"response: {response.text}")
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert result["message"] == "Tag test-tag created successfully"
|
||||
|
|
@ -158,3 +161,93 @@ async def test_delete_tag():
|
|||
|
||||
# Verify _save_tags_config was called without the deleted tag
|
||||
mock_save_tags.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_deployments_by_model_id():
|
||||
"""
|
||||
Test get_deployments_by_model when model is found by model_id
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.proxy.management_endpoints.tag_management_endpoints import (
|
||||
get_deployments_by_model,
|
||||
)
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
|
||||
# Create a mock router
|
||||
mock_router = Mock()
|
||||
|
||||
# Setup mock to return deployment by model_id
|
||||
mock_deployment = Deployment(
|
||||
model_name="gpt-3.5-turbo",
|
||||
litellm_params=LiteLLM_Params(model="gpt-3.5-turbo"),
|
||||
model_info=ModelInfo(),
|
||||
)
|
||||
mock_router.get_deployment.return_value = mock_deployment
|
||||
|
||||
result = await get_deployments_by_model("model-123", mock_router)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0] == mock_deployment
|
||||
mock_router.get_deployment.assert_called_once_with(model_id="model-123")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_deployments_by_model_name():
|
||||
"""
|
||||
Test get_deployments_by_model when model is found by model_name
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.proxy.management_endpoints.tag_management_endpoints import (
|
||||
get_deployments_by_model,
|
||||
)
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
# Create a mock router
|
||||
mock_router = Mock()
|
||||
|
||||
# Setup mock to not find by model_id but find by model_name
|
||||
mock_router.get_deployment.return_value = None
|
||||
mock_router.get_model_list.return_value = [
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo", "api_key": "test-key"},
|
||||
"model_info": {"id": "model-1", "description": "Test model"},
|
||||
}
|
||||
]
|
||||
|
||||
result = await get_deployments_by_model("gpt-3.5-turbo", mock_router)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].model_name == "gpt-3.5-turbo"
|
||||
assert isinstance(result[0], Deployment)
|
||||
mock_router.get_deployment.assert_called_once_with(model_id="gpt-3.5-turbo")
|
||||
mock_router.get_model_list.assert_called_once_with(model_name="gpt-3.5-turbo")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_deployments_by_model_not_found():
|
||||
"""
|
||||
Test get_deployments_by_model when model is not found
|
||||
"""
|
||||
from unittest.mock import Mock
|
||||
|
||||
from litellm.proxy.management_endpoints.tag_management_endpoints import (
|
||||
get_deployments_by_model,
|
||||
)
|
||||
|
||||
# Create a mock router
|
||||
mock_router = Mock()
|
||||
|
||||
# Setup mock to not find model by either method
|
||||
mock_router.get_deployment.return_value = None
|
||||
mock_router.get_model_list.return_value = None
|
||||
|
||||
result = await get_deployments_by_model("nonexistent-model", mock_router)
|
||||
|
||||
assert len(result) == 0
|
||||
assert result == []
|
||||
mock_router.get_deployment.assert_called_once_with(model_id="nonexistent-model")
|
||||
mock_router.get_model_list.assert_called_once_with(model_name="nonexistent-model")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue