Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 2 additions & 3 deletions python/packages/core/agent_framework/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
import re
import sys
import warnings
from asyncio import iscoroutine
from collections.abc import (
AsyncGenerator,
AsyncIterable,
Expand Down Expand Up @@ -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:
Expand Down
29 changes: 24 additions & 5 deletions python/packages/core/tests/core/test_types.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
Loading