diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 84313c6a940..a60f1460495 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -18,7 +18,7 @@ "limit": 40 }, "reportDeprecated": { - "limit": 201 + "limit": 197 }, "reportDuplicateImport": { "limit": 19 diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 018df2b237c..e2376cf8dd0 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -2465,13 +2465,9 @@ class OpenAIAssistantsAPI(BaseLLM): **message_data, ) - response_obj: OpenAIMessage | None = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage.model_validate(thread_message.dict()) - else: - response_obj = OpenAIMessage.model_validate(thread_message.dict()) - return response_obj + return OpenAIMessage.model_validate(thread_message.model_dump()) # fmt: off @@ -2544,13 +2540,9 @@ class OpenAIAssistantsAPI(BaseLLM): **message_data, ) - response_obj: OpenAIMessage | None = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage.model_validate(thread_message.dict()) - else: - response_obj = OpenAIMessage.model_validate(thread_message.dict()) - return response_obj + return OpenAIMessage.model_validate(thread_message.model_dump()) async def async_get_messages( self, diff --git a/tests/test_litellm/llms/openai/test_openai_assistants.py b/tests/test_litellm/llms/openai/test_openai_assistants.py index 36c988749ee..7e679a8a52c 100644 --- a/tests/test_litellm/llms/openai/test_openai_assistants.py +++ b/tests/test_litellm/llms/openai/test_openai_assistants.py @@ -3,7 +3,7 @@ import pytest from openai import AsyncOpenAI, OpenAI from litellm.llms.openai.openai import OpenAIAssistantsAPI -from litellm.types.llms.openai import Thread +from litellm.types.llms.openai import OpenAIMessage, Thread _THREAD_PAYLOAD = { "id": "thread_123", @@ -13,6 +13,20 @@ _THREAD_PAYLOAD = { "unexpected_upstream_field": "kept", } +_MESSAGE_PAYLOAD = { + "id": "msg_123", + "object": "thread.message", + "created_at": 1700000000, + "thread_id": "thread_123", + "role": "assistant", + "status": "in_progress", + "content": [{"type": "text", "text": {"value": "hi", "annotations": []}}], + "metadata": {"origin": "unit-test"}, + "run_id": "run_123", + "assistant_id": "asst_123", + "unexpected_upstream_field": "kept", +} + _COMMON_ARGS = { "api_key": "test-key", "api_base": "https://api.openai.com/v1", @@ -22,23 +36,26 @@ _COMMON_ARGS = { } -def _handler(request: httpx.Request) -> httpx.Response: - return httpx.Response(200, json=_THREAD_PAYLOAD) +def _transport(payload: dict) -> httpx.MockTransport: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json=payload) + + return httpx.MockTransport(handler) -def _async_client() -> AsyncOpenAI: +def _async_client(payload: dict = _THREAD_PAYLOAD) -> AsyncOpenAI: """A real AsyncOpenAI wired to a mock transport, so the SDK's own response parsing runs and the handler under test receives exactly what production would.""" return AsyncOpenAI( api_key="test-key", - http_client=httpx.AsyncClient(transport=httpx.MockTransport(_handler)), + http_client=httpx.AsyncClient(transport=_transport(payload)), ) -def _sync_client() -> OpenAI: +def _sync_client(payload: dict = _THREAD_PAYLOAD) -> OpenAI: return OpenAI( api_key="test-key", - http_client=httpx.Client(transport=httpx.MockTransport(_handler)), + http_client=httpx.Client(transport=_transport(payload)), ) @@ -73,6 +90,64 @@ async def test_async_thread_responses_preserve_declared_fields(): _assert_thread(retrieved) +def _assert_message(result: object) -> None: + assert isinstance(result, OpenAIMessage) + assert result.id == "msg_123" + assert result.thread_id == "thread_123" + assert result.role == "assistant" + assert result.metadata == {"origin": "unit-test"} + assert result.model_dump()["unexpected_upstream_field"] == "kept" + + +@pytest.mark.asyncio +async def test_a_add_message_preserves_fields_and_defaults_status(): + api = OpenAIAssistantsAPI() + + result = await api.a_add_message( + thread_id="thread_123", + message_data={"role": "user", "content": "hi"}, + client=_async_client(_MESSAGE_PAYLOAD), + **_COMMON_ARGS, + ) + + _assert_message(result) + assert result.status == "in_progress" + + without_status = {k: v for k, v in _MESSAGE_PAYLOAD.items() if k != "status"} + defaulted = await api.a_add_message( + thread_id="thread_123", + message_data={"role": "user", "content": "hi"}, + client=_async_client(without_status), + **_COMMON_ARGS, + ) + + assert defaulted.status == "completed" + + +def test_sync_add_message_preserves_fields_and_defaults_status(): + api = OpenAIAssistantsAPI() + + result = api.add_message( + thread_id="thread_123", + message_data={"role": "user", "content": "hi"}, + client=_sync_client(_MESSAGE_PAYLOAD), + **_COMMON_ARGS, + ) + + _assert_message(result) + assert result.status == "in_progress" + + without_status = {k: v for k, v in _MESSAGE_PAYLOAD.items() if k != "status"} + defaulted = api.add_message( + thread_id="thread_123", + message_data={"role": "user", "content": "hi"}, + client=_sync_client(without_status), + **_COMMON_ARGS, + ) + + assert defaulted.status == "completed" + + def test_sync_thread_responses_preserve_declared_fields(): api = OpenAIAssistantsAPI() diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 45b3ec68ca1..573970a5561 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16707 + "limit": 16701 }, "LIT011": { "limit": 5591