diff --git a/packages/llmpane-py/README.md b/packages/llmpane-py/README.md index ec99c59..43808b4 100644 --- a/packages/llmpane-py/README.md +++ b/packages/llmpane-py/README.md @@ -231,6 +231,59 @@ from llmpane.agent import ToolUseMetadata, ToolStatus # {"metadata": {"tools": [{"name": "get_weather", "status": "completed", "result": "Sunny..."}]}} ``` +### Conversation Identity + +Conversation IDs are server-owned. Pass an existing ID to continue a +conversation; an unknown or absent ID starts a new one under a +server-generated ID, so a client can never choose the key its data is stored +under. Read the real ID back from the terminal chunk: + +```python +async for chunk in session.run(request.conversation_id, request.message): + if chunk.done: + conversation_id = chunk.conversation_id # authoritative +``` + +A `ConversationStore` is persistence, not authorization — it does not check +*who* may read a conversation. Pass a store already scoped to the current +user, or check ownership before calling `run()`. + +### Errors + +Stream errors are classified into a stable `ErrorCode`, a safe display +message, and a retryability hint. Raw exception text is **not** sent to the +client: provider exceptions routinely embed prompt content, tool output, +response bodies, or credentials. + +The original exception is logged with its traceback through the standard +logging module, so you keep full diagnostics: + +```python +import logging + +logging.getLogger("llmpane.errors").setLevel(logging.ERROR) +``` + +For local debugging you can opt back into raw text. Do not enable this in +production — it returns provider exception messages to every client that +triggers an error: + +```python +from llmpane.errors import set_expose_raw_errors + +set_expose_raw_errors(True) # or LLMPANE_EXPOSE_RAW_ERRORS=1 +``` + +### Terminal Chunk Guarantees + +The final chunk is emitted only after the assistant message has been +persisted, and its `message_id` is the ID actually stored — so a client can +rely on that ID existing server-side. If persistence fails, you receive a +classified error instead of a success terminal. + +If the client disconnects mid-stream, the user message is already durable and +partial assistant text is discarded rather than silently stored. + ## Custom Metadata Use generics to add type-safe custom metadata: diff --git a/packages/llmpane-py/llmpane/agent/session.py b/packages/llmpane-py/llmpane/agent/session.py index cab9447..9460124 100644 --- a/packages/llmpane-py/llmpane/agent/session.py +++ b/packages/llmpane-py/llmpane/agent/session.py @@ -3,11 +3,13 @@ from __future__ import annotations import base64 +import uuid from collections.abc import AsyncGenerator from typing import TYPE_CHECKING, Any from llmpane.agent.models import ToolUseMetadata from llmpane.agent.pydantic_ai import stream_agent +from llmpane.errors import create_error_chunk from llmpane.limits import ContentLimits, validate_message_content from llmpane.models import ( ChatMessage, @@ -86,8 +88,19 @@ async def run( ) -> AsyncGenerator[StreamChunk[ToolUseMetadata], None]: """Run a chat turn with automatic persistence. + Conversation identity is server-owned. An unknown ``conversation_id`` + starts a fresh conversation under a server-generated ID rather than + being created under the caller's ID, so a client cannot choose the key + its data is stored under. Read the real ID back from the terminal + chunk's ``conversation_id``. + + Note that a ``ConversationStore`` is persistence, not authorization. It + does not check who may read a conversation; pass a store (or a scoped + wrapper) that is already limited to the current user. + Args: - conversation_id: Optional conversation ID (created if None/not found) + conversation_id: Existing conversation to continue. Unknown or + absent values start a new server-generated conversation. message: User message content (string or list of content parts for multimodal) Yields: @@ -96,12 +109,13 @@ async def run( if self.content_limits is not None: validate_message_content(message, self.content_limits) - # Get or create conversation + # Resume only a conversation that already exists. An unknown ID is + # never persisted under the caller-supplied value. conv = None if conversation_id: conv = await self.store.get_conversation(conversation_id) - if not conv: - conv = await self.store.create_conversation(conversation_id) + if conv is None: + conv = await self.store.create_conversation() # Convert conversation history to Pydantic AI format BEFORE adding new message history = self._to_pydantic_history(conv.messages) @@ -126,30 +140,52 @@ async def run( limits=self.run_limits, ) ) + # The terminal chunk is held back until persistence has succeeded, so a + # storage failure can never be reported to the client as success. If + # the consumer stops iterating first, the terminal is simply never + # emitted and no partial assistant message is written. + terminal: StreamChunk[ToolUseMetadata] | None = None async for chunk in stream: accumulated += chunk.delta - # Add conversation ID to the final chunk if chunk.done: - yield StreamChunk( - delta=chunk.delta, - metadata=chunk.metadata, - done=True, - error=chunk.error, - error_info=chunk.error_info, - message_id=chunk.message_id, - conversation_id=conv.id, - usage=chunk.usage, - ) - else: - yield chunk - - # Persist assistant message (after streaming completes) + terminal = chunk + break + yield chunk + + if terminal is None: + # The agent stream ended without a terminal event. + yield self._terminal_error( + RuntimeError("Agent stream ended without a final chunk"), conv.id + ) + return + + if terminal.error or terminal.error_info: + # Already a classified failure; pass it through with identity + # attached and persist nothing. + yield terminal.model_copy(update={"conversation_id": conv.id}) + return + + # One ID, assigned before persistence, reported after it. + message_id = terminal.message_id or f"msg_{uuid.uuid4().hex[:12]}" if accumulated: assistant_msg: ChatMessage[Any] = ChatMessage( + id=message_id, role=MessageRole.ASSISTANT, content=accumulated, ) - await self.store.add_message(conv.id, assistant_msg) + try: + await self.store.add_message(conv.id, assistant_msg) + except Exception as exc: + yield self._terminal_error(exc, conv.id) + return + + yield terminal.model_copy(update={"message_id": message_id, "conversation_id": conv.id}) + + @staticmethod + def _terminal_error(exc: Exception, conversation_id: str) -> StreamChunk[Any]: + """Build a classified terminal error chunk carrying conversation identity.""" + chunk = create_error_chunk(exc) + return chunk.model_copy(update={"conversation_id": conversation_id}) def _to_pydantic_prompt(self, content: MessageContent) -> str | list[Any]: """Convert llmpane MessageContent to Pydantic AI prompt format. diff --git a/packages/llmpane-py/llmpane/errors.py b/packages/llmpane-py/llmpane/errors.py index 462d316..66b160f 100644 --- a/packages/llmpane-py/llmpane/errors.py +++ b/packages/llmpane-py/llmpane/errors.py @@ -1,12 +1,25 @@ -"""Error classification utilities for llmpane.""" +"""Error classification utilities for llmpane. + +Classification reads the raw exception, but what reaches the client is +deliberately decoupled from it. Provider exceptions routinely embed request +content, prompt text, tool output, or response bodies, so the raw string is +logged internally and never serialized into a `StreamChunk` unless an operator +explicitly opts in for local debugging. +""" from __future__ import annotations +import logging +import os import re from typing import Any from llmpane.models import ErrorCode, StreamChunk, StreamError +# Applications configure this through the standard logging module, e.g. +# logging.getLogger("llmpane.errors").setLevel(logging.DEBUG). +logger = logging.getLogger("llmpane.errors") + # User-friendly messages for each error type ERROR_MESSAGES: dict[ErrorCode, str] = { ErrorCode.RATE_LIMIT: "The AI service is temporarily busy. Please wait a moment and try again.", @@ -28,107 +41,102 @@ ErrorCode.MODEL_UNAVAILABLE, } +_EXPOSE_RAW_ERRORS = os.environ.get("LLMPANE_EXPOSE_RAW_ERRORS", "").lower() in ("1", "true") -def classify_exception(exc: Exception) -> StreamError: - """Classify an exception into a structured StreamError. - This function inspects the exception type and message to determine - the appropriate error code and user-friendly message. +def set_expose_raw_errors(enabled: bool) -> None: + """Include raw exception text in client-facing error details. + + This is a development-only aid. Raw provider exceptions can contain prompt + content, tool output, or response bodies, so leaving this enabled in + production leaks that data to every client that triggers an error. + + Also settable via the ``LLMPANE_EXPOSE_RAW_ERRORS=1`` environment variable. + """ + global _EXPOSE_RAW_ERRORS + _EXPOSE_RAW_ERRORS = enabled + + +def expose_raw_errors() -> bool: + """Whether raw exception text is currently included in error details.""" + return _EXPOSE_RAW_ERRORS + + +def _build_details(exc: Exception) -> dict[str, Any]: + """Build client-safe error details. + + The exception's class name is included because it is a type identifier + rather than message content. The message itself is only included under the + explicit development opt-in. + """ + details: dict[str, Any] = {"exception_type": type(exc).__name__} + if _EXPOSE_RAW_ERRORS: + details["original_error"] = str(exc) + return details + + +def _classify(exc: Exception) -> tuple[ErrorCode, bool, int | None]: + """Map an exception to a semantic code, retryability, and retry hint. + + Classification inspects the raw text, but nothing derived from it beyond a + numeric retry-after hint escapes this function. """ exc_str = str(exc).lower() - exc_type = type(exc).__name__ - # Check for rate limit errors (429) if "429" in exc_str or ("rate" in exc_str and "limit" in exc_str): - retry_after = _extract_retry_after(exc_str) - return StreamError( - code=ErrorCode.RATE_LIMIT, - message=ERROR_MESSAGES[ErrorCode.RATE_LIMIT], - is_retryable=True, - retry_after_seconds=retry_after or 30, - details={"original_error": str(exc)}, - ) - - # Check for resource exhausted (Gemini-style rate limit) + return ErrorCode.RATE_LIMIT, True, _extract_retry_after(exc_str) or 30 + + # Gemini-style rate limit if "resource" in exc_str and "exhausted" in exc_str: - retry_after = _extract_retry_after(exc_str) - return StreamError( - code=ErrorCode.RATE_LIMIT, - message=ERROR_MESSAGES[ErrorCode.RATE_LIMIT], - is_retryable=True, - retry_after_seconds=retry_after or 30, - details={"original_error": str(exc)}, - ) - - # Check for timeout errors + return ErrorCode.RATE_LIMIT, True, _extract_retry_after(exc_str) or 30 + if "timeout" in exc_str or "timed out" in exc_str: - return StreamError( - code=ErrorCode.TIMEOUT, - message=ERROR_MESSAGES[ErrorCode.TIMEOUT], - is_retryable=True, - details={"original_error": str(exc)}, - ) - - # Check for server errors (5xx) + return ErrorCode.TIMEOUT, True, None + if any(code in exc_str for code in ["500", "502", "503", "504"]): - return StreamError( - code=ErrorCode.SERVER_ERROR, - message=ERROR_MESSAGES[ErrorCode.SERVER_ERROR], - is_retryable=True, - details={"original_error": str(exc)}, - ) - - # Check for authentication errors + return ErrorCode.SERVER_ERROR, True, None + if "401" in exc_str or "403" in exc_str or "unauthorized" in exc_str: - return StreamError( - code=ErrorCode.AUTH, - message=ERROR_MESSAGES[ErrorCode.AUTH], - is_retryable=False, - details={"original_error": str(exc)}, - ) - - # Check for validation errors (400) + return ErrorCode.AUTH, False, None + if "400" in exc_str or "validation" in exc_str or "invalid" in exc_str: - return StreamError( - code=ErrorCode.VALIDATION, - message=ERROR_MESSAGES[ErrorCode.VALIDATION], - is_retryable=False, - details={"original_error": str(exc)}, - ) - - # Check for content filter + return ErrorCode.VALIDATION, False, None + if "content" in exc_str and ("filter" in exc_str or "policy" in exc_str): - return StreamError( - code=ErrorCode.CONTENT_FILTERED, - message=ERROR_MESSAGES[ErrorCode.CONTENT_FILTERED], - is_retryable=False, - details={"original_error": str(exc)}, - ) - - # Check for connection errors + return ErrorCode.CONTENT_FILTERED, False, None + if "connection" in exc_str or "network" in exc_str: - return StreamError( - code=ErrorCode.NETWORK_ERROR, - message=ERROR_MESSAGES[ErrorCode.NETWORK_ERROR], - is_retryable=True, - details={"original_error": str(exc)}, - ) - - # Check for model unavailable + return ErrorCode.NETWORK_ERROR, True, None + if "model" in exc_str and ("unavailable" in exc_str or "not found" in exc_str): - return StreamError( - code=ErrorCode.MODEL_UNAVAILABLE, - message=ERROR_MESSAGES[ErrorCode.MODEL_UNAVAILABLE], - is_retryable=True, - details={"original_error": str(exc)}, - ) - - # Default: unknown error + return ErrorCode.MODEL_UNAVAILABLE, True, None + + # Unknown failures stay retryable so a transient blip is not fatal. + return ErrorCode.UNKNOWN, True, None + + +def classify_exception(exc: Exception) -> StreamError: + """Classify an exception into a structured, client-safe StreamError. + + The original exception is logged at ERROR level with a traceback so + operators keep full diagnostics; the returned value carries only a stable + code, a safe message, retryability, and sanitized details. + """ + code, is_retryable, retry_after = _classify(exc) + + logger.error( + "llmpane classified a stream failure as %s (retryable=%s)", + code.value, + is_retryable, + exc_info=exc, + ) + return StreamError( - code=ErrorCode.UNKNOWN, - message=ERROR_MESSAGES[ErrorCode.UNKNOWN], - is_retryable=True, # Allow retry for unknown errors - details={"original_error": str(exc), "exception_type": exc_type}, + code=code, + message=ERROR_MESSAGES[code], + is_retryable=is_retryable, + retry_after_seconds=retry_after, + details=_build_details(exc), ) @@ -141,10 +149,7 @@ def _extract_retry_after(error_message: str) -> int | None: def create_error_chunk(exc: Exception) -> StreamChunk[Any]: - """Create a StreamChunk with structured error information. - - This is a convenience function for creating error responses. - """ + """Create a StreamChunk with structured, client-safe error information.""" error_info = classify_exception(exc) return StreamChunk( error=error_info.message, # Legacy field for backward compatibility diff --git a/packages/llmpane-py/tests/test_stream_safety.py b/packages/llmpane-py/tests/test_stream_safety.py new file mode 100644 index 0000000..608bee0 --- /dev/null +++ b/packages/llmpane-py/tests/test_stream_safety.py @@ -0,0 +1,341 @@ +"""Tests for client-facing error safety, conversation identity, and terminal persistence.""" + +from typing import Any + +import pytest + +from llmpane import ErrorCode +from llmpane.errors import classify_exception, create_error_chunk, set_expose_raw_errors +from llmpane.models import ChatMessage, MessageRole +from llmpane.store import InMemoryStore + +PYDANTIC_AI_AVAILABLE = True +try: + import pydantic_ai # noqa: F401 +except ImportError: + PYDANTIC_AI_AVAILABLE = False + +# Stand-in for the kind of content a provider exception can carry: prompt text, +# tool output, response bodies, credentials in a URL. +SECRET = "sk-live-9f8a7b6c-CUSTOMER-SSN-078-05-1120" + + +@pytest.fixture(autouse=True) +def _reset_raw_error_exposure(): + """Keep the module-level opt-in from leaking between tests.""" + set_expose_raw_errors(False) + yield + set_expose_raw_errors(False) + + +class TestErrorRedaction: + """Raw exception text must not reach the client by default.""" + + def test_secret_absent_from_classified_error(self) -> None: + """A secret in the exception message does not survive classification.""" + error = classify_exception(RuntimeError(f"429 rate limit hit for {SECRET}")) + assert SECRET not in error.model_dump_json() + + def test_secret_absent_from_serialized_chunk(self) -> None: + """A secret in the exception message never appears in the wire payload.""" + chunk = create_error_chunk(RuntimeError(f"Connection failed: {SECRET}")) + assert SECRET not in chunk.model_dump_json() + + @pytest.mark.parametrize( + ("message", "expected"), + [ + ("429 rate limit exceeded", ErrorCode.RATE_LIMIT), + ("resource exhausted", ErrorCode.RATE_LIMIT), + ("request timed out", ErrorCode.TIMEOUT), + ("503 service unavailable", ErrorCode.SERVER_ERROR), + ("401 unauthorized", ErrorCode.AUTH), + ("400 validation failed", ErrorCode.VALIDATION), + ("content filter triggered", ErrorCode.CONTENT_FILTERED), + ("connection reset", ErrorCode.NETWORK_ERROR), + ("model unavailable", ErrorCode.MODEL_UNAVAILABLE), + ("something inexplicable", ErrorCode.UNKNOWN), + ], + ) + def test_semantic_codes_are_stable(self, message: str, expected: ErrorCode) -> None: + """Redaction does not change how failures are classified.""" + assert classify_exception(RuntimeError(message)).code == expected + + def test_codes_stable_even_when_secret_present(self) -> None: + """Classification still works when the message also carries a secret.""" + error = classify_exception(RuntimeError(f"429 rate limit for {SECRET}")) + assert error.code == ErrorCode.RATE_LIMIT + assert error.is_retryable is True + assert error.retry_after_seconds == 30 + + def test_exception_type_is_still_reported(self) -> None: + """The exception class name is a type identifier, not content, so it stays.""" + error = classify_exception(TimeoutError("timed out")) + assert error.details is not None + assert error.details["exception_type"] == "TimeoutError" + + def test_retry_after_hint_survives_redaction(self) -> None: + """A numeric retry hint is extracted without echoing the message.""" + error = classify_exception(RuntimeError(f"429 retry after 12 seconds {SECRET}")) + assert error.retry_after_seconds == 12 + assert SECRET not in error.model_dump_json() + + def test_user_facing_message_is_the_canned_one(self) -> None: + """The client sees a stable message, not the provider's.""" + error = classify_exception(RuntimeError(f"429 {SECRET}")) + assert error.message == ( + "The AI service is temporarily busy. Please wait a moment and try again." + ) + + def test_opt_in_exposes_raw_error(self) -> None: + """The development opt-in restores the raw text.""" + set_expose_raw_errors(True) + error = classify_exception(RuntimeError(f"boom {SECRET}")) + assert error.details is not None + assert SECRET in error.details["original_error"] + + def test_opt_in_is_off_by_default(self) -> None: + """Default configuration omits the raw text entirely.""" + error = classify_exception(RuntimeError("boom")) + assert error.details is not None + assert "original_error" not in error.details + + def test_original_exception_is_logged(self, caplog: pytest.LogCaptureFixture) -> None: + """Operators keep full diagnostics through the logging module.""" + with caplog.at_level("ERROR", logger="llmpane.errors"): + classify_exception(RuntimeError(f"boom {SECRET}")) + assert SECRET in caplog.text + + +def _mock_agent(events: list[Any] | None = None): + """Build a stand-in agent whose stream yields the given pydantic-ai events.""" + from unittest.mock import MagicMock + + agent = MagicMock() + + async def run_stream_events(*args: Any, **kwargs: Any): + for event in events or []: + yield event + + agent.run_stream_events = run_stream_events + return agent + + +@pytest.mark.skipif(not PYDANTIC_AI_AVAILABLE, reason="pydantic-ai not installed") +class TestConversationIdentity: + """Conversation IDs are server-owned.""" + + async def test_unknown_id_is_not_persisted(self) -> None: + """An arbitrary caller-supplied ID never becomes a stored key.""" + from llmpane.agent import ChatSession + + store = InMemoryStore() + session = ChatSession(agent=_mock_agent(), store=store) + + async for _ in session.run("attacker-chosen-id", "Hello"): + pass + + assert await store.get_conversation("attacker-chosen-id") is None + conversations = await store.list_conversations() + assert len(conversations) == 1 + assert conversations[0].id != "attacker-chosen-id" + + async def test_new_conversation_gets_server_generated_id(self) -> None: + """A fresh conversation uses the store's own ID scheme.""" + from llmpane.agent import ChatSession + + store = InMemoryStore() + session = ChatSession(agent=_mock_agent(), store=store) + + async for _ in session.run(None, "Hello"): + pass + + conversations = await store.list_conversations() + assert conversations[0].id.startswith("conv_") + + async def test_known_id_is_resumed(self) -> None: + """An existing conversation is continued, not replaced.""" + from llmpane.agent import ChatSession + + store = InMemoryStore() + existing = await store.create_conversation("conv_existing") + session = ChatSession(agent=_mock_agent(), store=store) + + async for _ in session.run(existing.id, "Hello"): + pass + + assert len(await store.list_conversations()) == 1 + resumed = await store.get_conversation("conv_existing") + assert resumed is not None + assert len(resumed.messages) == 1 + + +@pytest.mark.skipif(not PYDANTIC_AI_AVAILABLE, reason="pydantic-ai not installed") +class TestTerminalPersistence: + """The terminal chunk and stored history must agree.""" + + @staticmethod + def _text_events() -> list[Any]: + from pydantic_ai.messages import PartStartEvent, TextPart + + return [PartStartEvent(index=0, part=TextPart(content="Hello there"))] + + async def _run(self, store: InMemoryStore, events: list[Any] | None = None) -> list[Any]: + from llmpane.agent import ChatSession + + session = ChatSession(agent=_mock_agent(events), store=store) + return [chunk async for chunk in session.run(None, "Hi")] + + async def test_terminal_message_id_matches_persisted_id(self) -> None: + """The ID the client receives is the ID actually stored.""" + store = InMemoryStore() + chunks = await self._run(store, self._text_events()) + + terminal = chunks[-1] + assert terminal.done is True + assert terminal.message_id is not None + + conversation = (await store.list_conversations())[0] + assistant = [m for m in conversation.messages if m.role == MessageRole.ASSISTANT] + assert len(assistant) == 1 + assert assistant[0].id == terminal.message_id + + async def test_terminal_reports_conversation_id(self) -> None: + """The terminal carries the server-owned conversation ID.""" + store = InMemoryStore() + chunks = await self._run(store, self._text_events()) + + conversation = (await store.list_conversations())[0] + assert chunks[-1].conversation_id == conversation.id + + async def test_assistant_content_is_persisted(self) -> None: + """Accumulated text reaches the store.""" + store = InMemoryStore() + await self._run(store, self._text_events()) + + conversation = (await store.list_conversations())[0] + assistant = [m for m in conversation.messages if m.role == MessageRole.ASSISTANT] + assert assistant[0].content == "Hello there" + + async def test_persistence_failure_yields_error_not_success(self) -> None: + """A storage failure cannot be reported to the client as success.""" + from llmpane.agent import ChatSession + + class FailingStore(InMemoryStore): + def __init__(self) -> None: + super().__init__() + self._calls = 0 + + async def add_message( + self, conversation_id: str, message: ChatMessage[Any] + ) -> ChatMessage[Any]: + self._calls += 1 + # Let the user message through, fail persisting the assistant. + if self._calls > 1: + raise RuntimeError("database connection lost") + return await super().add_message(conversation_id, message) + + store = FailingStore() + session = ChatSession(agent=_mock_agent(self._text_events()), store=store) + chunks = [chunk async for chunk in session.run(None, "Hi")] + + terminal = chunks[-1] + assert terminal.done is True + assert terminal.error_info is not None + assert terminal.error_info.code == ErrorCode.NETWORK_ERROR + assert terminal.message_id is None + + async def test_persistence_failure_does_not_leak_raw_error(self) -> None: + """The storage failure is classified, not echoed verbatim.""" + from llmpane.agent import ChatSession + + class FailingStore(InMemoryStore): + async def add_message( + self, conversation_id: str, message: ChatMessage[Any] + ) -> ChatMessage[Any]: + if message.role == MessageRole.ASSISTANT: + raise RuntimeError(f"connection to {SECRET} failed") + return await super().add_message(conversation_id, message) + + session = ChatSession(agent=_mock_agent(self._text_events()), store=FailingStore()) + chunks = [chunk async for chunk in session.run(None, "Hi")] + + assert SECRET not in chunks[-1].model_dump_json() + + async def test_agent_error_passes_through_without_persisting(self) -> None: + """A failed run reports the error and stores no assistant message.""" + from llmpane.agent import ChatSession + + class FailingAgent: + async def run_stream_events(self, *args: Any, **kwargs: Any): + raise RuntimeError("429 rate limit") + yield + + store = InMemoryStore() + session = ChatSession(agent=FailingAgent(), store=store) + chunks = [chunk async for chunk in session.run(None, "Hi")] + + assert chunks[-1].error_info is not None + assert chunks[-1].error_info.code == ErrorCode.RATE_LIMIT + + conversation = (await store.list_conversations())[0] + assert [m for m in conversation.messages if m.role == MessageRole.ASSISTANT] == [] + + async def test_error_terminal_still_reports_conversation_id(self) -> None: + """Clients can still correlate a failed turn with its conversation.""" + from llmpane.agent import ChatSession + + class FailingAgent: + async def run_stream_events(self, *args: Any, **kwargs: Any): + raise RuntimeError("boom") + yield + + store = InMemoryStore() + session = ChatSession(agent=FailingAgent(), store=store) + chunks = [chunk async for chunk in session.run(None, "Hi")] + + conversation = (await store.list_conversations())[0] + assert chunks[-1].conversation_id == conversation.id + + +@pytest.mark.skipif(not PYDANTIC_AI_AVAILABLE, reason="pydantic-ai not installed") +class TestCancellation: + """Abandoning the stream has a defined persistence outcome.""" + + @staticmethod + def _many_text_events() -> list[Any]: + from pydantic_ai.messages import PartDeltaEvent, TextPartDelta + + return [ + PartDeltaEvent(index=0, delta=TextPartDelta(content_delta=f"part{i} ")) + for i in range(5) + ] + + async def test_user_message_survives_cancellation(self) -> None: + """The user's message is persisted before streaming begins.""" + from llmpane.agent import ChatSession + + store = InMemoryStore() + session = ChatSession(agent=_mock_agent(self._many_text_events()), store=store) + + stream = session.run(None, "Hi") + await stream.__anext__() + await stream.aclose() + + conversation = (await store.list_conversations())[0] + user = [m for m in conversation.messages if m.role == MessageRole.USER] + assert len(user) == 1 + assert user[0].content == "Hi" + + async def test_partial_assistant_text_is_not_persisted(self) -> None: + """Partial output is discarded rather than silently stored.""" + from llmpane.agent import ChatSession + + store = InMemoryStore() + session = ChatSession(agent=_mock_agent(self._many_text_events()), store=store) + + stream = session.run(None, "Hi") + await stream.__anext__() + await stream.aclose() + + conversation = (await store.list_conversations())[0] + assert [m for m in conversation.messages if m.role == MessageRole.ASSISTANT] == []