mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Direct Anthropic http image inlining, Vertex AI Anthropic forced base64 conversion, Ollama completion image download and the watsonx GPT-OSS Hugging Face chat template lookup all ran synchronous HTTP inside the async request path. Each provider config now transforms through async_inline_remote_media on the async path, the Anthropic handler awaits the config's async_transform_request before dispatch and in the Rust fallback, and watsonx text exposes async_transform_request and always awaits ahf_chat_template Resolves LIT-7028 Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
577 lines
21 KiB
Python
577 lines
21 KiB
Python
import json
|
|
|
|
from typing import Optional
|
|
from unittest.mock import Mock, patch
|
|
|
|
import pytest
|
|
|
|
import litellm
|
|
from litellm import completion
|
|
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
|
|
|
|
|
@pytest.fixture
|
|
def watsonx_chat_completion_call():
|
|
def _call(
|
|
model="watsonx/my-test-model",
|
|
messages=None,
|
|
api_key="test_api_key",
|
|
space_id: Optional[str] = None,
|
|
headers=None,
|
|
client=None,
|
|
patch_token_call=True,
|
|
):
|
|
if messages is None:
|
|
messages = [{"role": "user", "content": "Hello, how are you?"}]
|
|
if client is None:
|
|
client = HTTPHandler()
|
|
|
|
if patch_token_call:
|
|
mock_response = Mock()
|
|
mock_response.json.return_value = {
|
|
"access_token": "mock_access_token",
|
|
"expires_in": 3600,
|
|
}
|
|
mock_response.raise_for_status = Mock() # No-op to simulate no exception
|
|
|
|
with (
|
|
patch.object(client, "post") as mock_post,
|
|
patch.object(
|
|
litellm.module_level_client, "post", return_value=mock_response
|
|
) as mock_get,
|
|
):
|
|
try:
|
|
completion(
|
|
model=model,
|
|
messages=messages,
|
|
api_key=api_key,
|
|
headers=headers or {},
|
|
client=client,
|
|
space_id=space_id,
|
|
)
|
|
except Exception as e:
|
|
print(e)
|
|
|
|
return mock_post, mock_get
|
|
else:
|
|
with patch.object(client, "post") as mock_post:
|
|
try:
|
|
completion(
|
|
model=model,
|
|
messages=messages,
|
|
api_key=api_key,
|
|
headers=headers or {},
|
|
client=client,
|
|
space_id=space_id,
|
|
)
|
|
except Exception as e:
|
|
print(e)
|
|
return mock_post, None
|
|
|
|
return _call
|
|
|
|
|
|
def test_watsonx_deployment_model_id_not_in_payload(
|
|
monkeypatch, watsonx_chat_completion_call
|
|
):
|
|
"""Test that deployment models do not include 'model_id' in the request payload"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
model = "watsonx/deployment/test-deployment-id"
|
|
messages = [{"role": "user", "content": "Test message"}]
|
|
|
|
mock_post, _ = watsonx_chat_completion_call(model=model, messages=messages)
|
|
|
|
assert mock_post.call_count == 1
|
|
json_data = json.loads(mock_post.call_args.kwargs["data"])
|
|
# Ensure model_id is not in the payload for deployment models
|
|
assert "model_id" not in json_data or json_data["model_id"] is None
|
|
# Ensure project_id is also not in the payload for deployment models
|
|
assert "project_id" not in json_data or json_data["project_id"] is None
|
|
|
|
|
|
def test_watsonx_regular_model_includes_model_id(
|
|
monkeypatch, watsonx_chat_completion_call
|
|
):
|
|
"""Test that regular models include 'model_id' in the request payload"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
model = "watsonx/regular-model"
|
|
messages = [{"role": "user", "content": "Test message"}]
|
|
|
|
mock_post, _ = watsonx_chat_completion_call(model=model, messages=messages)
|
|
|
|
assert mock_post.call_count == 1
|
|
json_data = json.loads(mock_post.call_args.kwargs["data"])
|
|
# Ensure model_id is included in the payload for regular models
|
|
assert "model_id" in json_data
|
|
assert json_data["model_id"] == "regular-model" # Provider prefix is stripped
|
|
# Ensure project_id is also included for regular models
|
|
assert "project_id" in json_data
|
|
|
|
|
|
@pytest.fixture
|
|
def watsonx_completion_call():
|
|
def _call(
|
|
model="watsonx_text/my-test-model",
|
|
prompt="Hello, how are you?",
|
|
api_key="test_api_key",
|
|
space_id: Optional[str] = None,
|
|
headers=None,
|
|
client=None,
|
|
patch_token_call=True,
|
|
):
|
|
if client is None:
|
|
client = HTTPHandler()
|
|
|
|
if patch_token_call:
|
|
mock_response = Mock()
|
|
mock_response.json.return_value = {
|
|
"access_token": "mock_access_token",
|
|
"expires_in": 3600,
|
|
}
|
|
mock_response.raise_for_status = Mock()
|
|
|
|
with (
|
|
patch.object(client, "post") as mock_post,
|
|
patch.object(
|
|
litellm.module_level_client, "post", return_value=mock_response
|
|
) as mock_get,
|
|
):
|
|
try:
|
|
litellm.text_completion(
|
|
model=model,
|
|
prompt=prompt,
|
|
api_key=api_key,
|
|
headers=headers or {},
|
|
client=client,
|
|
space_id=space_id,
|
|
)
|
|
except Exception as e:
|
|
print(e)
|
|
|
|
return mock_post, mock_get
|
|
else:
|
|
with patch.object(client, "post") as mock_post:
|
|
try:
|
|
litellm.text_completion(
|
|
model=model,
|
|
prompt=prompt,
|
|
api_key=api_key,
|
|
headers=headers or {},
|
|
client=client,
|
|
space_id=space_id,
|
|
)
|
|
except Exception as e:
|
|
print(e)
|
|
return mock_post, None
|
|
|
|
return _call
|
|
|
|
|
|
def test_watsonx_completion_deployment_model_id_not_in_payload(
|
|
monkeypatch, watsonx_completion_call
|
|
):
|
|
"""Test that deployment models do not include 'model_id' in completion request payload"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
model = "watsonx_text/deployment/test-deployment-id"
|
|
prompt = "Test prompt"
|
|
|
|
mock_post, _ = watsonx_completion_call(model=model, prompt=prompt)
|
|
|
|
assert mock_post.call_count == 1
|
|
json_data = json.loads(mock_post.call_args.kwargs["data"])
|
|
# Ensure model_id is not in the payload for deployment models
|
|
assert "model_id" not in json_data
|
|
# Ensure project_id is also not in the payload for deployment models
|
|
assert "project_id" not in json_data
|
|
|
|
|
|
def test_watsonx_completion_regular_model_includes_model_id(
|
|
monkeypatch, watsonx_completion_call
|
|
):
|
|
"""Test that regular models include 'model_id' in completion request payload"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
model = "watsonx_text/regular-model"
|
|
prompt = "Test prompt"
|
|
|
|
mock_post, _ = watsonx_completion_call(model=model, prompt=prompt)
|
|
|
|
assert mock_post.call_count == 1
|
|
json_data = json.loads(mock_post.call_args.kwargs["data"])
|
|
# Ensure model_id is included in the payload for regular models
|
|
assert "model_id" in json_data
|
|
assert json_data["model_id"] == "regular-model" # Provider prefix is stripped
|
|
# Ensure project_id is also included for regular models
|
|
assert "project_id" in json_data
|
|
|
|
|
|
def test_watsonx_gpt_oss_prompt_transformation(monkeypatch):
|
|
"""
|
|
Test that gpt-oss-120b model transforms messages to proper format instead of simple concatenation.
|
|
|
|
This test calls litellm.completion (sync) and verifies what gets sent in the final POST request body.
|
|
Input messages should be transformed using the HuggingFace chat template from openai/gpt-oss-120b,
|
|
not just concatenated as "You are chatgpt Hi there".
|
|
"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
|
|
# Test with gpt-oss model using watsonx_text provider (text generation endpoint)
|
|
model = "watsonx_text/openai/gpt-oss-120b"
|
|
|
|
# Input messages
|
|
messages = [
|
|
{"role": "system", "content": "You are chatgpt"},
|
|
{"role": "user", "content": "Hi there"},
|
|
]
|
|
|
|
client = HTTPHandler()
|
|
|
|
# Mock HuggingFace template fetch to make test deterministic and avoid network flakiness.
|
|
# The test verifies that prompt transformation occurs (not simple concatenation), not the exact
|
|
# HuggingFace template format. Using a mock template that produces the correct format is sufficient.
|
|
#
|
|
# Mock template that produces gpt-oss-120b-like format.
|
|
# Note: This is a simplified version of the actual template. The real template is more complex
|
|
# (adds metadata, handles tools, thinking messages, etc.), but this captures the key aspects:
|
|
# - Converts system role to developer (matching real template behavior)
|
|
# - Uses the same tag structure (<|start|>, <|message|>, <|end|>)
|
|
# - Preserves message content
|
|
mock_tokenizer_config = {
|
|
"status": "success",
|
|
"tokenizer": {
|
|
"chat_template": "{% for message in messages %}{% if message['role'] == 'system' %}<|start|>developer<|message|>{% else %}<|start|>{{ message['role'] }}<|message|>{% endif %}{{ message['content'] }}<|end|>{% endfor %}",
|
|
"bos_token": None,
|
|
"eos_token": None,
|
|
},
|
|
}
|
|
|
|
# Isolate known_tokenizer_config so parallel tests don't interfere.
|
|
# monkeypatch.setitem restores the original value on teardown.
|
|
hf_model = "openai/gpt-oss-120b"
|
|
monkeypatch.setitem(litellm.known_tokenizer_config, hf_model, mock_tokenizer_config)
|
|
|
|
# Mock IAM token generation to avoid real HTTP calls.
|
|
mock_token_response = Mock()
|
|
mock_token_response.json.return_value = {
|
|
"access_token": "mock_access_token",
|
|
"expires_in": 3600,
|
|
}
|
|
mock_token_response.raise_for_status = Mock()
|
|
|
|
with (
|
|
patch.object(client, "post") as mock_post,
|
|
patch.object(
|
|
litellm.module_level_client, "post", return_value=mock_token_response
|
|
),
|
|
):
|
|
try:
|
|
completion(
|
|
model=model,
|
|
messages=messages,
|
|
api_key="test_api_key",
|
|
client=client,
|
|
)
|
|
except Exception as e:
|
|
print(f"Caught expected exception: {e}")
|
|
|
|
# Verify the POST was called
|
|
assert (
|
|
mock_post.call_count == 1
|
|
), f"POST should have been called exactly once, got {mock_post.call_count}"
|
|
|
|
# Get the request body
|
|
call_args = mock_post.call_args
|
|
assert "data" in call_args.kwargs, "call_args.kwargs should contain 'data'"
|
|
json_data = json.loads(call_args.kwargs["data"])
|
|
|
|
# Verify the transformed input is in the request
|
|
assert "input" in json_data, "Request should have 'input' field"
|
|
transformed_prompt = json_data["input"]
|
|
|
|
# Verify it's NOT simple concatenation
|
|
simple_concat = "You are chatgpt Hi there"
|
|
assert transformed_prompt != simple_concat, (
|
|
f"Prompt should not be simple concatenation.\n"
|
|
f"Expected: Chat template with <|start|> tags\n"
|
|
f"Got: {transformed_prompt}"
|
|
)
|
|
|
|
# Verify it contains proper chat template formatting
|
|
assert "<|start|>" in transformed_prompt, "Prompt should contain <|start|> tag"
|
|
assert "<|message|>" in transformed_prompt, "Prompt should contain <|message|> tag"
|
|
assert "<|end|>" in transformed_prompt, "Prompt should contain <|end|> tag"
|
|
assert (
|
|
"You are chatgpt" in transformed_prompt
|
|
), "Prompt should contain system message content"
|
|
assert (
|
|
"Hi there" in transformed_prompt
|
|
), "Prompt should contain user message content"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.xdist_group("watsonx_heavy")
|
|
async def test_watsonx_gpt_oss_uses_async_http_handler():
|
|
"""
|
|
Test that verifies async HTTP client is used when fetching HuggingFace templates.
|
|
"""
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
|
|
_aget_chat_template_file,
|
|
)
|
|
|
|
# Mock the async HTTP client
|
|
mock_async_client = MagicMock()
|
|
mock_get = AsyncMock()
|
|
mock_async_client.get = mock_get
|
|
|
|
# Create mock response for chat template file
|
|
mock_response = MagicMock()
|
|
mock_response.status_code = 200
|
|
mock_response.content = b"test template content"
|
|
mock_get.return_value = mock_response
|
|
|
|
# Test the async function directly
|
|
with patch(
|
|
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler.get_async_httpx_client",
|
|
return_value=mock_async_client,
|
|
):
|
|
result = await _aget_chat_template_file(hf_model_name="test/model")
|
|
|
|
# Verify async HTTP client was called
|
|
assert mock_get.called, "Async HTTP client's get method should be called"
|
|
assert mock_get.await_count > 0, "Async HTTP client's get should be awaited"
|
|
|
|
# Verify it was called with HuggingFace URL
|
|
call_args = mock_get.call_args
|
|
assert call_args is not None, "get should have been called with arguments"
|
|
called_url = call_args.kwargs.get("url", "")
|
|
assert (
|
|
"huggingface.co/test/model" in called_url
|
|
), f"Should call HuggingFace API for test/model, got: {called_url}"
|
|
assert result["status"] == "success", "Should return success status"
|
|
|
|
|
|
@pytest.mark.parametrize("tokenizer_config_cached", [False, True], ids=["tokenizer_config", "cached_config_jinja"])
|
|
async def test_watsonx_text_gpt_oss_async_completion_fetches_hf_template_off_the_event_loop(
|
|
monkeypatch, tokenizer_config_cached
|
|
):
|
|
import httpx
|
|
|
|
from litellm._uuid import uuid
|
|
from litellm.litellm_core_utils.prompt_templates import huggingface_template_handler
|
|
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
|
|
|
hf_model = f"openai/gpt-oss-{uuid.uuid4()}"
|
|
chat_template = "{% for m in messages %}<|{{ m['role'] }}|>{{ m['content'] }}{% endfor %}"
|
|
if tokenizer_config_cached:
|
|
cached_config = {"status": "success", "tokenizer": {"bos_token": None, "eos_token": None}}
|
|
monkeypatch.setattr(litellm, "known_tokenizer_config", {hf_model: cached_config})
|
|
expected_fetch = f"https://huggingface.co/{hf_model}/raw/main/chat_template.jinja"
|
|
else:
|
|
monkeypatch.setattr(litellm, "known_tokenizer_config", {})
|
|
expected_fetch = f"https://huggingface.co/{hf_model}/raw/main/tokenizer_config.json"
|
|
hf_fetched = []
|
|
captured = {}
|
|
|
|
def forbid_sync_client():
|
|
raise AssertionError("sync HuggingFace fetch ran on the request path")
|
|
|
|
async def serve_hf_file(url, **kwargs):
|
|
hf_fetched.append(url)
|
|
if url.endswith(".jinja"):
|
|
return httpx.Response(200, content=chat_template.encode())
|
|
return httpx.Response(200, json={"chat_template": chat_template, "bos_token": None, "eos_token": None})
|
|
|
|
monkeypatch.setattr(huggingface_template_handler, "_get_httpx_client", forbid_sync_client)
|
|
monkeypatch.setattr(huggingface_template_handler, "get_async_httpx_client", lambda **kwargs: Mock(get=serve_hf_file))
|
|
|
|
def handle(request):
|
|
captured["body"] = json.loads(request.content)
|
|
return httpx.Response(
|
|
200,
|
|
json={
|
|
"model_id": hf_model,
|
|
"results": [
|
|
{
|
|
"generated_text": "Hi",
|
|
"generated_token_count": 1,
|
|
"input_token_count": 1,
|
|
"stop_reason": "eos_token",
|
|
}
|
|
],
|
|
},
|
|
)
|
|
|
|
client = AsyncHTTPHandler()
|
|
client.client = httpx.AsyncClient(transport=httpx.MockTransport(handle))
|
|
|
|
response = await litellm.acompletion(
|
|
model=f"watsonx_text/{hf_model}",
|
|
messages=[{"role": "user", "content": "Hi there"}],
|
|
api_base="https://test-api.watsonx.ai",
|
|
project_id="test-project-id",
|
|
token="test-token",
|
|
client=client,
|
|
)
|
|
|
|
assert response.choices[0].message.content == "Hi"
|
|
assert hf_fetched == [expected_fetch]
|
|
assert captured["body"]["input"] == "<|user|>Hi there"
|
|
|
|
|
|
def test_watsonx_chat_completion_with_reasoning_effort(monkeypatch):
|
|
"""
|
|
Test that 'reasoning_effort' is correctly passed through to the WatsonX API payload.
|
|
"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
|
|
model = "watsonx/openai/gpt-oss-120b"
|
|
messages = [{"role": "user", "content": "Test message"}]
|
|
|
|
client = HTTPHandler()
|
|
|
|
# Mock the token generation call
|
|
mock_token_response = Mock()
|
|
mock_token_response.json.return_value = {
|
|
"access_token": "mock_access_token",
|
|
"expires_in": 3600,
|
|
}
|
|
mock_token_response.raise_for_status = Mock()
|
|
|
|
# Call litellm.completion with the new parameter
|
|
with (
|
|
patch.object(client, "post") as mock_post,
|
|
patch.object(
|
|
litellm.module_level_client, "post", return_value=mock_token_response
|
|
),
|
|
):
|
|
try:
|
|
completion(
|
|
model=model,
|
|
messages=messages,
|
|
api_key="test_api_key",
|
|
client=client,
|
|
reasoning_effort="low",
|
|
)
|
|
except Exception as e:
|
|
print(f"Caught expected exception: {e}")
|
|
|
|
# Verify the parameter is in the final request payload
|
|
assert (
|
|
mock_post.call_count == 1
|
|
), "The completion endpoint should have been called once."
|
|
|
|
# Get the JSON data sent in the POST request
|
|
request_kwargs = mock_post.call_args.kwargs
|
|
json_data = json.loads(request_kwargs["data"])
|
|
|
|
print("\nRequest payload sent to WatsonX API:")
|
|
print(json.dumps(json_data, indent=2))
|
|
|
|
# Check for the parameter at the top level of the payload
|
|
assert (
|
|
"reasoning_effort" in json_data
|
|
), "'reasoning_effort' should be at the top level of the payload."
|
|
assert (
|
|
json_data["reasoning_effort"] == "low"
|
|
), "The value of 'reasoning_effort' should be 'low'."
|
|
|
|
|
|
def test_watsonx_zen_api_key_from_client(monkeypatch, watsonx_chat_completion_call):
|
|
"""
|
|
Test that zen_api_key can be passed from client code and is used in Authorization header.
|
|
"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
|
|
model = "watsonx/ibm/granite-3-3-8b-instruct"
|
|
messages = [{"role": "user", "content": "What is your favorite color?"}]
|
|
|
|
client = HTTPHandler()
|
|
|
|
zen_api_key = "U1ZDLWQo="
|
|
|
|
# No need to patch token call since zen_api_key should skip token generation
|
|
with patch.object(client, "post") as mock_post:
|
|
try:
|
|
completion(
|
|
model=model,
|
|
messages=messages,
|
|
api_key="test_api_key",
|
|
client=client,
|
|
zen_api_key=zen_api_key,
|
|
)
|
|
except Exception as e:
|
|
print(f"Caught expected exception: {e}")
|
|
|
|
# Verify the request was made
|
|
assert (
|
|
mock_post.call_count == 1
|
|
), "The completion endpoint should have been called once."
|
|
|
|
# Get the headers sent in the POST request
|
|
request_kwargs = mock_post.call_args.kwargs
|
|
headers = request_kwargs["headers"]
|
|
|
|
print("\nHeaders sent to WatsonX API:")
|
|
print(json.dumps(dict(headers), indent=2))
|
|
|
|
# Verify Authorization header uses ZenApiKey format
|
|
assert "Authorization" in headers, "Authorization header should be present."
|
|
assert headers["Authorization"] == f"ZenApiKey {zen_api_key}", (
|
|
f"Authorization header should use ZenApiKey format. "
|
|
f"Expected: 'ZenApiKey {zen_api_key}', Got: '{headers['Authorization']}'"
|
|
)
|
|
|
|
|
|
def test_watsonx_zen_api_key_from_env(monkeypatch, watsonx_chat_completion_call):
|
|
"""
|
|
Test that zen_api_key from environment variable is used in Authorization header.
|
|
"""
|
|
monkeypatch.setenv("WATSONX_PROJECT_ID", "test-project-id")
|
|
monkeypatch.setenv("WATSONX_API_BASE", "https://test-api.watsonx.ai")
|
|
|
|
zen_api_key = "U1ZDLWxpdG--==="
|
|
monkeypatch.setenv("WATSONX_ZENAPIKEY", zen_api_key)
|
|
|
|
model = "watsonx/ibm/granite-3-3-8b-instruct"
|
|
messages = [{"role": "user", "content": "What is your favorite color?"}]
|
|
|
|
client = HTTPHandler()
|
|
|
|
# No need to patch token call since zen_api_key should skip token generation
|
|
with patch.object(client, "post") as mock_post:
|
|
try:
|
|
completion(
|
|
model=model,
|
|
messages=messages,
|
|
api_key="test_api_key",
|
|
client=client,
|
|
)
|
|
except Exception as e:
|
|
print(f"Caught expected exception: {e}")
|
|
|
|
# Verify the request was made
|
|
assert (
|
|
mock_post.call_count == 1
|
|
), "The completion endpoint should have been called once."
|
|
|
|
# Get the headers sent in the POST request
|
|
request_kwargs = mock_post.call_args.kwargs
|
|
headers = request_kwargs["headers"]
|
|
|
|
print("\nHeaders sent to WatsonX API:")
|
|
print(json.dumps(dict(headers), indent=2))
|
|
|
|
# Verify Authorization header uses ZenApiKey format
|
|
assert "Authorization" in headers, "Authorization header should be present."
|
|
assert headers["Authorization"] == f"ZenApiKey {zen_api_key}", (
|
|
f"Authorization header should use ZenApiKey format. "
|
|
f"Expected: 'ZenApiKey {zen_api_key}', Got: '{headers['Authorization']}'"
|
|
)
|