diff --git a/.sampo/changesets/langchain-agent-middleware.md b/.sampo/changesets/langchain-agent-middleware.md new file mode 100644 index 000000000..db17390c8 --- /dev/null +++ b/.sampo/changesets/langchain-agent-middleware.md @@ -0,0 +1,5 @@ +--- +pypi/posthog: minor +--- + +Add LangChain v1 agent middleware for AI observability. diff --git a/examples/example-ai-langchain/README.md b/examples/example-ai-langchain/README.md index 5596a84ea..6e60f33e9 100644 --- a/examples/example-ai-langchain/README.md +++ b/examples/example-ai-langchain/README.md @@ -14,6 +14,7 @@ uv sync ## Examples - **callback_handler.py** - PostHog callback handler with tool calling +- **agent_middleware.py** - LangChain v1 agent middleware with model and tool tracking - **otel.py** - OpenTelemetry instrumentation exporting to PostHog ## Run @@ -21,5 +22,10 @@ uv sync ```bash source .env uv run python callback_handler.py +uv run python agent_middleware.py uv run python otel.py ``` + +Use either `PostHogMiddleware` or `CallbackHandler` for an agent invocation, not +both, to avoid recording duplicate events. Put `PostHogMiddleware` last in the +middleware list so it records the final model selection and each retry attempt. diff --git a/examples/example-ai-langchain/agent_middleware.py b/examples/example-ai-langchain/agent_middleware.py new file mode 100644 index 000000000..8b015e9b0 --- /dev/null +++ b/examples/example-ai-langchain/agent_middleware.py @@ -0,0 +1,34 @@ +"""Track a LangChain v1 agent with PostHog middleware.""" + +import os + +from langchain.agents import create_agent +from langchain_core.messages import HumanMessage +from langchain_core.tools import tool +from langchain_openai import ChatOpenAI + +from posthog import Posthog +from posthog.ai.langchain.middleware import PostHogMiddleware + + +@tool +def get_weather(city: str) -> str: + """Get the weather for a city.""" + return f"It is sunny in {city}." + + +client = Posthog( + os.environ["POSTHOG_API_KEY"], + host=os.environ.get("POSTHOG_HOST", "https://us.i.posthog.com"), +) +agent = create_agent( + model=ChatOpenAI(model="gpt-4.1-mini"), + tools=[get_weather], + middleware=[PostHogMiddleware(client, distinct_id="example-user")], +) + +result = agent.invoke( + {"messages": [HumanMessage(content="What is the weather in Berlin?")]} +) +print(result["messages"][-1].content) +client.shutdown() diff --git a/posthog/ai/langchain/callbacks.py b/posthog/ai/langchain/callbacks.py index c5fe44eb6..3b9e188b0 100644 --- a/posthog/ai/langchain/callbacks.py +++ b/posthog/ai/langchain/callbacks.py @@ -537,8 +537,11 @@ def _capture_trace_or_span( run: SpanMetadata, outputs: Any, parent_run_id: Optional[UUID], + event_name_override: Optional[str] = None, ): - event_name = "$ai_trace" if parent_run_id is None else "$ai_span" + event_name = event_name_override or ( + "$ai_trace" if parent_run_id is None else "$ai_span" + ) event_properties = { "$ai_trace_id": trace_id, "$ai_input_state": with_privacy_mode( @@ -616,6 +619,7 @@ def _capture_generation( run: GenerationMetadata, output: Union[LLMResult, BaseException], parent_run_id: Optional[UUID] = None, + include_parent_id: bool = True, ): # The served tier comes from the response, because a requested tier can be refused. model_params = run.model_params @@ -629,7 +633,6 @@ def _capture_generation( "$ai_trace_id": trace_id, "$ai_span_id": run_id, "$ai_span_name": run.name, - "$ai_parent_id": parent_run_id, "$ai_provider": run.provider, "$ai_model": run.model, "$ai_model_parameters": model_params, @@ -645,6 +648,8 @@ def _capture_generation( "$ai_base_url": run.base_url, "$ai_framework": "langchain", } + if include_parent_id: + event_properties["$ai_parent_id"] = parent_run_id warn_if_posthog_ai_gateway(run.base_url) diff --git a/posthog/ai/langchain/middleware.py b/posthog/ai/langchain/middleware.py new file mode 100644 index 000000000..34783747e --- /dev/null +++ b/posthog/ai/langchain/middleware.py @@ -0,0 +1,373 @@ +"""LangChain v1 agent middleware for PostHog AI observability.""" + +import asyncio +import logging +import time +from collections.abc import Awaitable, Callable, Mapping +from typing import Annotated, Any, Optional, Union +from uuid import UUID, uuid4 + +from langchain.agents.middleware.types import ( + AgentMiddleware, + AgentState, + ModelRequest, + ModelResponse, + PrivateStateAttr, + ToolCallRequest, +) +from langchain_core.messages import AIMessage, BaseMessage, ToolMessage +from langchain_core.outputs import ChatGeneration, LLMResult +from langchain_core.utils.function_calling import convert_to_openai_tool +from langgraph.errors import GraphBubbleUp +from langgraph.types import Command +from ...client import Client +from .callbacks import ( + CallbackHandler, + GenerationMetadata, + SpanMetadata, + _convert_message_to_dict, +) + +log = logging.getLogger("posthog") + + +class _PostHogMiddlewareState(AgentState[Any], total=False): # type: ignore[call-arg] + """Private state used to correlate one agent invocation.""" + + _posthog_root_id: Annotated[UUID | None, PrivateStateAttr] + _posthog_root_start_time: Annotated[float | None, PrivateStateAttr] + _posthog_root_input: Annotated[dict[str, Any] | None, PrivateStateAttr] + + +class PostHogMiddleware(AgentMiddleware[_PostHogMiddlewareState, Any, Any]): + """Capture LangChain v1 agent, model, and tool activity in PostHog. + + Use either this middleware or :class:`CallbackHandler` for an agent invocation, + not both, to avoid duplicate events. Put this middleware last in the middleware + list so it records the final model selection and each retry attempt. + """ + + state_schema = _PostHogMiddlewareState + + def __init__( + self, + client: Optional[Client] = None, + *, + distinct_id: Optional[Union[str, int, UUID]] = None, + trace_id: Optional[Union[str, int, float, UUID]] = None, + properties: Optional[dict[str, Any]] = None, + privacy_mode: bool = False, + groups: Optional[dict[str, Any]] = None, + ) -> None: + self._callback = CallbackHandler( + client, + distinct_id=distinct_id, + trace_id=trace_id, + properties=properties, + privacy_mode=privacy_mode, + groups=groups, + ) + + def before_agent( + self, state: _PostHogMiddlewareState, runtime: Any + ) -> dict[str, Any]: + return self._safely_call(self._start_agent, state) or {} + + async def abefore_agent( + self, state: _PostHogMiddlewareState, runtime: Any + ) -> dict[str, Any]: + return self._safely_call(self._start_agent, state) or {} + + def after_agent( + self, state: _PostHogMiddlewareState, runtime: Any + ) -> dict[str, Any]: + self._safely_call(self._finish_agent, state) + return self._clear_agent_state() + + async def aafter_agent( + self, state: _PostHogMiddlewareState, runtime: Any + ) -> dict[str, Any]: + await self._asafely_call(self._finish_agent, state) + return self._clear_agent_state() + + def wrap_model_call( + self, + request: ModelRequest[Any], + handler: Callable[[ModelRequest[Any]], ModelResponse[Any]], + ) -> ModelResponse[Any]: + run_id = self._start_model(request) + try: + response = handler(request) + except BaseException as error: + self._safely_call(self._finish_model, request.state, run_id, error, False) + raise + self._safely_call(self._finish_model, request.state, run_id, response, True) + return response + + async def awrap_model_call( + self, + request: ModelRequest[Any], + handler: Callable[[ModelRequest[Any]], Awaitable[ModelResponse[Any]]], + ) -> ModelResponse[Any]: + run_id = self._start_model(request) + try: + response = await handler(request) + except BaseException as error: + await self._asafely_call( + self._finish_model, request.state, run_id, error, False + ) + raise + await self._asafely_call( + self._finish_model, request.state, run_id, response, True + ) + return response + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], + ) -> ToolMessage | Command[Any]: + run_id = self._start_tool(request) + try: + response = handler(request) + except GraphBubbleUp: + if run_id is not None: + self._callback._runs.pop(run_id, None) + raise + except BaseException as error: + self._safely_call(self._finish_tool, request.state, run_id, error, False) + raise + output: Any = response + if isinstance(response, ToolMessage) and response.status == "error": + output = self._returned_tool_error(response) + self._safely_call(self._finish_tool, request.state, run_id, output, True) + return response + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]], + ) -> ToolMessage | Command[Any]: + run_id = self._start_tool(request) + try: + response = await handler(request) + except GraphBubbleUp: + if run_id is not None: + self._callback._runs.pop(run_id, None) + raise + except BaseException as error: + await self._asafely_call( + self._finish_tool, request.state, run_id, error, False + ) + raise + output: Any = response + if isinstance(response, ToolMessage) and response.status == "error": + output = self._returned_tool_error(response) + await self._asafely_call(self._finish_tool, request.state, run_id, output, True) + return response + + def _start_agent(self, state: Mapping[str, Any]) -> dict[str, Any]: + return { + "_posthog_root_id": uuid4(), + "_posthog_root_start_time": time.time(), + "_posthog_root_input": ( + None if self._privacy_mode_enabled() else _public_state(state) + ), + } + + @staticmethod + def _clear_agent_state() -> dict[str, Any]: + return { + "_posthog_root_id": None, + "_posthog_root_start_time": None, + "_posthog_root_input": None, + } + + def _finish_agent(self, state: Mapping[str, Any]) -> None: + root_id = state.get("_posthog_root_id") + start_time = state.get("_posthog_root_start_time") + root_input = state.get("_posthog_root_input") + if not isinstance(root_id, UUID) or not isinstance(start_time, (int, float)): + return + + run = SpanMetadata( + name="agent", + input=root_input, + start_time=start_time, + end_time=time.time(), + ) + self._safely_call( + self._callback._capture_trace_or_span, + self._trace_id(root_id), + root_id, + run, + _public_state(state), + None, + ) + + def _start_model(self, request: ModelRequest[Any]) -> UUID | None: + run_id = uuid4() + try: + messages: list[BaseMessage] = [] + if request.system_message is not None: + messages.append(request.system_message) + messages.extend(request.messages) + + model_parameters = dict(request.model_settings) + invocation_parameters = request.model._get_invocation_params( + **request.model_settings + ) + if isinstance(invocation_parameters, dict): + model_parameters = {**invocation_parameters, **model_parameters} + if request.tools: + model_parameters["tools"] = [ + self._normalize_tool(tool) for tool in request.tools + ] + + metadata = request.model._get_ls_params(**request.model_settings) + self._callback._set_llm_metadata( + request.model.to_json(), + run_id, + [_convert_message_to_dict(message) for message in messages], + metadata=metadata if isinstance(metadata, dict) else None, + invocation_params=model_parameters, + ) + except Exception: + log.exception("Failed to prepare PostHog LangChain model telemetry") + self._callback._runs.pop(run_id, None) + return None + return run_id + + def _finish_model( + self, + state: Mapping[str, Any], + run_id: UUID | None, + result: ModelResponse[Any] | BaseException, + include_parent: bool, + ) -> None: + if run_id is None: + return + run = self._callback._pop_run_metadata(run_id) + if not isinstance(run, GenerationMetadata): + return + + output: LLMResult | BaseException + if isinstance(result, BaseException): + output = result + else: + generations = [ + ChatGeneration(message=message) + for message in result.result + if isinstance(message, AIMessage) + ] + if not generations: + return + metadata = generations[-1].message.response_metadata + output = LLMResult( + generations=[generations], + llm_output=metadata if isinstance(metadata, dict) else None, + ) + + root_id = _root_id(state) + self._safely_call( + self._callback._capture_generation, + self._trace_id(root_id or run_id), + run_id, + run, + output, + root_id if include_parent else None, + include_parent, + ) + + def _start_tool(self, request: ToolCallRequest) -> UUID | None: + run_id = uuid4() + try: + name = str(request.tool_call.get("name") or "tool") + serialized = request.tool.to_json() if request.tool is not None else None + self._callback._set_trace_or_span_metadata( + serialized, + request.tool_call.get("args"), + run_id, + _root_id(request.state), + name=name, + ) + except Exception: + log.exception("Failed to prepare PostHog LangChain tool telemetry") + self._callback._runs.pop(run_id, None) + return None + return run_id + + def _finish_tool( + self, + state: Mapping[str, Any], + run_id: UUID | None, + output: Any, + include_parent: bool, + ) -> None: + if run_id is None: + return + run = self._callback._pop_run_metadata(run_id) + if not isinstance(run, SpanMetadata) or isinstance(run, GenerationMetadata): + return + root_id = _root_id(state) + self._safely_call( + self._callback._capture_trace_or_span, + self._trace_id(root_id or run_id), + run_id, + run, + output, + root_id if include_parent else None, + "$ai_span", + ) + + def _trace_id(self, fallback: UUID) -> Any: + return self._callback._trace_id or fallback + + def _privacy_mode_enabled(self) -> bool: + return bool( + self._callback._privacy_mode + or getattr(self._callback._ph_client, "privacy_mode", False) + ) + + def _returned_tool_error(self, response: ToolMessage) -> RuntimeError: + message = ( + "Tool returned an error" + if self._privacy_mode_enabled() + else str(response.content) + ) + return RuntimeError(message) + + @staticmethod + def _normalize_tool(tool: Any) -> dict[str, Any]: + try: + return convert_to_openai_tool(tool) + except Exception: + if isinstance(tool, dict): + return tool + return {"name": getattr(tool, "name", tool.__class__.__name__)} + + @staticmethod + def _safely_call(function: Callable[..., Any], *args: Any) -> Any: + try: + return function(*args) + except Exception: + log.exception("Failed to capture PostHog LangChain telemetry") + + async def _asafely_call(self, function: Callable[..., Any], *args: Any) -> Any: + if getattr(self._callback._ph_client, "sync_mode", False): + return await asyncio.to_thread(self._safely_call, function, *args) + return self._safely_call(function, *args) + + +def _root_id(state: Mapping[str, Any]) -> UUID | None: + root_id = state.get("_posthog_root_id") + return root_id if isinstance(root_id, UUID) else None + + +def _public_state(state: Mapping[str, Any]) -> dict[str, Any]: + return { + key: value for key, value in state.items() if not key.startswith("_posthog_") + } + + +__all__ = ["PostHogMiddleware"] diff --git a/posthog/test/ai/langchain/test_middleware.py b/posthog/test/ai/langchain/test_middleware.py new file mode 100644 index 000000000..ed68bca85 --- /dev/null +++ b/posthog/test/ai/langchain/test_middleware.py @@ -0,0 +1,893 @@ +import asyncio +import subprocess +import sys +import threading +import time +from typing import Any, Callable, Sequence +from unittest.mock import MagicMock + +import pytest +from pydantic import PrivateAttr + +from langchain.agents import AgentState, create_agent +from langchain.agents.middleware import ModelRetryMiddleware +from langchain.agents.middleware.types import ( + ModelRequest, + ModelResponse, + ToolCallRequest, +) +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage +from langchain_core.outputs import ChatGeneration, ChatResult +from langchain_core.tools import BaseTool, tool +from langgraph.checkpoint.memory import InMemorySaver +from langgraph.types import Command, interrupt + +import posthog.ai.langchain.middleware as middleware_module +from posthog.ai.langchain import CallbackHandler +from posthog.ai.langchain.middleware import PostHogMiddleware + + +class StubAgentModel(BaseChatModel): + responses: list[AIMessage] + _response_index: int = PrivateAttr(default=0) + + @property + def _llm_type(self) -> str: + return "posthog-test-agent-model" + + def bind_tools( + self, + tools: Sequence[dict[str, Any] | type | Callable[..., Any] | BaseTool], + *, + tool_choice: str | None = None, + **kwargs: Any, + ) -> BaseChatModel: + return self + + def _generate( + self, + messages: list[Any], + stop: list[str] | None = None, + run_manager: Any = None, + **kwargs: Any, + ) -> ChatResult: + response = self.responses[self._response_index] + self._response_index += 1 + return ChatResult(generations=[ChatGeneration(message=response)]) + + def _get_ls_params(self, stop: list[str] | None = None, **kwargs: Any) -> Any: + return {"ls_model_name": "test-model", "ls_provider": "test-provider"} + + def _get_invocation_params( + self, stop: list[str] | None = None, **kwargs: Any + ) -> dict[str, Any]: + return {"temperature": 0.25, **kwargs} + + +class FailingAgentModel(StubAgentModel): + error: Exception + + def _generate( + self, + messages: list[Any], + stop: list[str] | None = None, + run_manager: Any = None, + **kwargs: Any, + ) -> ChatResult: + raise self.error + + +class RetryOnceAgentModel(StubAgentModel): + _attempt: int = PrivateAttr(default=0) + + def _generate( + self, + messages: list[Any], + stop: list[str] | None = None, + run_manager: Any = None, + **kwargs: Any, + ) -> ChatResult: + self._attempt += 1 + if self._attempt == 1: + raise RuntimeError("retryable failure") + return super()._generate(messages, stop, run_manager, **kwargs) + + +class CustomAgentState(AgentState): + tenant_id: str + workflow: str + + +@tool +def weather(city: str) -> str: + """Get the weather for a city.""" + return f"Sunny in {city}" + + +@tool +def request_approval() -> str: + """Request approval before continuing.""" + approved = interrupt("Approve?") + return "approved" if approved else "rejected" + + +def _client() -> MagicMock: + client = MagicMock() + client.privacy_mode = False + client.sync_mode = False + client.enable_exception_autocapture = False + return client + + +def _events(client: MagicMock) -> list[dict[str, Any]]: + return [call.kwargs for call in client.capture.call_args_list] + + +def _event(client: MagicMock, event_name: str) -> dict[str, Any]: + return next(event for event in _events(client) if event["event"] == event_name) + + +def _middleware_state(middleware: PostHogMiddleware) -> dict[str, Any]: + state = {"messages": [HumanMessage(content="Hello")]} + update = middleware.before_agent(state, None) + return {**state, **(update or {})} + + +def _model_request( + model: BaseChatModel, + state: dict[str, Any], + *, + tools: list[BaseTool | dict[str, Any]] | None = None, +) -> ModelRequest[Any]: + return ModelRequest( + model=model, + messages=[HumanMessage(content="Hello")], + system_message=SystemMessage(content="Be concise"), + tools=tools, + state=state, + runtime=None, + model_settings={"max_tokens": 100}, + ) + + +def _tool_request(state: dict[str, Any]) -> ToolCallRequest: + return ToolCallRequest( + tool_call={ + "id": "weather-call", + "name": "weather", + "args": {"city": "London"}, + "type": "tool_call", + }, + tool=weather, + state=state, + runtime=None, + ) + + +def test_instruments_sync_agent_model_and_tool_loop() -> None: + client = _client() + agent = create_agent( + model=StubAgentModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + { + "id": "weather-call", + "name": "weather", + "args": {"city": "London"}, + } + ], + ), + AIMessage( + content="It is sunny", + usage_metadata={ + "input_tokens": 10, + "output_tokens": 4, + "total_tokens": 14, + }, + response_metadata={"finish_reason": "stop"}, + ), + ] + ), + tools=[weather], + middleware=[ + PostHogMiddleware( + client, + distinct_id="user-id", + trace_id="trace-id", + properties={"environment": "test"}, + groups={"company": "posthog"}, + ) + ], + ) + + result = agent.invoke({"messages": [HumanMessage(content="Hello")]}) + + assert result["messages"][-1].content == "It is sunny" + assert not any(key.startswith("_posthog_") for key in result) + events = _events(client) + assert [event["event"] for event in events] == [ + "$ai_generation", + "$ai_span", + "$ai_generation", + "$ai_trace", + ] + trace = events[-1] + assert trace["distinct_id"] == "user-id" + assert trace["groups"] == {"company": "posthog"} + assert all(event["properties"]["$ai_trace_id"] == "trace-id" for event in events) + assert all(event["properties"]["environment"] == "test" for event in events) + assert all( + event["properties"].get("$ai_parent_id") == trace["properties"]["$ai_span_id"] + for event in events[:-1] + ) + + generations = [event for event in events if event["event"] == "$ai_generation"] + assert generations[0]["properties"]["$ai_tools"] == [ + { + "type": "function", + "function": { + "name": "weather", + "description": "Get the weather for a city.", + "parameters": { + "properties": {"city": {"type": "string"}}, + "required": ["city"], + "type": "object", + }, + }, + } + ] + assert generations[-1]["properties"]["$ai_model"] == "test-model" + assert generations[-1]["properties"]["$ai_provider"] == "test-provider" + assert generations[-1]["properties"]["$ai_input_tokens"] == 10 + assert generations[-1]["properties"]["$ai_output_tokens"] == 4 + assert generations[-1]["properties"]["$ai_stop_reason"] == "stop" + + +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +@pytest.mark.asyncio +async def test_interrupt_resume_does_not_capture_tool_error( + asynchronous: bool, +) -> None: + client = _client() + client.enable_exception_autocapture = True + middleware = PostHogMiddleware(client) + agent = create_agent( + model=StubAgentModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + { + "id": "approval-call", + "name": "request_approval", + "args": {}, + } + ], + ), + AIMessage(content="Done"), + ] + ), + tools=[request_approval], + middleware=[middleware], + checkpointer=InMemorySaver(), + ) + config = {"configurable": {"thread_id": f"interrupt-{asynchronous}"}} + agent_input = {"messages": [HumanMessage(content="Continue the task")]} + + if asynchronous: + paused = await agent.ainvoke(agent_input, config=config) + else: + paused = agent.invoke(agent_input, config=config) + + assert paused["__interrupt__"][0].value == "Approve?" + assert [event["event"] for event in _events(client)] == ["$ai_generation"] + client.capture_exception.assert_not_called() + assert middleware._callback._runs == {} + + if asynchronous: + completed = await agent.ainvoke(Command(resume=True), config=config) + else: + completed = agent.invoke(Command(resume=True), config=config) + + assert completed["messages"][-1].content == "Done" + events = _events(client) + assert [event["event"] for event in events] == [ + "$ai_generation", + "$ai_span", + "$ai_generation", + "$ai_trace", + ] + assert all("$ai_is_error" not in event["properties"] for event in events) + assert all("$ai_error" not in event["properties"] for event in events) + client.capture_exception.assert_not_called() + assert middleware._callback._runs == {} + + +@pytest.mark.asyncio +async def test_instruments_async_agent_and_preserves_custom_state() -> None: + client = _client() + agent = create_agent( + model=StubAgentModel( + responses=[ + AIMessage( + content="Done", + usage_metadata={ + "input_tokens": 2, + "output_tokens": 1, + "total_tokens": 3, + }, + ) + ] + ), + tools=[], + middleware=[PostHogMiddleware(client)], + state_schema=CustomAgentState, + ) + + result = await agent.ainvoke( + { + "messages": [HumanMessage(content="Hello")], + "tenant_id": "tenant-123", + "workflow": "support", + } + ) + + assert result["tenant_id"] == "tenant-123" + assert result["workflow"] == "support" + assert not any(key.startswith("_posthog_") for key in result) + trace = _event(client, "$ai_trace")["properties"] + assert trace["$ai_input_state"]["tenant_id"] == "tenant-123" + assert trace["$ai_input_state"]["workflow"] == "support" + assert trace["$ai_output_state"]["tenant_id"] == "tenant-123" + assert trace["$ai_output_state"]["workflow"] == "support" + assert not any(key.startswith("_posthog_") for key in trace["$ai_input_state"]) + assert not any(key.startswith("_posthog_") for key in trace["$ai_output_state"]) + + +@pytest.mark.asyncio +async def test_async_hooks_do_not_capture_on_the_event_loop_thread() -> None: + client = _client() + client.sync_mode = True + event_loop_thread = threading.get_ident() + capture_threads: list[int] = [] + client.capture.side_effect = lambda **_: capture_threads.append( + threading.get_ident() + ) + middleware = PostHogMiddleware(client) + state = {"messages": [HumanMessage(content="Hello")]} + state.update(await middleware.abefore_agent(state, None)) + + model_request = _model_request(StubAgentModel(responses=[]), state) + + async def model_handler(_: ModelRequest[Any]) -> ModelResponse[Any]: + return ModelResponse(result=[AIMessage(content="Done")]) + + await middleware.awrap_model_call(model_request, model_handler) + + tool_request = _tool_request(state) + + async def tool_handler(_: ToolCallRequest) -> ToolMessage: + return ToolMessage( + content="Sunny", + tool_call_id="weather-call", + name="weather", + ) + + await middleware.awrap_tool_call(tool_request, tool_handler) + await middleware.aafter_agent(state, None) + + assert len(capture_threads) == 3 + assert all(thread_id != event_loop_thread for thread_id in capture_threads) + + +@pytest.mark.parametrize("privacy_source", ["middleware", "client"]) +def test_completed_private_state_is_redacted_and_cleared_from_checkpoint( + privacy_source: str, +) -> None: + client = _client() + client.privacy_mode = privacy_source == "client" + checkpointer = InMemorySaver() + agent = create_agent( + model=StubAgentModel(responses=[AIMessage(content="Done")]), + tools=[], + middleware=[ + PostHogMiddleware( + client, + privacy_mode=privacy_source == "middleware", + ) + ], + checkpointer=checkpointer, + ) + secret = "customer-secret-agent-input" + config = {"configurable": {"thread_id": f"privacy-{privacy_source}"}} + + result = agent.invoke( + {"messages": [HumanMessage(content=secret)]}, + config=config, + ) + + assert not any(key.startswith("_posthog_") for key in result) + checkpoints = list(checkpointer.list(config)) + assert checkpoints + for checkpoint in checkpoints: + root_input = checkpoint.checkpoint["channel_values"].get("_posthog_root_input") + assert secret not in repr(root_input) + + latest = checkpointer.get_tuple(config) + assert latest is not None + values = latest.checkpoint["channel_values"] + for key in ( + "_posthog_root_id", + "_posthog_root_start_time", + "_posthog_root_input", + ): + assert values.get(key) is None + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_model_wrapper_preserves_exact_response_and_error( + asynchronous: bool, +) -> None: + client = _client() + middleware = PostHogMiddleware(client) + state = _middleware_state(middleware) + request = _model_request(StubAgentModel(responses=[]), state, tools=[weather]) + response = ModelResponse( + result=[ + AIMessage( + content="Done", + usage_metadata={ + "input_tokens": 3, + "output_tokens": 2, + "total_tokens": 5, + }, + ) + ] + ) + + if asynchronous: + + async def handler(inner_request: ModelRequest[Any]) -> ModelResponse[Any]: + assert inner_request is request + return response + + actual = await middleware.awrap_model_call(request, handler) + else: + + def handler(inner_request: ModelRequest[Any]) -> ModelResponse[Any]: + assert inner_request is request + return response + + actual = middleware.wrap_model_call(request, handler) + + assert actual is response + generation = _event(client, "$ai_generation")["properties"] + assert generation["$ai_input"] == [ + {"role": "system", "content": "Be concise"}, + {"role": "user", "content": "Hello"}, + ] + assert generation["$ai_model_parameters"]["max_tokens"] == 100 + assert generation["$ai_model_parameters"]["temperature"] == 0.25 + + client.reset_mock() + error = RuntimeError("model failed") + if asynchronous: + + async def failing_handler( + inner_request: ModelRequest[Any], + ) -> ModelResponse[Any]: + raise error + + with pytest.raises(RuntimeError) as raised: + await middleware.awrap_model_call(request, failing_handler) + else: + + def failing_handler(inner_request: ModelRequest[Any]) -> ModelResponse[Any]: + raise error + + with pytest.raises(RuntimeError) as raised: + middleware.wrap_model_call(request, failing_handler) + + assert raised.value is error + failed_generation = _event(client, "$ai_generation")["properties"] + assert failed_generation["$ai_is_error"] is True + assert "model failed" in failed_generation["$ai_error"] + assert "$ai_parent_id" not in failed_generation + + +def test_model_tool_normalization_falls_back_per_tool( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = _client() + middleware = PostHogMiddleware(client) + state = _middleware_state(middleware) + fallback_tool = { + "type": "function", + "function": { + "name": "fallback_tool", + "description": "Already serialized", + "parameters": {"type": "object", "properties": {}}, + }, + } + request = _model_request( + StubAgentModel(responses=[]), + state, + tools=[weather, fallback_tool], + ) + real_converter = middleware_module.convert_to_openai_tool + + def flaky_converter(candidate: Any) -> dict[str, Any]: + if candidate is fallback_tool: + raise ValueError("unsupported tool") + return real_converter(candidate) + + monkeypatch.setattr( + middleware_module, + "convert_to_openai_tool", + flaky_converter, + ) + response = ModelResponse(result=[AIMessage(content="Done")]) + + assert middleware.wrap_model_call(request, lambda _: response) is response + + generation = _event(client, "$ai_generation")["properties"] + assert generation["$ai_tools"][0]["function"]["name"] == "weather" + assert generation["$ai_tools"][1] is fallback_tool + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_tool_wrapper_captures_returned_and_raised_errors( + asynchronous: bool, +) -> None: + client = _client() + middleware = PostHogMiddleware(client) + state = _middleware_state(middleware) + request = _tool_request(state) + response = ToolMessage( + content="Invalid arguments", + tool_call_id="weather-call", + name="weather", + status="error", + ) + + if asynchronous: + + async def handler(inner_request: ToolCallRequest) -> ToolMessage: + assert inner_request is request + return response + + actual = await middleware.awrap_tool_call(request, handler) + else: + + def handler(inner_request: ToolCallRequest) -> ToolMessage: + assert inner_request is request + return response + + actual = middleware.wrap_tool_call(request, handler) + + assert actual is response + span = _event(client, "$ai_span")["properties"] + assert span["$ai_is_error"] is True + assert "Invalid arguments" in span["$ai_error"] + assert "$ai_parent_id" in span + + client.reset_mock() + error = RuntimeError("tool failed") + if asynchronous: + + async def failing_handler(inner_request: ToolCallRequest) -> ToolMessage: + raise error + + with pytest.raises(RuntimeError) as raised: + await middleware.awrap_tool_call(request, failing_handler) + else: + + def failing_handler(inner_request: ToolCallRequest) -> ToolMessage: + raise error + + with pytest.raises(RuntimeError) as raised: + middleware.wrap_tool_call(request, failing_handler) + + assert raised.value is error + failed_span = _event(client, "$ai_span")["properties"] + assert failed_span["$ai_is_error"] is True + assert "tool failed" in failed_span["$ai_error"] + assert "$ai_parent_id" not in failed_span + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("privacy_source", ["middleware", "client"]) +@pytest.mark.asyncio +async def test_returned_tool_error_respects_privacy_mode( + asynchronous: bool, + privacy_source: str, +) -> None: + client = _client() + client.privacy_mode = privacy_source == "client" + middleware = PostHogMiddleware( + client, + privacy_mode=privacy_source == "middleware", + ) + state = _middleware_state(middleware) + request = _tool_request(state) + secret = "customer-secret-tool-error" + response = ToolMessage( + content=secret, + tool_call_id="weather-call", + name="weather", + status="error", + ) + + if asynchronous: + + async def handler(inner_request: ToolCallRequest) -> ToolMessage: + assert inner_request is request + return response + + actual = await middleware.awrap_tool_call(request, handler) + else: + + def handler(inner_request: ToolCallRequest) -> ToolMessage: + assert inner_request is request + return response + + actual = middleware.wrap_tool_call(request, handler) + + assert actual is response + span = _event(client, "$ai_span")["properties"] + assert span["$ai_is_error"] is True + assert secret not in str(span.get("$ai_error", "")) + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_tool_wrapper_preserves_command_result(asynchronous: bool) -> None: + client = _client() + middleware = PostHogMiddleware(client) + state = _middleware_state(middleware) + request = ToolCallRequest( + tool_call={ + "id": "weather-call", + "name": "weather", + "args": {"city": "London"}, + "type": "tool_call", + }, + tool=None, + state=state, + runtime=None, + ) + result = Command( + update={ + "messages": [ + ToolMessage( + content="Sunny", + tool_call_id="weather-call", + name="weather", + ) + ] + } + ) + + if asynchronous: + + async def handler(inner_request: ToolCallRequest) -> Command[Any]: + assert inner_request is request + return result + + actual = await middleware.awrap_tool_call(request, handler) + else: + + def handler(inner_request: ToolCallRequest) -> Command[Any]: + assert inner_request is request + return result + + actual = middleware.wrap_tool_call(request, handler) + + assert actual is result + span = _event(client, "$ai_span")["properties"] + assert "$ai_is_error" not in span + assert span["$ai_span_name"] == "weather" + assert span["$ai_parent_id"] + + +def test_terminal_model_error_has_no_dangling_root_trace() -> None: + client = _client() + error = RuntimeError("model failed") + agent = create_agent( + model=FailingAgentModel(responses=[], error=error), + tools=[], + middleware=[PostHogMiddleware(client)], + ) + + with pytest.raises(RuntimeError) as raised: + agent.invoke({"messages": [HumanMessage(content="Hello")]}) + + assert raised.value is error + events = _events(client) + assert [event["event"] for event in events] == ["$ai_generation"] + assert "$ai_parent_id" not in events[0]["properties"] + + +def test_recovered_model_attempt_is_linked_to_completed_root() -> None: + client = _client() + middleware = PostHogMiddleware(client) + state = _middleware_state(middleware) + request = _model_request(StubAgentModel(responses=[]), state) + error = RuntimeError("retryable failure") + + def fail(_: ModelRequest[Any]) -> ModelResponse[Any]: + raise error + + with pytest.raises(RuntimeError) as raised: + middleware.wrap_model_call(request, fail) + assert raised.value is error + + response = ModelResponse(result=[AIMessage(content="Recovered")]) + assert middleware.wrap_model_call(request, lambda _: response) is response + middleware.after_agent(state, None) + + failed_generation, successful_generation, trace = _events(client) + root_id = trace["properties"]["$ai_span_id"] + assert failed_generation["event"] == "$ai_generation" + assert failed_generation["properties"]["$ai_is_error"] is True + assert "$ai_parent_id" not in failed_generation["properties"] + assert successful_generation["event"] == "$ai_generation" + assert successful_generation["properties"]["$ai_parent_id"] == root_id + assert trace["event"] == "$ai_trace" + assert all( + event["properties"]["$ai_trace_id"] == trace["properties"]["$ai_trace_id"] + for event in (failed_generation, successful_generation, trace) + ) + + +def test_model_retry_captures_each_attempt_when_posthog_is_innermost() -> None: + client = _client() + agent = create_agent( + model=RetryOnceAgentModel(responses=[AIMessage(content="Recovered")]), + tools=[], + middleware=[ + ModelRetryMiddleware( + max_retries=1, + retry_on=(RuntimeError,), + initial_delay=0, + jitter=False, + ), + PostHogMiddleware(client), + ], + ) + + result = agent.invoke({"messages": [HumanMessage(content="Hello")]}) + + assert result["messages"][-1].content == "Recovered" + failed_generation, successful_generation, trace = _events(client) + assert failed_generation["event"] == "$ai_generation" + assert failed_generation["properties"]["$ai_is_error"] is True + assert "$ai_parent_id" not in failed_generation["properties"] + assert successful_generation["event"] == "$ai_generation" + assert ( + successful_generation["properties"]["$ai_parent_id"] + == trace["properties"]["$ai_span_id"] + ) + assert trace["event"] == "$ai_trace" + assert all( + event["properties"]["$ai_trace_id"] == trace["properties"]["$ai_trace_id"] + for event in (failed_generation, successful_generation, trace) + ) + + +def test_concurrent_invocations_keep_unique_roots_with_explicit_trace_id() -> None: + client = _client() + middleware = PostHogMiddleware(client, trace_id="shared-trace") + states = [_middleware_state(middleware), _middleware_state(middleware)] + requests = [_model_request(StubAgentModel(responses=[]), state) for state in states] + + async def run( + request: ModelRequest[Any], state: dict[str, Any], content: str + ) -> None: + async def handler(inner_request: ModelRequest[Any]) -> ModelResponse[Any]: + await asyncio.sleep(0) + return ModelResponse(result=[AIMessage(content=content)]) + + await middleware.awrap_model_call(request, handler) + middleware.after_agent(state, None) + + async def run_concurrently() -> None: + await asyncio.gather( + run(requests[0], states[0], "A"), + run(requests[1], states[1], "B"), + ) + + asyncio.run(run_concurrently()) + + events = _events(client) + traces = [event for event in events if event["event"] == "$ai_trace"] + generations = [event for event in events if event["event"] == "$ai_generation"] + assert len(traces) == 2 + assert len(generations) == 2 + assert all( + event["properties"]["$ai_trace_id"] == "shared-trace" for event in events + ) + root_ids = {trace["properties"]["$ai_span_id"] for trace in traces} + assert len(root_ids) == 2 + assert { + generation["properties"]["$ai_parent_id"] for generation in generations + } == root_ids + + +def test_privacy_mode_redacts_agent_and_model_content() -> None: + client = _client() + middleware = PostHogMiddleware(client, privacy_mode=True) + state = _middleware_state(middleware) + request = _model_request(StubAgentModel(responses=[]), state) + + middleware.wrap_model_call( + request, + lambda _: ModelResponse(result=[AIMessage(content="private output")]), + ) + middleware.after_agent(state, None) + + generation = _event(client, "$ai_generation")["properties"] + trace = _event(client, "$ai_trace")["properties"] + assert generation["$ai_input"] is None + assert generation["$ai_output_choices"] is None + assert trace["$ai_input_state"] is None + assert trace["$ai_output_state"] is None + + +def test_capture_failure_never_changes_agent_result() -> None: + client = _client() + client.capture.side_effect = RuntimeError("telemetry unavailable") + middleware = PostHogMiddleware(client) + state = _middleware_state(middleware) + request = _model_request(StubAgentModel(responses=[]), state) + response = ModelResponse(result=[AIMessage(content="Done")]) + + assert middleware.wrap_model_call(request, lambda _: response) is response + cleanup = middleware.after_agent(state, None) + assert cleanup is not None + assert cleanup + assert all(key.startswith("_posthog_") for key in cleanup) + assert all(value is None for value in cleanup.values()) + + +def test_root_latency_starts_before_agent_execution( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = _client() + middleware = PostHogMiddleware(client) + timestamps = iter([100.0, 103.0]) + monkeypatch.setattr(time, "time", lambda: next(timestamps)) + + state = _middleware_state(middleware) + middleware.after_agent(state, None) + + trace = _event(client, "$ai_trace")["properties"] + assert trace["$ai_latency"] == 3.0 + + +def test_callback_integration_remains_available() -> None: + assert CallbackHandler.__module__ == "posthog.ai.langchain.callbacks" + + +def test_callback_only_import_does_not_require_langchain_package() -> None: + script = """ +import builtins + +original_import = builtins.__import__ + +def import_without_langchain(name, *args, **kwargs): + if name == "langchain" or name.startswith("langchain."): + raise AssertionError(f"unexpected top-level LangChain import: {name}") + return original_import(name, *args, **kwargs) + +builtins.__import__ = import_without_langchain +from posthog.ai.langchain import CallbackHandler +assert CallbackHandler.__module__ == "posthog.ai.langchain.callbacks" +""" + + subprocess.run([sys.executable, "-c", script], check=True) diff --git a/pyproject.toml b/pyproject.toml index 12c546325..fb6875ba2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -83,6 +83,7 @@ test = [ "langgraph-checkpoint>=4.1.1", "tiktoken>=0.12.0", "langchain-core>=1.0", + "langchain>=1.3.9", "langchain-community>=0.4", "langchain-openai>=1.1.14", "langchain-anthropic>=1.0", @@ -148,4 +149,3 @@ pytest_add_cli_args = ["--timeout=30"] dev = [ "claude-agent-sdk>=0.1.50", ] - diff --git a/references/public_api_snapshot.txt b/references/public_api_snapshot.txt index 26c37dc44..36b28e434 100644 --- a/references/public_api_snapshot.txt +++ b/references/public_api_snapshot.txt @@ -416,6 +416,7 @@ attribute posthog.ai.langchain.callbacks.SpanMetadata.latency: float attribute posthog.ai.langchain.callbacks.SpanMetadata.name: str attribute posthog.ai.langchain.callbacks.SpanMetadata.start_time: float attribute posthog.ai.langchain.callbacks.log = logging.getLogger('posthog') +attribute posthog.ai.langchain.middleware.PostHogMiddleware.state_schema = _PostHogMiddlewareState attribute posthog.ai.openai.openai.OpenAI.beta: WrappedBeta attribute posthog.ai.openai.openai.OpenAI.chat: WrappedChat attribute posthog.ai.openai.openai.OpenAI.embeddings: WrappedEmbeddings @@ -966,6 +967,7 @@ class posthog.ai.langchain.callbacks.CallbackHandler(client: Optional[Client] = class posthog.ai.langchain.callbacks.GenerationMetadata(name: str, start_time: float, end_time: Optional[float], input: Optional[Any], provider: Optional[str] = None, model: Optional[str] = None, model_params: Optional[Dict[str, Any]] = None, base_url: Optional[str] = None, tools: Optional[List[Dict[str, Any]]] = None, posthog_properties: Optional[Dict[str, Any]] = None) class posthog.ai.langchain.callbacks.ModelUsage(input_tokens: Optional[int], output_tokens: Optional[int], cache_write_tokens: Optional[int], cache_read_tokens: Optional[int], reasoning_tokens: Optional[int], cache_write_5m_tokens: Optional[int] = None, cache_write_1h_tokens: Optional[int] = None) class posthog.ai.langchain.callbacks.SpanMetadata(name: str, start_time: float, end_time: Optional[float], input: Optional[Any]) +class posthog.ai.langchain.middleware.PostHogMiddleware(client: Optional[Client] = None, *, distinct_id: Optional[Union[str, int, UUID]] = None, trace_id: Optional[Union[str, int, float, UUID]] = None, properties: Optional[dict[str, Any]] = None, privacy_mode: bool = False, groups: Optional[dict[str, Any]] = None) class posthog.ai.openai.openai.OpenAI(posthog_client: Optional[PostHogClient] = None, **kwargs) class posthog.ai.openai.openai.WrappedBeta class posthog.ai.openai.openai.WrappedBetaChat @@ -1310,6 +1312,14 @@ method posthog.ai.langchain.callbacks.CallbackHandler.on_retriever_start(seriali method posthog.ai.langchain.callbacks.CallbackHandler.on_tool_end(output: str, *, run_id: UUID, parent_run_id: Optional[UUID] = None, **kwargs: Any) -> Any method posthog.ai.langchain.callbacks.CallbackHandler.on_tool_error(error: BaseException, *, run_id: UUID, parent_run_id: Optional[UUID] = None, tags: Optional[list[str]] = None, **kwargs: Any) -> Any method posthog.ai.langchain.callbacks.CallbackHandler.on_tool_start(serialized: Optional[Dict[str, Any]], input_str: str, *, run_id: UUID, parent_run_id: Optional[UUID] = None, metadata: Optional[Dict[str, Any]] = None, **kwargs: Any) -> Any +method posthog.ai.langchain.middleware.PostHogMiddleware.aafter_agent(state: _PostHogMiddlewareState, runtime: Any) -> dict[str, Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.abefore_agent(state: _PostHogMiddlewareState, runtime: Any) -> dict[str, Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.after_agent(state: _PostHogMiddlewareState, runtime: Any) -> dict[str, Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.awrap_model_call(request: ModelRequest[Any], handler: Callable[[ModelRequest[Any]], Awaitable[ModelResponse[Any]]]) -> ModelResponse[Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.awrap_tool_call(request: ToolCallRequest, handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]]) -> ToolMessage | Command[Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.before_agent(state: _PostHogMiddlewareState, runtime: Any) -> dict[str, Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.wrap_model_call(request: ModelRequest[Any], handler: Callable[[ModelRequest[Any]], ModelResponse[Any]]) -> ModelResponse[Any] +method posthog.ai.langchain.middleware.PostHogMiddleware.wrap_tool_call(request: ToolCallRequest, handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]]) -> ToolMessage | Command[Any] method posthog.ai.openai.openai.WrappedBetaCompletions.parse(posthog_distinct_id: Optional[str] = None, posthog_trace_id: Optional[str] = None, posthog_properties: Optional[Dict[str, Any]] = None, posthog_privacy_mode: bool = False, posthog_groups: Optional[Dict[str, Any]] = None, posthog_provider_override: Optional[str] = None, **kwargs: Any) method posthog.ai.openai.openai.WrappedCompletions.create(posthog_distinct_id: Optional[str] = None, posthog_trace_id: Optional[str] = None, posthog_properties: Optional[Dict[str, Any]] = None, posthog_privacy_mode: bool = False, posthog_groups: Optional[Dict[str, Any]] = None, posthog_provider_override: Optional[str] = None, **kwargs: Any) method posthog.ai.openai.openai.WrappedCompletions.parse(posthog_distinct_id: Optional[str] = None, posthog_trace_id: Optional[str] = None, posthog_properties: Optional[Dict[str, Any]] = None, posthog_privacy_mode: bool = False, posthog_groups: Optional[Dict[str, Any]] = None, posthog_provider_override: Optional[str] = None, **kwargs: Any) @@ -1497,6 +1507,7 @@ module posthog.ai.gemini.gemini_async module posthog.ai.gemini.gemini_converter module posthog.ai.langchain module posthog.ai.langchain.callbacks +module posthog.ai.langchain.middleware module posthog.ai.media module posthog.ai.openai module posthog.ai.openai.openai diff --git a/uv.lock b/uv.lock index af2e6f181..afec1731a 100644 --- a/uv.lock +++ b/uv.lock @@ -2781,6 +2781,7 @@ test = [ { name = "gevent", marker = "implementation_name == 'cpython'" }, { name = "google-genai" }, { name = "jsonschema" }, + { name = "langchain" }, { name = "langchain-anthropic" }, { name = "langchain-community" }, { name = "langchain-core" }, @@ -2828,6 +2829,7 @@ requires-dist = [ { name = "httpx", marker = "extra == 'async'", specifier = ">=0.27.0,<1.0" }, { name = "jsonschema", marker = "extra == 'test'", specifier = ">=4.0" }, { name = "langchain", marker = "extra == 'langchain'", specifier = ">=1.3.9" }, + { name = "langchain", marker = "extra == 'test'", specifier = ">=1.3.9" }, { name = "langchain-anthropic", marker = "extra == 'test'", specifier = ">=1.0" }, { name = "langchain-community", marker = "extra == 'test'", specifier = ">=0.4" }, { name = "langchain-core", marker = "extra == 'test'", specifier = ">=1.0" },