diff --git a/tests/test_litellm/llms/watsonx/test_watsonx.py b/tests/test_litellm/llms/watsonx/test_watsonx.py index 1ab21ac6dc8..6cb12a6ac62 100644 --- a/tests/test_litellm/llms/watsonx/test_watsonx.py +++ b/tests/test_litellm/llms/watsonx/test_watsonx.py @@ -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, (