Skip to content

Commit 44e6b14

Browse files
fix(ai): support async with on async streaming responses (Fixes #393) (#645)
1 parent 5b4a91a commit 44e6b14

10 files changed

Lines changed: 454 additions & 4 deletions

File tree

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
---
2+
pypi/posthog: patch
3+
---
4+
5+
Fix async streaming responses from the AI wrappers (OpenAI, Anthropic, Gemini) so they support `async with` as well as `async for`. Previously, consuming a stream via `async with` (e.g. with pydantic-ai) raised `TypeError: 'async_generator' object does not support the asynchronous context manager protocol`.

‎posthog/ai/anthropic/anthropic_async.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
from typing import Any, Dict, List, Optional
1212

1313
from posthog import setup
14+
from posthog.ai.stream import AsyncStreamWrapper
1415
from posthog.ai.types import StreamingContentBlock, TokenUsage, ToolInProgress
1516
from posthog.ai.utils import (
1617
call_llm_and_track_usage_async,
@@ -225,7 +226,7 @@ async def generator():
225226
stop_reason=stop_reason,
226227
)
227228

228-
return generator()
229+
return AsyncStreamWrapper(generator(), stream=response)
229230

230231
async def _capture_streaming_event(
231232
self,

‎posthog/ai/gemini/gemini_async.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import uuid
44
from typing import Any, Dict, Optional
55

6+
from posthog.ai.stream import AsyncStreamWrapper
67
from posthog.ai.types import TokenUsage, StreamingEventData
78
from posthog.ai.utils import merge_system_prompt
89

@@ -354,7 +355,7 @@ async def async_generator():
354355
stop_reason=stop_reason,
355356
)
356357

357-
return async_generator()
358+
return AsyncStreamWrapper(async_generator(), stream=response)
358359

359360
def _capture_streaming_event(
360361
self,

‎posthog/ai/openai/openai_async.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
import uuid
33
from typing import Any, Dict, List, Optional
44

5+
from posthog.ai.stream import AsyncStreamWrapper
56
from posthog.ai.types import TokenUsage
67

78
try:
@@ -221,7 +222,7 @@ async def async_generator():
221222
stop_reason=stop_reason,
222223
)
223224

224-
return async_generator()
225+
return AsyncStreamWrapper(async_generator(), stream=response)
225226

226227
async def _capture_streaming_event(
227228
self,
@@ -515,7 +516,7 @@ async def async_generator():
515516
stop_reason=stop_reason,
516517
)
517518

518-
return async_generator()
519+
return AsyncStreamWrapper(async_generator(), stream=response)
519520

520521
async def _capture_streaming_event(
521522
self,

‎posthog/ai/stream.py‎

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
"""Shared async streaming utilities for PostHog AI wrappers."""
2+
3+
from typing import Any, AsyncGenerator, Generic, Optional, TypeVar
4+
5+
T = TypeVar("T")
6+
7+
8+
class AsyncStreamWrapper(Generic[T]):
9+
"""Adds the async context manager protocol to a PostHog streaming generator.
10+
11+
The OpenAI and Anthropic SDK streams support both ``async for`` and
12+
``async with``. PostHog's wrappers returned a bare async generator, which
13+
only supports ``async for``, so ``async with response:`` (used by
14+
pydantic-ai) raised a TypeError. This wraps the tracking generator and,
15+
when given the original provider stream, closes it and proxies attribute
16+
access (e.g. ``.response``) to it.
17+
"""
18+
19+
def __init__(
20+
self,
21+
generator: AsyncGenerator[T, None],
22+
stream: Optional[Any] = None,
23+
) -> None:
24+
self._generator = generator
25+
self._stream = stream
26+
27+
def __aiter__(self) -> "AsyncStreamWrapper[T]":
28+
return self
29+
30+
async def __anext__(self) -> T:
31+
return await self._generator.__anext__()
32+
33+
async def __aenter__(self) -> "AsyncStreamWrapper[T]":
34+
return self
35+
36+
async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> bool:
37+
# Close the generator first so its `finally` captures the event, even on
38+
# early exit. try/finally still closes the provider stream if that raises.
39+
try:
40+
await self._generator.aclose()
41+
finally:
42+
if self._stream is not None:
43+
close = getattr(self._stream, "aclose", None) or getattr(
44+
self._stream, "close", None
45+
)
46+
if close is not None:
47+
await close()
48+
49+
return False
50+
51+
# aclose/asend/athrow belong to the generator; provider streams expose
52+
# close(), not these. Forwarding aclose() keeps it firing the event.
53+
_GENERATOR_METHODS = ("aclose", "asend", "athrow")
54+
55+
def __getattr__(self, name: str) -> Any:
56+
# Proxy only public attributes (e.g. `.response`) to the provider stream.
57+
if name.startswith("_"):
58+
raise AttributeError(name)
59+
if name in self._GENERATOR_METHODS:
60+
return getattr(self._generator, name)
61+
target = self._stream if self._stream is not None else self._generator
62+
return getattr(target, name)

‎posthog/test/ai/anthropic/test_anthropic.py‎

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from anthropic.types import Message, Usage
1111

1212
from posthog.ai.anthropic import Anthropic, AsyncAnthropic
13+
from posthog.test.ai.utils import RecordingAsyncStream
1314

1415
ANTHROPIC_AVAILABLE = True
1516
except ImportError:
@@ -1421,3 +1422,74 @@ def test_integration_stop_reason(mock_client):
14211422
assert props["$ai_stop_reason"] in ("end_turn", "max_tokens")
14221423
assert props["$ai_provider"] == "anthropic"
14231424
assert props["$ai_input_tokens"] > 0
1425+
1426+
1427+
def _anthropic_stream_events():
1428+
final = MockStreamEvent("message_delta")
1429+
final.usage = MockUsage(
1430+
input_tokens=10,
1431+
output_tokens=5,
1432+
cache_read_input_tokens=0,
1433+
cache_creation_input_tokens=0,
1434+
)
1435+
return [
1436+
MockStreamEvent("message_start"),
1437+
MockStreamEvent("content_block_delta", text="Hi"),
1438+
final,
1439+
]
1440+
1441+
1442+
@pytest.mark.asyncio
1443+
async def test_async_messages_create_streaming_supports_async_with(mock_client):
1444+
"""Regression test for #393: messages.create(stream=True) must support
1445+
`async with`."""
1446+
1447+
async def mock_async_create(**kwargs):
1448+
return RecordingAsyncStream(_anthropic_stream_events())
1449+
1450+
with patch(
1451+
"anthropic.resources.messages.AsyncMessages.create",
1452+
side_effect=mock_async_create,
1453+
):
1454+
client = AsyncAnthropic(posthog_client=mock_client)
1455+
response = await client.messages.create(
1456+
model="claude-3-opus-20240229",
1457+
messages=[{"role": "user", "content": "Foo"}],
1458+
stream=True,
1459+
max_tokens=1,
1460+
)
1461+
1462+
async with response as stream:
1463+
events = [event async for event in stream]
1464+
1465+
assert len(events) == 3
1466+
assert mock_client.capture.call_count == 1
1467+
1468+
1469+
@pytest.mark.asyncio
1470+
async def test_async_messages_streaming_early_exit_closes_provider_stream(mock_client):
1471+
"""Breaking out early must close the underlying Anthropic stream and still
1472+
capture the event."""
1473+
source = RecordingAsyncStream(_anthropic_stream_events())
1474+
1475+
async def mock_async_create(**kwargs):
1476+
return source
1477+
1478+
with patch(
1479+
"anthropic.resources.messages.AsyncMessages.create",
1480+
side_effect=mock_async_create,
1481+
):
1482+
client = AsyncAnthropic(posthog_client=mock_client)
1483+
response = await client.messages.create(
1484+
model="claude-3-opus-20240229",
1485+
messages=[{"role": "user", "content": "Foo"}],
1486+
stream=True,
1487+
max_tokens=1,
1488+
)
1489+
1490+
async with response as stream:
1491+
async for _ in stream:
1492+
break
1493+
1494+
assert source.closed is True
1495+
assert mock_client.capture.call_count == 1

‎posthog/test/ai/gemini/test_gemini_async.py‎

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1110,3 +1110,43 @@ async def test_async_embed_content_integration_batch(mock_client):
11101110

11111111
assert response.embeddings is not None
11121112
assert len(response.embeddings) == len(inputs)
1113+
1114+
1115+
async def test_async_client_streaming_supports_async_with(
1116+
mock_client, mock_google_genai_client
1117+
):
1118+
"""Regression test for #393: generate_content_stream must support `async with`."""
1119+
1120+
async def mock_streaming_response():
1121+
chunk = MagicMock()
1122+
chunk.text = "Hi"
1123+
usage = MagicMock()
1124+
usage.prompt_token_count = 5
1125+
usage.candidates_token_count = 3
1126+
usage.cached_content_token_count = 0
1127+
usage.thoughts_token_count = 0
1128+
chunk.usage_metadata = usage
1129+
yield chunk
1130+
1131+
mock_google_genai_client.aio.models.generate_content_stream = AsyncMock(
1132+
return_value=mock_streaming_response()
1133+
)
1134+
1135+
client = AsyncClient(api_key="test-key", posthog_client=mock_client)
1136+
1137+
response = await client.models.generate_content_stream(
1138+
model="gemini-2.0-flash",
1139+
contents=["Hi"],
1140+
posthog_distinct_id="test-id",
1141+
)
1142+
1143+
chunks = []
1144+
async with response as stream:
1145+
async for chunk in stream:
1146+
chunks.append(chunk)
1147+
1148+
assert len(chunks) == 1
1149+
assert mock_client.capture.call_count == 1
1150+
call_args = mock_client.capture.call_args[1]
1151+
assert call_args["event"] == "$ai_generation"
1152+
assert call_args["properties"]["$ai_provider"] == "gemini"

‎posthog/test/ai/openai/test_openai.py‎

Lines changed: 98 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@
3939
from posthog.ai.openai import OpenAI
4040
from posthog.ai.openai.openai_async import AsyncOpenAI
4141
from posthog.ai.openai.wrapper_utils import reset_fallback_warnings
42+
from posthog.test.ai.utils import RecordingAsyncStream
4243

4344
OPENAI_AVAILABLE = True
4445
except ImportError:
@@ -2352,3 +2353,100 @@ def test_integration_stop_reason(mock_client):
23522353
assert props["$ai_stop_reason"] in ("stop", "length")
23532354
assert props["$ai_provider"] == "openai"
23542355
assert props["$ai_input_tokens"] > 0
2356+
2357+
2358+
@pytest.mark.asyncio
2359+
async def test_async_chat_streaming_supports_async_with(
2360+
mock_client, streaming_tool_call_chunks
2361+
):
2362+
"""Regression test for #393: chat completions stream=True must support
2363+
`async with` (the protocol pydantic-ai relies on)."""
2364+
2365+
async def mock_create(self, **kwargs):
2366+
return RecordingAsyncStream(streaming_tool_call_chunks)
2367+
2368+
with patch(
2369+
"openai.resources.chat.completions.AsyncCompletions.create", new=mock_create
2370+
):
2371+
client = AsyncOpenAI(api_key="test-key", posthog_client=mock_client)
2372+
2373+
response = await client.chat.completions.create(
2374+
model="gpt-4",
2375+
messages=[{"role": "user", "content": "Hi"}],
2376+
stream=True,
2377+
posthog_distinct_id="test-id",
2378+
)
2379+
2380+
chunks = []
2381+
async with response as stream:
2382+
async for chunk in stream:
2383+
chunks.append(chunk)
2384+
2385+
assert chunks == streaming_tool_call_chunks
2386+
assert mock_client.capture.call_count == 1
2387+
call_args = mock_client.capture.call_args[1]
2388+
props = call_args["properties"]
2389+
assert call_args["event"] == "$ai_generation"
2390+
assert props["$ai_provider"] == "openai"
2391+
assert props["$ai_model"] == "gpt-4"
2392+
2393+
2394+
@pytest.mark.asyncio
2395+
async def test_async_responses_streaming_supports_async_with(mock_client):
2396+
"""Regression test for #393: responses stream=True must support
2397+
`async with`."""
2398+
from unittest.mock import MagicMock
2399+
2400+
chunk = MagicMock()
2401+
chunk.type = "response.text.delta"
2402+
chunk.text = "hello"
2403+
2404+
async def mock_create(self, **kwargs):
2405+
return RecordingAsyncStream([chunk])
2406+
2407+
with patch("openai.resources.responses.AsyncResponses.create", new=mock_create):
2408+
client = AsyncOpenAI(api_key="test-key", posthog_client=mock_client)
2409+
2410+
response = await client.responses.create(
2411+
model="gpt-4o-mini",
2412+
input=[{"role": "user", "content": "Hi"}],
2413+
stream=True,
2414+
posthog_distinct_id="test-id",
2415+
)
2416+
2417+
async with response as stream:
2418+
received = [c async for c in stream]
2419+
2420+
assert received == [chunk]
2421+
assert mock_client.capture.call_count == 1
2422+
2423+
2424+
@pytest.mark.asyncio
2425+
async def test_async_chat_streaming_early_exit_closes_provider_stream(
2426+
mock_client, streaming_tool_call_chunks
2427+
):
2428+
"""Breaking out of the stream early must close the underlying provider
2429+
stream (release the HTTP connection) and still capture the event."""
2430+
source = RecordingAsyncStream(streaming_tool_call_chunks)
2431+
2432+
async def mock_create(self, **kwargs):
2433+
return source
2434+
2435+
with patch(
2436+
"openai.resources.chat.completions.AsyncCompletions.create", new=mock_create
2437+
):
2438+
client = AsyncOpenAI(api_key="test-key", posthog_client=mock_client)
2439+
2440+
response = await client.chat.completions.create(
2441+
model="gpt-4",
2442+
messages=[{"role": "user", "content": "Hi"}],
2443+
stream=True,
2444+
posthog_distinct_id="test-id",
2445+
)
2446+
2447+
async with response as stream:
2448+
async for _ in stream:
2449+
break
2450+
2451+
assert source.closed is True
2452+
assert mock_client.capture.call_count == 1

0 commit comments

Comments
 (0)