litellm/tests/test_litellm/llms/watsonx/test_watsonx.py
devin-ai-integration[bot] 0f886d9006
perf: move Anthropic, Vertex Anthropic, Ollama and HF template fetches off the event loop (#40311)
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>
2026-09-08 15:53:43 -07:00

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']}'"
)