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:
yuneng-jiang 2026-03-09 15:32:30 -07:00
parent 2e34870df9
commit 74ed6a16ac

View file

@ -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, (