diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index 1ab85a5b124..eff9adfb243 100644 --- a/litellm/assistants/main.py +++ b/litellm/assistants/main.py @@ -846,6 +846,15 @@ def get_messages( ### RUNS ### +def arun_thread_stream( + *, + event_handler: Optional[AssistantEventHandler] = None, + **kwargs, +) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]: + kwargs["arun_thread"] = True + return run_thread(stream=True, event_handler=event_handler, **kwargs) # type: ignore + + async def arun_thread( custom_llm_provider: Literal["openai", "azure"], thread_id: str, diff --git a/litellm/llms/azure.py b/litellm/llms/azure.py index c907e3b0e76..f0279d0d756 100644 --- a/litellm/llms/azure.py +++ b/litellm/llms/azure.py @@ -1990,7 +1990,7 @@ class AzureAssistantsAPI(BaseLLM): return response - async def async_run_thread_stream( + def async_run_thread_stream( self, client: AsyncAzureOpenAI, thread_id: str, diff --git a/litellm/llms/openai.py b/litellm/llms/openai.py index 69e510ae730..dec86d35d95 100644 --- a/litellm/llms/openai.py +++ b/litellm/llms/openai.py @@ -2534,7 +2534,7 @@ class OpenAIAssistantsAPI(BaseLLM): return response - async def async_run_thread_stream( + def async_run_thread_stream( self, client: AsyncOpenAI, thread_id: str, diff --git a/litellm/tests/test_assistants.py b/litellm/tests/test_assistants.py index cf4adf0fa4b..5f565f67ccd 100644 --- a/litellm/tests/test_assistants.py +++ b/litellm/tests/test_assistants.py @@ -19,6 +19,7 @@ from litellm.llms.openai import ( AsyncCursorPage, SyncCursorPage, AssistantEventHandler, + AsyncAssistantEventHandler, ) from typing_extensions import override @@ -131,29 +132,6 @@ async def test_add_message_litellm(sync_mode, provider): assert isinstance(added_message, Message) -class EventHandler(AssistantEventHandler): - @override - def on_text_created(self, text) -> None: - print(f"\nassistant > ", end="", flush=True) - - @override - def on_text_delta(self, delta, snapshot): - print(delta.value, end="", flush=True) - - def on_tool_call_created(self, tool_call): - print(f"\nassistant > {tool_call.type}\n", flush=True) - - def on_tool_call_delta(self, delta, snapshot): - if delta.type == "code_interpreter": - if delta.code_interpreter.input: - print(delta.code_interpreter.input, end="", flush=True) - if delta.code_interpreter.outputs: - print(f"\n\noutput >", flush=True) - for output in delta.code_interpreter.outputs: - if output.type == "logs": - print(f"\n{output.logs}", flush=True) - - @pytest.mark.parametrize( "provider", [ @@ -163,8 +141,11 @@ class EventHandler(AssistantEventHandler): ) # @pytest.mark.parametrize( "sync_mode", - [False, True], -) # + [ + True, + False, + ], +) @pytest.mark.parametrize( "is_streaming", [True, False], @@ -202,9 +183,7 @@ async def test_aarun_thread_litellm(sync_mode, provider, is_streaming): added_message = litellm.add_message(**data) if is_streaming: - run = litellm.run_thread_stream( - assistant_id=assistant_id, event_handler=EventHandler(), **data - ) + run = litellm.run_thread_stream(assistant_id=assistant_id, **data) with run as run: assert isinstance(run, AssistantEventHandler) print(run) @@ -225,11 +204,13 @@ async def test_aarun_thread_litellm(sync_mode, provider, is_streaming): added_message = await litellm.a_add_message(**data) if is_streaming: - run = litellm.run_thread_stream( - assistant_id=assistant_id, event_handler=EventHandler(), **data - ) - with run as run: - assert isinstance(run, AssistantEventHandler) + run = litellm.arun_thread_stream(assistant_id=assistant_id, **data) + async with run as run: + print(f"run: {run}") + assert isinstance( + run, + AsyncAssistantEventHandler, + ) print(run) run.until_done() else: