diff --git a/python/packages/core/agent_framework/_types.py b/python/packages/core/agent_framework/_types.py index 7b0efa634e..e9b1fcc31b 100644 --- a/python/packages/core/agent_framework/_types.py +++ b/python/packages/core/agent_framework/_types.py @@ -9,7 +9,6 @@ import re import sys import warnings -from asyncio import iscoroutine from collections.abc import ( AsyncGenerator, AsyncIterable, @@ -3379,8 +3378,8 @@ async def _get_stream(self) -> AsyncIterable[UpdateT]: if hasattr(self._stream_source, "__aiter__"): self._stream = self._stream_source # type: ignore[assignment] else: - if not iscoroutine(self._stream_source): - self._stream = self._stream_source # type: ignore[assignment] + if not isawaitable(self._stream_source): + self._stream = self._stream_source else: self._stream = await self._stream_source if isinstance(self._stream, ResponseStream) and self._wrap_inner: diff --git a/python/packages/core/tests/core/test_types.py b/python/packages/core/tests/core/test_types.py index f7669221bd..9589f1b032 100644 --- a/python/packages/core/tests/core/test_types.py +++ b/python/packages/core/tests/core/test_types.py @@ -1,9 +1,10 @@ # Copyright (c) Microsoft. All rights reserved. +import asyncio import base64 import json import warnings -from collections.abc import AsyncIterable, Sequence +from collections.abc import AsyncIterable, Awaitable, Sequence from dataclasses import dataclass from datetime import datetime, timezone from typing import Any, Literal, cast @@ -4529,13 +4530,22 @@ async def async_map(update: ChatResponseUpdate) -> ChatResponseUpdate: assert collected == ["async_update_0", "async_update_1"] - async def test_from_awaitable(self) -> None: + @pytest.mark.parametrize("source_kind", ["coroutine", "task", "future"]) + async def test_from_awaitable(self, source_kind: str) -> None: """from_awaitable() wraps an awaitable ResponseStream.""" async def get_stream() -> ResponseStream[ChatResponseUpdate, ChatResponse]: return ResponseStream(_generate_updates(2), finalizer=_combine_updates) - outer = ResponseStream.from_awaitable(get_stream()) + source: Awaitable[ResponseStream[ChatResponseUpdate, ChatResponse]] + if source_kind == "task": + source = asyncio.create_task(get_stream()) + elif source_kind == "future": + source = asyncio.get_running_loop().create_future() + source.set_result(await get_stream()) + else: + source = get_stream() + outer = ResponseStream.from_awaitable(source) collected: list[str] = [] async for update in outer: @@ -4614,13 +4624,22 @@ def finalizer(updates: list[ChatResponseUpdate]) -> ChatResponse: class TestResponseStreamAwaitableSource: """Tests for ResponseStream with awaitable stream sources.""" - async def test_awaitable_stream_source(self) -> None: + @pytest.mark.parametrize("source_kind", ["coroutine", "task", "future"]) + async def test_awaitable_stream_source(self, source_kind: str) -> None: """ResponseStream can accept an awaitable that resolves to an async iterable.""" async def get_stream() -> AsyncIterable[ChatResponseUpdate]: return _generate_updates(2) - stream = ResponseStream(get_stream(), finalizer=_combine_updates) + source: Awaitable[AsyncIterable[ChatResponseUpdate]] + if source_kind == "task": + source = asyncio.create_task(get_stream()) + elif source_kind == "future": + source = asyncio.get_running_loop().create_future() + source.set_result(await get_stream()) + else: + source = get_stream() + stream = ResponseStream(source, finalizer=_combine_updates) collected: list[str] = [] async for update in stream: