Skip to content

Commit 4916988

Browse files
committed
refactor(ai): split PostHogClaudeSDKClient into client.py
1 parent 794a8f7 commit 4916988

3 files changed

Lines changed: 261 additions & 231 deletions

File tree

‎posthog/ai/claude_agent_sdk/__init__.py‎

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -14,10 +14,8 @@
1414
"Please install the Claude Agent SDK to use this feature: 'pip install claude-agent-sdk'"
1515
)
1616

17-
from posthog.ai.claude_agent_sdk.processor import (
18-
PostHogClaudeAgentProcessor,
19-
PostHogClaudeSDKClient,
20-
)
17+
from posthog.ai.claude_agent_sdk.client import PostHogClaudeSDKClient
18+
from posthog.ai.claude_agent_sdk.processor import PostHogClaudeAgentProcessor
2119

2220
__all__ = [
2321
"PostHogClaudeAgentProcessor",
Lines changed: 259 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,259 @@
1+
"""PostHog-instrumented ClaudeSDKClient for stateful multi-turn conversations.
2+
3+
Wraps claude_agent_sdk.ClaudeSDKClient to automatically emit $ai_generation,
4+
$ai_span, and $ai_trace events across multiple conversation turns.
5+
"""
6+
7+
import logging
8+
import time
9+
import uuid
10+
from typing import Any, Callable, Dict, List, Optional, Union
11+
12+
try:
13+
from claude_agent_sdk import (
14+
AssistantMessage,
15+
ClaudeSDKClient,
16+
ResultMessage,
17+
ToolUseBlock,
18+
UserMessage,
19+
)
20+
from claude_agent_sdk.types import ClaudeAgentOptions, StreamEvent
21+
except ImportError:
22+
raise ModuleNotFoundError(
23+
"Please install the Claude Agent SDK to use this feature: 'pip install claude-agent-sdk'"
24+
)
25+
26+
from posthog.ai.claude_agent_sdk.processor import (
27+
PostHogClaudeAgentProcessor,
28+
_GenerationTracker,
29+
_ensure_serializable,
30+
)
31+
from posthog.client import Client
32+
33+
log = logging.getLogger("posthog")
34+
35+
36+
class PostHogClaudeSDKClient:
37+
"""Wraps ClaudeSDKClient for stateful multi-turn conversations with PostHog instrumentation.
38+
39+
Usage:
40+
async with PostHogClaudeSDKClient(options, posthog_client=ph, posthog_distinct_id="user") as client:
41+
await client.query("Hello")
42+
async for msg in client.receive_response():
43+
... # turn 1, emits $ai_generation events
44+
await client.query("Follow up")
45+
async for msg in client.receive_response():
46+
... # turn 2, same trace, has conversation history
47+
"""
48+
49+
def __init__(
50+
self,
51+
options: Optional["ClaudeAgentOptions"] = None,
52+
transport: Any = None,
53+
*,
54+
posthog_client: Optional[Client] = None,
55+
posthog_distinct_id: Optional[
56+
Union[str, Callable[["ResultMessage"], Optional[str]]]
57+
] = None,
58+
posthog_trace_id: Optional[str] = None,
59+
posthog_properties: Optional[Dict[str, Any]] = None,
60+
posthog_privacy_mode: bool = False,
61+
posthog_groups: Optional[Dict[str, Any]] = None,
62+
):
63+
from dataclasses import replace as dc_replace
64+
65+
# Ensure partial messages for per-generation tracking
66+
if options is None:
67+
options = ClaudeAgentOptions(include_partial_messages=True)
68+
elif not options.include_partial_messages:
69+
options = dc_replace(options, include_partial_messages=True)
70+
71+
self._client = ClaudeSDKClient(options, transport)
72+
self._processor = PostHogClaudeAgentProcessor(
73+
client=posthog_client,
74+
distinct_id=posthog_distinct_id,
75+
privacy_mode=posthog_privacy_mode,
76+
groups=posthog_groups,
77+
properties=posthog_properties or {},
78+
)
79+
self._trace_id = posthog_trace_id or str(uuid.uuid4())
80+
self._distinct_id = posthog_distinct_id
81+
self._extra_props = posthog_properties or {}
82+
self._privacy = posthog_privacy_mode
83+
self._groups = posthog_groups or {}
84+
85+
# Shared state across turns
86+
self._tracker = _GenerationTracker()
87+
self._generation_index = 0
88+
self._current_generation_span_id: Optional[str] = None
89+
self._current_input: Optional[List[Dict[str, Any]]] = None
90+
self._next_input: Optional[List[Dict[str, Any]]] = None
91+
self._pending_output: List[Dict[str, Any]] = []
92+
self._query_start = time.time()
93+
94+
async def connect(self, prompt: Any = None) -> None:
95+
await self._client.connect(prompt)
96+
97+
async def query(self, prompt: str, session_id: str = "default") -> None:
98+
# Track the prompt as input for the next generation
99+
self._current_input = [{"role": "user", "content": prompt}]
100+
await self._client.query(prompt, session_id)
101+
102+
async def receive_response(self):
103+
"""Instrumented receive_response -- yields all messages, emits PostHog events."""
104+
async for message in self._client.receive_response():
105+
try:
106+
if isinstance(message, StreamEvent):
107+
self._tracker.process_stream_event(message)
108+
109+
if self._tracker.has_completed_generation():
110+
gen = self._tracker.pop_generation()
111+
self._generation_index += 1
112+
self._current_generation_span_id = gen.span_id
113+
self._processor._emit_generation(
114+
gen,
115+
self._trace_id,
116+
self._generation_index,
117+
self._current_input,
118+
self._pending_output or None,
119+
self._distinct_id,
120+
self._extra_props,
121+
self._privacy,
122+
self._groups,
123+
)
124+
self._current_input = self._next_input
125+
self._next_input = None
126+
self._pending_output = []
127+
128+
elif isinstance(message, AssistantMessage):
129+
self._tracker.set_model(message.model)
130+
parent_id = (
131+
self._tracker.current_span_id
132+
or self._current_generation_span_id
133+
)
134+
output_content: List[Dict[str, Any]] = []
135+
for block in message.content:
136+
if isinstance(block, ToolUseBlock):
137+
self._processor._emit_tool_span(
138+
block,
139+
self._trace_id,
140+
parent_id,
141+
self._distinct_id,
142+
self._extra_props,
143+
self._privacy,
144+
self._groups,
145+
)
146+
output_content.append(
147+
{
148+
"type": "function",
149+
"function": {
150+
"name": block.name,
151+
"arguments": block.input,
152+
},
153+
}
154+
)
155+
elif hasattr(block, "text"):
156+
output_content.append({"type": "text", "text": block.text})
157+
if output_content:
158+
self._pending_output = [
159+
{"role": "assistant", "content": output_content}
160+
]
161+
162+
elif isinstance(message, UserMessage):
163+
content = message.content
164+
if isinstance(content, str):
165+
self._next_input = [{"role": "user", "content": content}]
166+
elif isinstance(content, list):
167+
formatted: List[Dict[str, Any]] = []
168+
for block in content:
169+
if hasattr(block, "tool_use_id"):
170+
formatted.append(
171+
{
172+
"type": "tool_result",
173+
"tool_use_id": block.tool_use_id,
174+
"content": str(block.content)[:500]
175+
if block.content
176+
else None,
177+
}
178+
)
179+
elif hasattr(block, "text"):
180+
formatted.append({"type": "text", "text": block.text})
181+
if formatted:
182+
self._next_input = [{"role": "user", "content": formatted}]
183+
184+
elif isinstance(message, ResultMessage):
185+
if not self._tracker.had_any_stream_events:
186+
self._processor._emit_generation_from_result(
187+
message,
188+
self._trace_id,
189+
self._tracker.last_model,
190+
self._query_start,
191+
self._current_input,
192+
self._pending_output,
193+
self._distinct_id,
194+
self._extra_props,
195+
self._privacy,
196+
self._groups,
197+
)
198+
# Don't emit trace here -- wait for disconnect/close
199+
# so multi-turn sessions get one trace at the end
200+
201+
except Exception as e:
202+
log.debug(f"PostHog instrumentation error (non-fatal): {e}")
203+
204+
yield message
205+
206+
async def disconnect(self) -> None:
207+
# Emit the trace event covering the entire session
208+
try:
209+
latency = time.time() - self._query_start
210+
resolved_id = self._processor._resolve_distinct_id(self._distinct_id)
211+
212+
properties: Dict[str, Any] = {
213+
"$ai_trace_id": self._trace_id,
214+
"$ai_trace_name": "claude_agent_sdk_session",
215+
"$ai_provider": "anthropic",
216+
"$ai_framework": "claude-agent-sdk",
217+
"$ai_latency": latency,
218+
**self._extra_props,
219+
}
220+
221+
if resolved_id is None:
222+
properties["$process_person_profile"] = False
223+
224+
self._processor._capture_event(
225+
"$ai_trace",
226+
properties,
227+
resolved_id or self._trace_id,
228+
self._groups,
229+
)
230+
231+
try:
232+
ph = self._processor._client
233+
if hasattr(ph, "flush") and callable(ph.flush):
234+
ph.flush()
235+
except Exception as e:
236+
log.debug(f"Error flushing PostHog client: {e}")
237+
238+
except Exception as e:
239+
log.debug(f"PostHog trace emission error (non-fatal): {e}")
240+
241+
await self._client.disconnect()
242+
243+
# Delegate other methods
244+
async def interrupt(self) -> None:
245+
await self._client.interrupt()
246+
247+
async def set_permission_mode(self, mode: str) -> None:
248+
await self._client.set_permission_mode(mode)
249+
250+
async def set_model(self, model: Optional[str] = None) -> None:
251+
await self._client.set_model(model)
252+
253+
async def __aenter__(self) -> "PostHogClaudeSDKClient":
254+
await self.connect()
255+
return self
256+
257+
async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> bool:
258+
await self.disconnect()
259+
return False

0 commit comments

Comments
 (0)