test_init_responses_api_endpoints

This commit is contained in:
Ishaan Jaff 2025-04-24 21:43:59 -07:00
parent cb87dbbd51
commit 8ed3557ce7

View file

@ -6,7 +6,7 @@ from typing import Optional
from dotenv import load_dotenv
from fastapi import Request
from datetime import datetime
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, patch, MagicMock
sys.path.insert(
0, os.path.abspath("../..")
@ -553,9 +553,67 @@ def test_initialize_router_endpoints():
assert hasattr(router, "aanthropic_messages")
assert hasattr(router, "aresponses")
assert hasattr(router, "responses")
assert hasattr(router, "aget_responses")
assert hasattr(router, "adelete_responses")
# Verify the endpoints are callable
assert callable(router.amoderation)
assert callable(router.aanthropic_messages)
assert callable(router.aresponses)
assert callable(router.responses)
assert callable(router.aget_responses)
assert callable(router.adelete_responses)
@pytest.mark.asyncio
async def test_init_responses_api_endpoints():
"""
A simpler test for _init_responses_api_endpoints that focuses on the basic functionality
"""
from litellm.responses.utils import ResponsesAPIRequestUtils
# Create a router with a basic model
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "openai/test-model",
"api_key": "fake-api-key",
},
}
]
)
# Just mock the _ageneric_api_call_with_fallbacks method
router._ageneric_api_call_with_fallbacks = AsyncMock()
# Add a mock implementation of _get_model_id_from_response_id to the Router instance
ResponsesAPIRequestUtils.get_model_id_from_response_id = MagicMock(return_value=None)
# Call without a response_id (no model extraction should happen)
await router._init_responses_api_endpoints(
original_function=AsyncMock(),
thread_id="thread_xyz"
)
# Verify _ageneric_api_call_with_fallbacks was called but model wasn't changed
first_call_kwargs = router._ageneric_api_call_with_fallbacks.call_args.kwargs
assert "model" not in first_call_kwargs
assert first_call_kwargs["thread_id"] == "thread_xyz"
# Reset the mock
router._ageneric_api_call_with_fallbacks.reset_mock()
# Change the return value for the second call
ResponsesAPIRequestUtils.get_model_id_from_response_id.return_value = "claude-3-sonnet"
# Call with a response_id
await router._init_responses_api_endpoints(
original_function=AsyncMock(),
response_id="resp_claude_123"
)
# Verify model was updated in the kwargs
second_call_kwargs = router._ageneric_api_call_with_fallbacks.call_args.kwargs
assert second_call_kwargs["model"] == "claude-3-sonnet"
assert second_call_kwargs["response_id"] == "resp_claude_123"