diff --git a/.sampo/changesets/roguish-baron-tuonetar.md b/.sampo/changesets/roguish-baron-tuonetar.md new file mode 100644 index 00000000..097c6943 --- /dev/null +++ b/.sampo/changesets/roguish-baron-tuonetar.md @@ -0,0 +1,5 @@ +--- +pypi/posthog: patch +--- + +Preserve Anthropic messages.stream compatibility diff --git a/posthog/ai/anthropic/anthropic.py b/posthog/ai/anthropic/anthropic.py index a5a2299b..a25bedc4 100644 --- a/posthog/ai/anthropic/anthropic.py +++ b/posthog/ai/anthropic/anthropic.py @@ -10,6 +10,7 @@ import uuid from typing import Any, Dict, List, Optional +from ..stream import _StreamWrapper from posthog.ai.types import StreamingContentBlock, TokenUsage, ToolInProgress from posthog.ai.utils import ( call_llm_and_track_usage, @@ -115,37 +116,76 @@ def stream( **kwargs: Arguments passed to Anthropic's ``messages.create`` API. Returns: - A streaming iterator yielding Anthropic events. + Anthropic's native streaming context manager. """ if posthog_trace_id is None: posthog_trace_id = str(uuid.uuid4()) - return self._create_streaming( + # Construct the provider resource directly so older Anthropic versions, + # whose stream manager delegates through ``self.create(stream=True)``, + # cannot re-enter our tracked ``create`` override. + manager = Messages(self._client).stream(**kwargs) + request_attribute = "_MessageStreamManager__api_request" + request = getattr(manager, request_attribute, None) + if request is None: + return manager + + def tracked_request(): + start_time = time.time() + response = request() + return self._track_streaming_response( + response, + posthog_distinct_id, + posthog_trace_id, + posthog_properties, + posthog_privacy_mode, + posthog_groups, + kwargs, + start_time, + ) + + setattr(manager, request_attribute, tracked_request) + return manager + + def _create_streaming( + self, + posthog_distinct_id: Optional[str], + posthog_trace_id: Optional[str], + posthog_properties: Optional[Dict[str, Any]], + posthog_privacy_mode: bool, + posthog_groups: Optional[Dict[str, Any]], + **kwargs: Any, + ): + start_time = time.time() + response = super().create(**kwargs) + return self._track_streaming_response( + response, posthog_distinct_id, posthog_trace_id, posthog_properties, posthog_privacy_mode, posthog_groups, - **kwargs, + kwargs, + start_time, ) - def _create_streaming( + def _track_streaming_response( self, + response: Any, posthog_distinct_id: Optional[str], posthog_trace_id: Optional[str], posthog_properties: Optional[Dict[str, Any]], posthog_privacy_mode: bool, posthog_groups: Optional[Dict[str, Any]], - **kwargs: Any, + kwargs: Dict[str, Any], + start_time: float, ): - start_time = time.time() usage_stats: TokenUsage = TokenUsage(input_tokens=0, output_tokens=0) accumulated_content = "" content_blocks: List[StreamingContentBlock] = [] tools_in_progress: Dict[str, ToolInProgress] = {} current_text_block: Optional[StreamingContentBlock] = None stop_reason: Optional[str] = None - response = super().create(**kwargs) def generator(): nonlocal usage_stats @@ -224,7 +264,7 @@ def generator(): stop_reason=stop_reason, ) - return generator() + return _StreamWrapper(generator(), stream=response) def _capture_streaming_event( self, diff --git a/posthog/ai/anthropic/anthropic_async.py b/posthog/ai/anthropic/anthropic_async.py index 8ea07840..eb642de8 100644 --- a/posthog/ai/anthropic/anthropic_async.py +++ b/posthog/ai/anthropic/anthropic_async.py @@ -95,7 +95,7 @@ async def create( **kwargs, ) - async def stream( + def stream( self, posthog_distinct_id: Optional[str] = None, posthog_trace_id: Optional[str] = None, @@ -116,37 +116,76 @@ async def stream( **kwargs: Arguments passed to Anthropic's async ``messages.create`` API. Returns: - An async streaming iterator yielding Anthropic events. + Anthropic's native async streaming context manager, without awaiting. """ if posthog_trace_id is None: posthog_trace_id = str(uuid.uuid4()) - return await self._create_streaming( + # Construct the provider resource directly so older Anthropic versions, + # whose stream manager delegates through ``self.create(stream=True)``, + # cannot re-enter our tracked ``create`` override. + manager = AsyncMessages(self._client).stream(**kwargs) + request_attribute = "_AsyncMessageStreamManager__api_request" + request = getattr(manager, request_attribute, None) + if request is None: + return manager + + async def tracked_request(): + start_time = time.time() + response = await request + return self._track_streaming_response( + response, + posthog_distinct_id, + posthog_trace_id, + posthog_properties, + posthog_privacy_mode, + posthog_groups, + kwargs, + start_time, + ) + + setattr(manager, request_attribute, tracked_request()) + return manager + + async def _create_streaming( + self, + posthog_distinct_id: Optional[str], + posthog_trace_id: Optional[str], + posthog_properties: Optional[Dict[str, Any]], + posthog_privacy_mode: bool, + posthog_groups: Optional[Dict[str, Any]], + **kwargs: Any, + ): + start_time = time.time() + response = await super().create(**kwargs) + return self._track_streaming_response( + response, posthog_distinct_id, posthog_trace_id, posthog_properties, posthog_privacy_mode, posthog_groups, - **kwargs, + kwargs, + start_time, ) - async def _create_streaming( + def _track_streaming_response( self, + response: Any, posthog_distinct_id: Optional[str], posthog_trace_id: Optional[str], posthog_properties: Optional[Dict[str, Any]], posthog_privacy_mode: bool, posthog_groups: Optional[Dict[str, Any]], - **kwargs: Any, + kwargs: Dict[str, Any], + start_time: float, ): - start_time = time.time() usage_stats: TokenUsage = TokenUsage(input_tokens=0, output_tokens=0) accumulated_content = "" content_blocks: List[StreamingContentBlock] = [] tools_in_progress: Dict[str, ToolInProgress] = {} current_text_block: Optional[StreamingContentBlock] = None stop_reason: Optional[str] = None - response = await super().create(**kwargs) async def generator(): nonlocal usage_stats diff --git a/posthog/ai/stream.py b/posthog/ai/stream.py index 4ed8ca94..2499f6e6 100644 --- a/posthog/ai/stream.py +++ b/posthog/ai/stream.py @@ -1,10 +1,47 @@ """Shared async streaming utilities for PostHog AI wrappers.""" -from typing import Any, AsyncGenerator, Generic, Optional, TypeVar +from typing import Any, AsyncGenerator, Generator, Generic, Optional, TypeVar T = TypeVar("T") +class _StreamWrapper(Generic[T]): + """Preserves a provider stream's helpers while tracking its iteration.""" + + def __init__( + self, + generator: Generator[T, None, None], + stream: Any, + ) -> None: + self._generator = generator + self._stream = stream + + def __iter__(self) -> "_StreamWrapper[T]": + return self + + def __next__(self) -> T: + return next(self._generator) + + def __enter__(self) -> "_StreamWrapper[T]": + return self + + def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: + self.close() + + def close(self) -> None: + try: + self._generator.close() + finally: + self._stream.close() + + def __getattr__(self, name: str) -> Any: + if name.startswith("_"): + raise AttributeError(name) + if name in ("send", "throw"): + return getattr(self._generator, name) + return getattr(self._stream, name) + + class AsyncStreamWrapper(Generic[T]): """Adds the async context manager protocol to a PostHog streaming generator. @@ -34,6 +71,10 @@ async def __aenter__(self) -> "AsyncStreamWrapper[T]": return self async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> bool: + await self.close() + return False + + async def close(self) -> None: # Close the generator first so its `finally` captures the event, even on # early exit. try/finally still closes the provider stream if that raises. try: @@ -46,11 +87,10 @@ async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> bool: if close is not None: await close() - return False + async def aclose(self) -> None: + await self.close() - # aclose/asend/athrow belong to the generator; provider streams expose - # close(), not these. Forwarding aclose() keeps it firing the event. - _GENERATOR_METHODS = ("aclose", "asend", "athrow") + _GENERATOR_METHODS = ("asend", "athrow") def __getattr__(self, name: str) -> Any: # Proxy only public attributes (e.g. `.response`) to the provider stream. diff --git a/posthog/test/ai/anthropic/test_anthropic.py b/posthog/test/ai/anthropic/test_anthropic.py index 62cdc24d..0fb840f8 100644 --- a/posthog/test/ai/anthropic/test_anthropic.py +++ b/posthog/test/ai/anthropic/test_anthropic.py @@ -1,13 +1,32 @@ +import inspect import json import os -from unittest.mock import patch +from unittest.mock import AsyncMock, Mock, patch import pytest from posthog import identify_context, new_context try: - from anthropic.types import CacheCreation, Message, Usage + from anthropic.lib.streaming._messages import ( + AsyncMessageStreamManager, + MessageStreamManager, + ) + from anthropic.types import ( + CacheCreation, + Message, + MessageDeltaUsage, + RawContentBlockDeltaEvent, + RawContentBlockStartEvent, + RawContentBlockStopEvent, + RawMessageDeltaEvent, + RawMessageStartEvent, + RawMessageStopEvent, + TextBlock, + TextDelta, + Usage, + ) + from anthropic.types.raw_message_delta_event import Delta from posthog.ai.anthropic import Anthropic, AnthropicBedrock, AsyncAnthropic from posthog.test.ai.utils import RecordingAsyncStream @@ -1744,6 +1763,22 @@ def test_integration_stop_reason(mock_client): assert props["$ai_input_tokens"] > 0 +class RecordingStream: + def __init__(self, items): + self._items = iter(items) + self.closed = False + self.response = "provider-response" + + def __iter__(self): + return self + + def __next__(self): + return next(self._items) + + def close(self): + self.closed = True + + def _anthropic_stream_events(): final = MockStreamEvent("message_delta") final.usage = MockUsage( @@ -1759,6 +1794,125 @@ def _anthropic_stream_events(): ] +def _anthropic_raw_stream_events(): + message = Message( + id="message-id", + type="message", + role="assistant", + content=[], + model="claude-3-opus-20240229", + usage=Usage(input_tokens=10, output_tokens=0), + stop_reason=None, + stop_sequence=None, + ) + return [ + RawMessageStartEvent(type="message_start", message=message), + RawContentBlockStartEvent( + type="content_block_start", + index=0, + content_block=TextBlock(type="text", text=""), + ), + RawContentBlockDeltaEvent( + type="content_block_delta", + index=0, + delta=TextDelta(type="text_delta", text="Hi"), + ), + RawContentBlockStopEvent(type="content_block_stop", index=0), + RawMessageDeltaEvent( + type="message_delta", + delta=Delta(stop_reason="end_turn", stop_sequence=None), + usage=MessageDeltaUsage(output_tokens=5), + ), + RawMessageStopEvent(type="message_stop"), + ] + + +def test_messages_stream_preserves_native_manager_helpers_close_and_tracking( + mock_client, +): + source = RecordingStream(_anthropic_raw_stream_events()) + client = Anthropic(api_key="test-key", posthog_client=mock_client) + client.post = Mock(return_value=source) + + response = client.messages.stream( + model="claude-haiku-4-5", + messages=[{"role": "user", "content": "Foo"}], + max_tokens=1, + posthog_distinct_id="test-user", + ) + + assert isinstance(response, MessageStreamManager) + with response as stream: + assert stream.response == "provider-response" + text = list(stream.text_stream) + + assert text == ["Hi"] + assert source.closed is True + assert mock_client.capture.call_count == 1 + assert mock_client.capture.call_args.kwargs["distinct_id"] == "test-user" + assert "posthog_distinct_id" not in client.post.call_args.kwargs + + +def test_messages_stream_tolerates_provider_manager_internals_changing(mock_client): + manager = object() + client = Anthropic(api_key="test-key", posthog_client=mock_client) + + with patch("anthropic.resources.messages.Messages.stream", return_value=manager): + response = client.messages.stream( + model="claude-haiku-4-5", + messages=[{"role": "user", "content": "Foo"}], + max_tokens=1, + ) + + assert response is manager + + +@pytest.mark.asyncio +async def test_async_messages_stream_preserves_provider_contract_and_manager( + mock_client, +): + source = RecordingAsyncStream(_anthropic_raw_stream_events()) + client = AsyncAnthropic(posthog_client=mock_client) + client.post = AsyncMock(return_value=source) + + response = client.messages.stream( + model="claude-haiku-4-5", + messages=[{"role": "user", "content": "Foo"}], + max_tokens=1, + posthog_distinct_id="test-user", + ) + + assert isinstance(response, AsyncMessageStreamManager) + assert not inspect.isawaitable(response) + async with response as stream: + assert stream.response == "provider-response" + text = [chunk async for chunk in stream.text_stream] + + assert text == ["Hi"] + assert source.closed is True + assert mock_client.capture.call_count == 1 + assert mock_client.capture.call_args.kwargs["distinct_id"] == "test-user" + assert "posthog_distinct_id" not in client.post.call_args.kwargs + + +def test_async_messages_stream_tolerates_provider_manager_internals_changing( + mock_client, +): + manager = object() + client = AsyncAnthropic(api_key="test-key", posthog_client=mock_client) + + with patch( + "anthropic.resources.messages.AsyncMessages.stream", return_value=manager + ): + response = client.messages.stream( + model="claude-haiku-4-5", + messages=[{"role": "user", "content": "Foo"}], + max_tokens=1, + ) + + assert response is manager + + @pytest.mark.asyncio async def test_async_messages_create_streaming_supports_async_with(mock_client): """Regression test for #393: messages.create(stream=True) must support diff --git a/references/public_api_snapshot.txt b/references/public_api_snapshot.txt index d83fbb74..67a22411 100644 --- a/references/public_api_snapshot.txt +++ b/references/public_api_snapshot.txt @@ -1228,6 +1228,8 @@ method posthog.ai.otel.processor.PostHogSpanProcessor.shutdown() -> None method posthog.ai.prompts.Prompts.clear_cache(name: Optional[str] = None, *, version: Optional[int] = None) -> None method posthog.ai.prompts.Prompts.compile(prompt: str, variables: PromptVariables) -> str method posthog.ai.prompts.Prompts.get(name: str, *, with_metadata: Optional[bool] = None, cache_ttl_seconds: Optional[int] = None, fallback: Optional[str] = None, version: Optional[int] = None, label: Optional[str] = None) -> Union[str, PromptResult] +method posthog.ai.stream.AsyncStreamWrapper.aclose() -> None +method posthog.ai.stream.AsyncStreamWrapper.close() -> None method posthog.bucketed_rate_limiter.BucketedRateLimiter.consume_rate_limit(key: Hashable) -> bool method posthog.bucketed_rate_limiter.BucketedRateLimiter.stop() -> None method posthog.client.Client.alias(previous_id: str, distinct_id: Optional[str], timestamp: Optional[Union[datetime, str]] = None, uuid: Optional[str] = None, disable_geoip: Optional[bool] = None) -> Optional[str]