mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Fix flaky test_watsonx_gpt_oss_prompt_transformation
The test was flaky under pytest-xdist parallel execution because it used async acompletion (which runs completion() in a thread pool via run_in_executor) and relied on shared global state (known_tokenizer_config, iam_token_cache, module_level_client) that could be modified by other tests running in parallel. Failures were silently swallowed by a broad try/except, causing mock_post.call_count to remain 0. Fix: - Convert from async acompletion to sync completion, matching every other test in the file. The test's intent is verifying prompt transformation, not async behavior. - Use monkeypatch.setitem for known_tokenizer_config to ensure proper teardown isolation. - Remove unnecessary mock layers (async template fetchers, iam_token_cache pre-population, mock completion response) that were only needed for the async code path. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
2e34870df9
commit
74ed6a16ac
1 changed files with 19 additions and 101 deletions
|
|
@ -207,12 +207,11 @@ def test_watsonx_completion_regular_model_includes_model_id(
|
|||
assert "project_id" in json_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_watsonx_gpt_oss_prompt_transformation(monkeypatch):
|
||||
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 starts from litellm.acompletion and verifies what gets sent in the final POST request body.
|
||||
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".
|
||||
"""
|
||||
|
|
@ -228,39 +227,12 @@ async def test_watsonx_gpt_oss_prompt_transformation(monkeypatch):
|
|||
{"role": "user", "content": "Hi there"},
|
||||
]
|
||||
|
||||
# Mock the HTTP client
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
# Mock the token 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()
|
||||
|
||||
# Mock the completion call
|
||||
mock_completion_response = Mock()
|
||||
mock_completion_response.status_code = 200
|
||||
mock_completion_response.json.return_value = {
|
||||
"results": [
|
||||
{
|
||||
"generated_text": "Hello! How can I help you?",
|
||||
"generated_token_count": 10,
|
||||
"input_token_count": 5,
|
||||
"stop_reason": "stop", # Required field for response transformation
|
||||
}
|
||||
],
|
||||
"model_id": "openai/gpt-oss-120b",
|
||||
}
|
||||
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.
|
||||
from unittest.mock import patch
|
||||
|
||||
#
|
||||
# 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:
|
||||
|
|
@ -276,100 +248,46 @@ async def test_watsonx_gpt_oss_prompt_transformation(monkeypatch):
|
|||
},
|
||||
}
|
||||
|
||||
async def mock_aget_tokenizer_config(hf_model_name: str):
|
||||
return mock_tokenizer_config
|
||||
|
||||
async def mock_aget_chat_template_file(hf_model_name: str):
|
||||
# Return failure to use tokenizer_config instead
|
||||
return {"status": "failure"}
|
||||
|
||||
# Set cached tokenizer config directly to avoid race conditions with parallel tests.
|
||||
# When running with pytest-xdist (-n 16), another test might populate the cache between
|
||||
# clearing it and the actual usage. By setting the cache directly, we ensure the correct
|
||||
# template is always used regardless of test execution order.
|
||||
# Isolate known_tokenizer_config so parallel tests don't interfere.
|
||||
# monkeypatch.setitem restores the original value on teardown.
|
||||
hf_model = "openai/gpt-oss-120b"
|
||||
litellm.known_tokenizer_config[hf_model] = mock_tokenizer_config
|
||||
monkeypatch.setitem(litellm.known_tokenizer_config, hf_model, mock_tokenizer_config)
|
||||
|
||||
# Also create sync mock functions in case the fallback sync path is used
|
||||
def mock_get_tokenizer_config(hf_model_name: str):
|
||||
return mock_tokenizer_config
|
||||
|
||||
def mock_get_chat_template_file(hf_model_name: str):
|
||||
return {"status": "failure"}
|
||||
|
||||
# Async mock function for client.post to properly handle async method mocking
|
||||
async def mock_post_func(*args, **kwargs):
|
||||
return mock_completion_response
|
||||
|
||||
# Mock the token generation response to avoid actual API call
|
||||
mock_token_get_response = Mock()
|
||||
mock_token_get_response.json.return_value = {
|
||||
# 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_get_response.raise_for_status = Mock()
|
||||
mock_token_response.raise_for_status = Mock()
|
||||
|
||||
with patch.object(client, "post", side_effect=mock_post_func) as mock_post, patch.object(
|
||||
litellm.module_level_client, "post", return_value=mock_token_get_response
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._aget_tokenizer_config",
|
||||
side_effect=mock_aget_tokenizer_config,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._aget_chat_template_file",
|
||||
side_effect=mock_aget_chat_template_file,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._get_tokenizer_config",
|
||||
side_effect=mock_get_tokenizer_config,
|
||||
), patch(
|
||||
"litellm.litellm_core_utils.prompt_templates.huggingface_template_handler._get_chat_template_file",
|
||||
side_effect=mock_get_chat_template_file,
|
||||
with patch.object(client, "post") as mock_post, patch.object(
|
||||
litellm.module_level_client, "post", return_value=mock_token_response
|
||||
):
|
||||
try:
|
||||
# Call acompletion with messages
|
||||
await litellm.acompletion(
|
||||
completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_key="test_api_key",
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
# May fail due to incomplete mocking, but we should have captured the request
|
||||
print(f"Exception (may be expected): {e}")
|
||||
print(f"Caught expected exception: {e}")
|
||||
|
||||
# Verify the POST was called
|
||||
assert (
|
||||
mock_post.call_count >= 1
|
||||
), f"POST should have been called at least once, got {mock_post.call_count}"
|
||||
mock_post.call_count == 1
|
||||
), f"POST should have been called exactly once, got {mock_post.call_count}"
|
||||
|
||||
# Get the request body from the first call
|
||||
# Use call_args_list to be more robust - get the first call's arguments
|
||||
assert len(mock_post.call_args_list) > 0, "mock_post should have at least one call"
|
||||
call_args = mock_post.call_args_list[0]
|
||||
assert call_args is not None, "call_args should not be None"
|
||||
# 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"])
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"Input messages to litellm.acompletion:")
|
||||
print(json.dumps(messages, indent=2))
|
||||
print(f"\n{'='*80}")
|
||||
print(f"Final POST request body:")
|
||||
print(json.dumps(json_data, indent=2))
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# Verify the transformed input is in the request
|
||||
assert "input" in json_data, "Request should have 'input' field"
|
||||
transformed_prompt = json_data["input"]
|
||||
|
||||
# Verify transformation occurred
|
||||
assert transformed_prompt is not None, (
|
||||
"Prompt transformation failed - the template should have been applied to transform "
|
||||
"messages into the correct format for gpt-oss-120b."
|
||||
)
|
||||
|
||||
print(f"Transformed prompt: {repr(transformed_prompt)}")
|
||||
print(f"Prompt length: {len(transformed_prompt)}")
|
||||
|
||||
# Verify it's NOT simple concatenation
|
||||
simple_concat = "You are chatgpt Hi there"
|
||||
assert transformed_prompt != simple_concat, (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue