From 3f919be5d35fb46a093d85067826e5ae869d09bd Mon Sep 17 00:00:00 2001 From: Maxim Svistunov Date: Mon, 28 Sep 2026 11:49:04 +0200 Subject: [PATCH] LCORE-3908: give the turns we store ourselves one owner OGX appends a turn to the conversation only when it is handed the conversation parameter and runs the inference. In compacted mode that parameter is dropped, so lightspeed-stack appends the turn itself. Every endpoint made that write on its own, and nothing enforced it: two unrelated cleanups removed the calls, and for six weeks conversations stopped growing after a compaction without anything failing (LCORE-3883). The write now has one owner, PendingTurn in utils/pending_turn.py. An endpoint creates it from the parameters the request is sent with and the original input of the compaction result, and reports how the turn ended: store_completed, store_blocked or store_interrupted. The owner decides whether the turn is ours, which input is stored, and it stores a turn once: the first report settles the turn, every later one does nothing. It covers all seven places that stored a turn, not only the compacted ones, because they are the same duty: a turn OGX does not store. That is a conversation in compacted mode, a request a shield blocked on /v1/responses, a stream the client interrupted, a continuation from previous_response_id. store_compacted_turn is removed; no endpoint appends a turn except through the owner. The shield capabilities still store the turn they rejected themselves, from inside the agent run, and only when the model was handed the conversation, so never in compacted mode. /v1/query and /v1/streaming_query have no blocked turn of their own to store: the shield moderation that ran before the agent on those two endpoints is gone from main (#2780). A compacted request that loses its turn through a change in the code fails: - creating the owner for compacted parameters without the original input raises ValueError, which is how a caller that lost the hand-over looks; - ensure_settled raises TurnNotStoredError when a compacted request reaches the end of its handler and nobody tried to store its turn. /v1/query runs in a scope that makes the check when it is left; the streaming paths, /v1/responses and A2A make it after their write, in the function that calls the write, not in the one that performs it. The error is not a RuntimeError, because the endpoints report that one as an inference failure. The check does not cover a stream the client stops reading, a /v1/responses stream without a final response, and a failed write that is logged. They store nothing, as before. Every condition of the old call sites is kept: where the write happens, what is stored for a failed or incomplete turn, what a failed write does to the request (it fails /v1/query and /v1/responses, it is logged on /v1/streaming_query, A2A and the interrupt path), and the order relative to quota consumption. The guard that lets only one of stream end, cancellation handler and interrupt callback finish a turn is unchanged. What is stored, and when, does not change. Tests: - integration, tests/integration/endpoints/test_turn_persistence.py: for /v1/query, /v1/streaming_query, /v1/responses (blocking and streaming) and A2A, every way a turn can end (completed, blocked on /v1/responses, interrupted, failed, cut short by the model, abandoned by the client, continued from a previous response, store off), compacted and not; what a failed write does to the request on each endpoint; and what happens when the step that stores the turn is taken out of an endpoint. Each test runs the real handler and compares the conversation item by item afterwards, so a missing write, a second write and a write of the wrong input all fail it. 39 of the 44 tests pass on main unchanged. The five that fail there take the write out of an endpoint, which main does not notice. - unit: the owner, the hand-over from the compaction seam (omit_conversation and the original input are set together), and a scan of src that fails when a module other than the known writers mentions one of the functions that write to a conversation. - existing unit tests follow the new parameter. Those that asserted on the arguments of a patched helper in the agent and interrupt paths now assert on what is stored. Mocked request parameters got the fields the owner reads. The A2A helpers of the compaction tests moved to _compaction_helpers.py, where both test modules import them. The tests were checked by mutation: with the write removed from any of the four endpoints, a check removed, the write doubled, the settled state ignored, the interrupt guard ignored, the explicit rewrite stored in place of the input, a turn OGX stores stored by us as well, or the store flag ignored, tests fail. One check is not pinned by a test: the one after the write of a blocked /v1/responses stream. Docs: the design document describes the owner and has a table of what is stored per endpoint and per way a turn can end. --- .../conversation-compaction.md | 55 + docs/devel_doc/ARCHITECTURE.md | 2 +- src/app/endpoints/a2a.py | 81 +- src/app/endpoints/query.py | 25 +- src/app/endpoints/responses.py | 89 +- src/app/endpoints/streaming_query.py | 10 +- src/utils/agents/query.py | 25 +- src/utils/agents/streaming.py | 53 +- src/utils/conversation_compaction.py | 25 +- src/utils/pending_turn.py | 285 ++++ src/utils/stream_interrupts.py | 52 +- .../endpoints/_compaction_helpers.py | 96 ++ .../endpoints/test_compaction_a2a.py | 121 +- .../endpoints/test_turn_persistence.py | 1358 +++++++++++++++++ tests/unit/app/endpoints/test_a2a.py | 34 + tests/unit/app/endpoints/test_query.py | 28 + tests/unit/app/endpoints/test_query_otel.py | 4 + tests/unit/app/endpoints/test_responses.py | 40 +- tests/unit/utils/agents/test_query.py | 69 +- tests/unit/utils/agents/test_streaming.py | 129 +- .../utils/test_conversation_compaction.py | 11 - tests/unit/utils/test_pending_turn.py | 408 +++++ tests/unit/utils/test_stream_interrupts.py | 115 +- 23 files changed, 2719 insertions(+), 396 deletions(-) create mode 100644 src/utils/pending_turn.py create mode 100644 tests/integration/endpoints/test_turn_persistence.py create mode 100644 tests/unit/utils/test_pending_turn.py diff --git a/docs/design/conversation-compaction/conversation-compaction.md b/docs/design/conversation-compaction/conversation-compaction.md index 9abe1dbc0..28ca8b66b 100644 --- a/docs/design/conversation-compaction/conversation-compaction.md +++ b/docs/design/conversation-compaction/conversation-compaction.md @@ -315,6 +315,7 @@ Add `compaction` field to the root `Configuration` class. | `src/models/config.py` | Add `CompactionConfiguration` (near `ConversationHistoryConfiguration`) | | `src/configuration.py` | Add `compaction_configuration` property to `AppConfig` singleton | | `src/utils/conversation_compaction.py` | New module: `apply_compaction()` / `apply_compaction_blocking()`, `needs_compaction_path()`, marker helpers, per-conversation lock | +| `src/utils/pending_turn.py` | `PendingTurn`, the owner of every turn the endpoints store themselves: decides whether the turn is ours, stores it against the input as it arrived, and stores it once — LCORE-3908 | | `src/models/common/responses/responses_api_params.py` | `omit_conversation` flag — drops the `conversation` parameter from the request body in compacted mode | | `src/app/endpoints/query.py` | Call `apply_compaction_blocking()` after preparing params; store the turn in compacted mode | | `src/app/endpoints/streaming_query.py` | Compaction-aware SSE path that emits the `compaction` event before summarizing (R12) | @@ -347,6 +348,60 @@ When compaction is active, the endpoint builds explicit input, the `conversation` parameter is omitted (via `ResponsesApiParams.omit_conversation`), and the completed turn is appended to the conversation items afterward. +That append has one owner, `PendingTurn` in `src/utils/pending_turn.py` +(LCORE-3908). An endpoint creates it from the parameters the request is sent +with and the `original_input` of the `CompactionResult`, and reports how the +turn ended: `store_completed`, `store_blocked` or `store_interrupted`. +The first report settles the turn and every later one does nothing, so a turn is +stored once even when several paths of a request want to store it (the end of a +stream, the cancellation handler, the interrupt callback). `PendingTurn` also +covers the other turns OGX does not store: a request a shield blocked on +`/v1/responses`, an interrupted stream, a continuation from +`previous_response_id`. + +A compacted request that loses its turn through a change in the code fails. +Building a `PendingTurn` for compacted parameters without the original input +raises `ValueError`, and `ensure_settled()` raises `TurnNotStoredError` when a +compacted request reaches the end of its handler and nobody tried to store its +turn. `/v1/query` runs inside the `pending_turn()` scope, which makes that check +when it is left; the streaming paths, `/v1/responses` and A2A make it after +their write. Before this, a lost write was silent: the conversation stopped +growing and nothing failed (LCORE-3883). + +The check does not cover three endings, which store nothing and are rows of the +table below: a stream the client stops reading, a `/v1/responses` stream without +a final response, and a failed write where the failure is logged. + +The shield capabilities (question validity, Granite Guardian) are the one writer +outside the owner. They store the turn they rejected from inside the agent run, +and only when the model was handed the conversation, so never in compacted mode. + +What is stored, per endpoint and per way a turn can end. "OGX" means the +`conversation` parameter was sent and OGX stores the turn itself. The last +column says what a failed write does to the request. + +| Endpoint | Turn ended | Not compacted | Compacted | Failed write | +|---|---|---|---|---| +| `/v1/query` | completed | OGX | stored | request fails | +| | model call failed | not stored | not stored | | +| | run did not finish with success | OGX | not stored | | +| `/v1/streaming_query` | completed, also when the run did not finish with success | OGX | stored when the stream ends | logged | +| | interrupted by the client | stored, with the answer so far | stored, with the answer so far | logged | +| | client stopped reading | OGX | not stored | | +| `/v1/responses` | completed, incomplete or failed | OGX; stored when continuing from `previous_response_id` | stored | request fails; a stream ends before `[DONE]` | +| | blocked by a shield | stored | stored | request fails | +| | stream without a final response | not stored | not stored | | +| | `store: false` | not stored | never compacted | | +| A2A | completed | OGX | stored | logged | +| | agent run failed | not stored | not stored | | + +`tests/integration/endpoints/test_turn_persistence.py` pins the table. It runs +the real handlers and compares the conversation item by item after the request. +The compacted column is covered row by row; the other column for the rows where +lightspeed-stack stores the turn, and for a completed turn on each endpoint. The +last column is covered for a completed turn on each endpoint and for an +interrupted stream. + ## Fetching conversation history Use the same pattern as `conversations_v1.py:240-246`: diff --git a/docs/devel_doc/ARCHITECTURE.md b/docs/devel_doc/ARCHITECTURE.md index 2bcb080e5..1eeb76b1d 100644 --- a/docs/devel_doc/ARCHITECTURE.md +++ b/docs/devel_doc/ARCHITECTURE.md @@ -391,7 +391,7 @@ The compaction system is split into two layers: - Marker persistence (`[lightspeed:compaction-summary]` sentinel in conversation items) - `CompactionStartedEvent` emission for streaming progress indicators - `apply_compaction()` (async generator) — Main entry point used by all endpoints - - `store_compacted_turn()` — Appends user query + LLM output when in compacted mode + - `PendingTurn` (`utils/pending_turn.py`) — Owns the turns the endpoints store themselves: appends user query + LLM output when in compacted mode, once per request **Data Flow:** diff --git a/src/app/endpoints/a2a.py b/src/app/endpoints/a2a.py index 3f1ed2ad1..ad6a0b81d 100644 --- a/src/app/endpoints/a2a.py +++ b/src/app/endpoints/a2a.py @@ -33,7 +33,7 @@ ) from a2a.utils import new_agent_text_message, new_task from fastapi import APIRouter, Depends, HTTPException, Request, status -from ogx_client import ApiException +from ogx_client import ApiException, AsyncOgxClient from opentelemetry import trace from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError @@ -63,11 +63,7 @@ from models.common.responses.responses_api_params import ResponsesApiParams from models.config import Action from utils.agents.error_handler import map_agent_inference_error -from utils.conversation_compaction import ( - CompactionResult, - apply_compaction_blocking, - store_compacted_turn, -) +from utils.conversation_compaction import apply_compaction_blocking from utils.mcp.mcp_headers import McpHeaders, mcp_headers_dependency from utils.otel_tracing import ( SpanAttributes, @@ -76,6 +72,7 @@ anonymize_value, set_span_attributes, ) +from utils.pending_turn import PendingTurn from utils.pydantic_ai_helpers import build_agent, captured_output_items from utils.query import extract_provider_and_model_from_model_id from utils.responses import prepare_responses_params @@ -213,10 +210,39 @@ def _record_execution_span( span.set_attribute(SpanAttributes.INFERENCE_TIME, inference_time) -async def _persist_compacted_a2a_turn( - client: Any, +async def _compact_a2a_request( + client: AsyncOgxClient, responses_params: ResponsesApiParams, - compaction: CompactionResult, +) -> tuple[ResponsesApiParams, PendingTurn]: + """Compact the conversation of an A2A request if it nears the context window. + + A2A is not a browser SSE stream, so no progress event is emitted; the + blocking variant summarizes inline before the call. No conversation cache + is passed: the A2A executor has no resolved user_id for the + (user_id, conversation_id) cache key, so A2A runs in marker-only mode + (additive summaries, no persisted fold). + + Parameters: + client: OGX client. + responses_params: Prepared Responses API parameters. + + Returns: + The parameters to send, and the pending turn of the request, which + stores the turn when OGX does not (LCORE-3908). + """ + compaction = await apply_compaction_blocking( + client, + responses_params, + configuration.inference, + configuration.compaction, + ) + return compaction.params, PendingTurn.for_request( + client, compaction.params, compaction.original_input + ) + + +async def _persist_compacted_a2a_turn( + turn: PendingTurn, agent: Any, task_id: str, ) -> None: @@ -228,22 +254,13 @@ async def _persist_compacted_a2a_turn( context. Parameters: - client: OGX client used to write the conversation items. - responses_params: Prepared Responses API parameters. - compaction: Outcome of applying compaction. Nothing is written unless - the request was served in compacted mode. + turn: The pending turn of the request. Nothing is written when OGX + stores the turn itself. agent: The pydantic-ai agent whose model captured the output items. task_id: A2A task identifier, used for error reporting. """ - if not compaction.compacted or compaction.original_input is None: - return try: - await store_compacted_turn( - client, - responses_params.conversation, - compaction.original_input, - captured_output_items(agent), - ) + await turn.store_completed(captured_output_items(agent)) except Exception: # pylint: disable=broad-except # The caller already has its answer; the cost of the failure is that the # next turn in this context loses this one. @@ -479,19 +496,9 @@ async def _process_task_streaming( # pylint: disable=too-many-locals,too-many-s store=True, request_headers=self.request_headers, ) - # Compact the conversation if it is approaching the context window - # limit. A2A is not a browser SSE stream, so no progress event is - # emitted; the blocking variant summarizes inline before the call. - # No conversation cache is passed: the A2A executor has no resolved - # user_id for the (user_id, conversation_id) cache key, so A2A runs - # in marker-only mode (additive summaries, no persisted fold). - compaction = await apply_compaction_blocking( - client, - responses_params, - configuration.inference, - configuration.compaction, + responses_params, turn = await _compact_a2a_request( + client, responses_params ) - responses_params = compaction.params _record_model_span(span, responses_params.model) agent = build_agent( @@ -582,15 +589,15 @@ async def _process_task_streaming( # pylint: disable=too-many-locals,too-many-s ) return - await _persist_compacted_a2a_turn( - client, responses_params, compaction, agent, task_id - ) + await _persist_compacted_a2a_turn(turn, agent, task_id) + turn.ensure_settled() _record_execution_span( span, self._tool_call_names, self._run_result, - compaction.compacted, + # compacted mode: the conversation parameter was not sent + responses_params.omit_conversation, (datetime.now(UTC) - started_at).total_seconds(), ) diff --git a/src/app/endpoints/query.py b/src/app/endpoints/query.py index 0e3a7fb4c..b96e69408 100644 --- a/src/app/endpoints/query.py +++ b/src/app/endpoints/query.py @@ -47,6 +47,7 @@ root_span_turn_attributes, set_span_attributes, ) +from utils.pending_turn import pending_turn from utils.query import ( consume_query_tokens, store_query_results, @@ -259,16 +260,20 @@ async def _handle_query_with_tracing( if a.content_type in IMAGE_CONTENT_TYPES ] or None - # Retrieve response using Responses API - turn_summary = await retrieve_agent_response( - client, - responses_params, - endpoint_path, - compaction.original_input if compaction.compacted else None, - shield_ids=query_request.shield_ids, - no_tools=bool(query_request.no_tools), - image_attachments=image_attachments, - ) + # Retrieve response using Responses API. In compacted mode OGX does not + # store the turn; the scope fails the request if nobody tried to store it. + async with pending_turn( + client, responses_params, compaction.original_input + ) as turn: + turn_summary = await retrieve_agent_response( + client, + responses_params, + endpoint_path, + turn, + shield_ids=query_request.shield_ids, + no_tools=bool(query_request.no_tools), + image_attachments=image_attachments, + ) # Combine inline RAG results (BYOK + Solr) with tool-based RAG results for the transcript rag_chunks = inline_rag_context.rag_chunks diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index df12cbfe4..4dcf853bc 100644 --- a/src/app/endpoints/responses.py +++ b/src/app/endpoints/responses.py @@ -66,7 +66,6 @@ apply_compaction_blocking, configured_conversation_cache, ) -from utils.conversations import append_turn_items_to_conversation from utils.endpoints import ( check_configuration_loaded, resolve_response_context, @@ -84,6 +83,7 @@ root_span_turn_attributes, set_span_attributes, ) +from utils.pending_turn import PendingTurn from utils.prompts import get_system_prompt from utils.query import ( consume_query_tokens, @@ -363,38 +363,52 @@ def _raise_response_api_http_exception( raise HTTPException(**error_response.model_dump()) from error +def _pending_turn( + api_params: ResponsesApiParams, + context: ResponsesContext, +) -> PendingTurn: + """Return the pending turn of a request, which stores it when OGX does not. + + Args: + api_params: Responses API parameters of the request, after compaction. + context: Request-scoped Responses API context, carrying the input as + it arrived when the request is compacted. + + Returns: + The pending turn of the request. + """ + return PendingTurn.for_request( + context.client, api_params, context.compacted_original_input + ) + + async def _persist_blocked_response_turn( api_params: ResponsesApiParams, context: ResponsesContext, + turn: Optional[PendingTurn] = None, ) -> None: """Persist a shield-blocked refusal turn when response storage is enabled. + In compacted mode the conversation parameter was dropped and + ``api_params.input`` is the explicit-input rewrite, so the turn is stored + against the input as it arrived (LCORE-1572). + Args: api_params: Responses API parameters for the blocked request. context: Request-scoped Responses API context with moderation details. + turn: The pending turn of the request; created from the parameters + and the context when not given. """ - if api_params.store: - moderation_result = cast("ShieldModerationBlocked", context.moderation_result) - # In compacted mode the conversation parameter was dropped and - # api_params.input is the explicit-input rewrite, so persist the turn - # against the original user input instead (LCORE-1572). - user_input = ( - context.compacted_original_input - if context.compacted_original_input is not None - else api_params.input - ) - await append_turn_items_to_conversation( - client=context.client, - conversation_id=api_params.conversation, - user_input=user_input, - llm_output=[moderation_result.refusal_response], - ) + moderation_result = cast("ShieldModerationBlocked", context.moderation_result) + turn = turn or _pending_turn(api_params, context) + await turn.store_blocked(moderation_result.refusal_response) async def _append_previous_response_turn( api_params: ResponsesApiParams, context: ResponsesContext, output: Sequence[OpenAIResponseOutput], + turn: Optional[PendingTurn] = None, ) -> None: """Append the completed turn when OGX did not store it automatically. @@ -409,23 +423,11 @@ async def _append_previous_response_turn( api_params: Responses API parameters containing conversation details. context: Request-scoped Responses API context. output: Final output items from the Responses API object. + turn: The pending turn of the request; created from the parameters + and the context when not given. """ - if not api_params.store: - return - if context.compacted_original_input is not None: - await append_turn_items_to_conversation( - context.client, - api_params.conversation, - context.compacted_original_input, - output, - ) - elif api_params.previous_response_id: - await append_turn_items_to_conversation( - context.client, - api_params.conversation, - api_params.input, - output, - ) + turn = turn or _pending_turn(api_params, context) + await turn.store_completed(output) def _store_response_query_results( @@ -791,7 +793,9 @@ async def handle_streaming_response( turn_summary.id = context.moderation_result.moderation_id turn_summary.llm_response = context.moderation_result.message generator = shield_violation_generator(api_params, context) - await _persist_blocked_response_turn(api_params, context) + turn = _pending_turn(api_params, context) + await _persist_blocked_response_turn(api_params, context, turn) + turn.ensure_settled() queue_blocked_response_event( api_params, context, @@ -1243,13 +1247,17 @@ async def response_generator( context.inline_rag_context.rag_chunks + turn_summary.rag_chunks ) - # Explicitly append the turn to conversation if context passed by previous response + # Store the turn when OGX did not: in compacted mode, or when the request + # continues from a previous response. if latest_response_object: + turn = _pending_turn(api_params, context) await _append_previous_response_turn( api_params, context, latest_response_object.output, + turn, ) + turn.ensure_settled() async def generate_response( @@ -1331,6 +1339,7 @@ async def handle_non_streaming_response( inference_span: Optional[trace.Span] = None inference_start_time: Optional[float] = None inference_time: Optional[float] = None + turn = _pending_turn(api_params, context) # Fork: Get response object (blocked vs normal) if context.moderation_result.decision == "blocked": @@ -1343,7 +1352,7 @@ async def handle_non_streaming_response( usage=get_zero_usage(), **api_params.echoed_params(configuration.rag_id_mapping), ) - await _persist_blocked_response_turn(api_params, context) + await _persist_blocked_response_turn(api_params, context, turn) queue_blocked_response_event(api_params, context, output_text) else: inference_start_time = time.monotonic() @@ -1379,11 +1388,13 @@ async def handle_non_streaming_response( token_usage=token_usage, ) output_text = extract_text_from_response_items(api_response.output) - # Explicitly append the turn to conversation if context passed by previous response + # Store the turn when OGX did not: in compacted mode, or when the + # request continues from a previous response. await _append_previous_response_turn( api_params, context, api_response.output, + turn, ) except ( @@ -1401,6 +1412,10 @@ async def handle_non_streaming_response( ) _raise_response_api_http_exception(e, api_params, context, inference_span) + # The request has its answer. Fail it if its turn had to be stored here + # and nobody tried to. + turn.ensure_settled() + vector_store_ids = extract_vector_store_ids_from_tools(api_params.tools) turn_summary = build_turn_summary( api_response, diff --git a/src/app/endpoints/streaming_query.py b/src/app/endpoints/streaming_query.py index 28cbd04e7..bdcc4fd51 100644 --- a/src/app/endpoints/streaming_query.py +++ b/src/app/endpoints/streaming_query.py @@ -42,7 +42,6 @@ from models.common.query import Attachment from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams -from models.common.responses.types import ResponseInput from models.common.turn_summary import ContextStatus from models.config import Action from utils.agents.streaming import ( @@ -69,6 +68,7 @@ anonymize_value, set_span_attributes, ) +from utils.pending_turn import PendingTurn from utils.query import ( extract_provider_and_model_from_model_id, handle_known_apistatus_errors, @@ -418,7 +418,7 @@ async def generate_response_with_compaction( request_id=context.request_id, ) - compacted_original_input: Optional[ResponseInput] = None + turn: Optional[PendingTurn] = None context_status: ContextStatus = "full" try: async for item in apply_compaction( @@ -435,7 +435,9 @@ async def generate_response_with_compaction( yield stream_compaction_event(context.conversation_id) elif isinstance(item, CompactionResult): responses_params = item.params - compacted_original_input = item.original_input + turn = PendingTurn.for_request( + context.client, item.params, item.original_input + ) context_status = item.context_status generator, turn_summary = await retrieve_agent_response_generator( @@ -489,7 +491,7 @@ async def generate_response_with_compaction( background_topic_summary_tasks=_background_topic_summary_tasks, root_span=root_span, emit_start=False, - original_input=compacted_original_input, + turn=turn, context_status=context_status, ): yield event diff --git a/src/utils/agents/query.py b/src/utils/agents/query.py index 8b6db71c3..e3f1bacd4 100644 --- a/src/utils/agents/query.py +++ b/src/utils/agents/query.py @@ -27,7 +27,6 @@ from models.common.agents import AgentTurnAccumulator from models.common.query import Attachment from models.common.responses.responses_api_params import ResponsesApiParams -from models.common.responses.types import ResponseInput from models.common.turn_summary import TurnSummary from utils.agents.error_handler import map_agent_inference_error from utils.agents.tool_processor import ( @@ -39,7 +38,6 @@ from utils.conversation_compaction import ( agent_prompt_text, reject_image_attachments_in_compacted_mode, - store_compacted_turn, ) from utils.otel_tracing import ( SpanAttributes, @@ -48,6 +46,7 @@ llm_inference_span_attributes, set_span_attributes, ) +from utils.pending_turn import PendingTurn from utils.pydantic_ai_helpers import build_agent, captured_output_items from utils.query import ( build_multimodal_input, @@ -230,7 +229,7 @@ async def retrieve_agent_response( client: AsyncOgxClient, responses_params: ResponsesApiParams, endpoint_path: str, - original_input: Optional[ResponseInput] = None, + turn: Optional[PendingTurn] = None, no_tools: bool = False, image_attachments: Optional[list[Attachment]] = None, shield_ids: Optional[list[str]] = None, @@ -241,9 +240,10 @@ async def retrieve_agent_response( client: OGX client used when building the agent. responses_params: Prepared Responses API parameters. endpoint_path: Endpoint path used for metric labeling. - original_input: Original user input before the explicit-input rewrite. - Set only in compacted mode; when set, the completed turn is - appended to the conversation explicitly (LCORE-3883). + turn: The pending turn of the request, which stores the turn when OGX + does not (LCORE-3908). Required in compacted mode, where it + carries the input as it arrived; created from the parameters + otherwise. no_tools: Whether to skip tool processing. image_attachments: Image attachments for multimodal prompt construction. shield_ids: Optional list of shield names to run for this turn, mirroring @@ -253,7 +253,10 @@ async def retrieve_agent_response( Raises: HTTPException: On agent or provider failure. + ValueError: When the request is compacted and no pending turn is given. """ + if turn is None: + turn = PendingTurn.for_request(client, responses_params) with tracer.start_as_current_span("llm.inference") as span: # Extract provider and model from model_id provider_id, model_id = extract_provider_and_model_from_model_id( @@ -324,14 +327,8 @@ async def retrieve_agent_response( add_span_event(span, SpanEvents.LLM_INFERENCE_COMPLETED) # In compacted mode the conversation parameter was not sent, so OGX did - # not persist this turn. Append it ourselves to keep the recent-turn + # not persist this turn. It is stored here, to keep the recent-turn # buffer and the audit history intact for the next request (LCORE-3883). - if original_input is not None: - await store_compacted_turn( - client, - responses_params.conversation, - original_input, - turn_summary.output_items, - ) + await turn.store_completed(turn_summary.output_items) return turn_summary diff --git a/src/utils/agents/streaming.py b/src/utils/agents/streaming.py index 49637c92d..01641a5f4 100644 --- a/src/utils/agents/streaming.py +++ b/src/utils/agents/streaming.py @@ -44,7 +44,6 @@ TurnCompleteStreamPayload, ) from models.common.query import Attachment -from models.common.responses import ResponseInput from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams from models.common.turn_summary import ContextStatus, TurnSummary @@ -64,7 +63,6 @@ from utils.conversation_compaction import ( agent_prompt_text, reject_image_attachments_in_compacted_mode, - store_compacted_turn, ) from utils.otel_tracing import ( SpanAttributes, @@ -74,6 +72,7 @@ root_span_turn_attributes, set_span_attributes, ) +from utils.pending_turn import PendingTurn from utils.pydantic_ai_helpers import build_agent, captured_output_items from utils.query import ( build_multimodal_input, @@ -150,9 +149,8 @@ async def retrieve_agent_response_generator( async def _persist_compacted_turn( context: ResponseGeneratorContext, - responses_params: ResponsesApiParams, + turn: PendingTurn, turn_summary: TurnSummary, - original_input: Optional[ResponseInput], persist_guard: list[bool], ) -> None: """Append a completed compacted turn to the conversation (LCORE-3883). @@ -162,24 +160,18 @@ async def _persist_compacted_turn( recent-turn buffer and the audit history intact for the next request. Parameters: - context: Streaming request context, providing the OGX client. - responses_params: Prepared Responses API parameters. + context: Streaming request context, used for error reporting. + turn: The pending turn of the request. Nothing is written when OGX + stores the turn itself. turn_summary: Completed turn, carrying the captured output items. - original_input: The user input before the explicit-input rewrite. When - ``None`` the request was not compacted and nothing is written. - persist_guard: Single-element flag shared with the interrupt path, so a - turn is persisted at most once. + persist_guard: Single-element flag shared with the interrupt path, so + that only one of them finishes the turn. """ - if original_input is None or persist_guard[0]: + if not turn.ours or persist_guard[0]: return persist_guard[0] = True try: - await store_compacted_turn( - context.client, - responses_params.conversation, - original_input, - turn_summary.output_items, - ) + await turn.store_completed(turn_summary.output_items) except Exception: # pylint: disable=broad-except # The client already has its answer, so the stream still succeeds; the # cost of the failure is that the next request loses this turn. @@ -197,7 +189,7 @@ async def generate_agent_response( # pylint: disable=too-many-statements background_topic_summary_tasks: list[asyncio.Task[None]], root_span: trace.Span, emit_start: bool = True, - original_input: Optional[ResponseInput] = None, + turn: Optional[PendingTurn] = None, context_status: ContextStatus = "full", ) -> AsyncIterator[str]: """Wrap an agent SSE generator with cleanup logic. @@ -215,23 +207,31 @@ async def generate_agent_response( # pylint: disable=too-many-statements root_span: OpenTelemetry root span for this request. emit_start: Whether to emit the SSE start event. False when the caller (the compaction-aware wrapper) has already emitted it. - original_input: In compacted mode, the original user input before the - explicit-input rewrite. Used to persist the completed turn with its - structured input (preserving attachments); ``None`` otherwise. + turn: The pending turn of the request, which stores the turn when OGX + does not (LCORE-3908). Required in compacted mode, where it + carries the input as it arrived; created from the parameters + otherwise. context_status: Whether the conversation context was sent in full ("full") or older turns were replaced by a summary ("summarized"). Reported to the client in the SSE end event. Yields: SSE-formatted strings from the wrapped generator. + + Raises: + ValueError: When the request is compacted and no pending turn is given. + TurnNotStoredError: When a compacted stream ran to its end and nobody + tried to store its turn. """ media_type = context.query_request.media_type or MEDIA_TYPE_JSON + if turn is None: + turn = PendingTurn.for_request(context.client, responses_params) persist_guard = register_interrupt_callback( context, responses_params, turn_summary, background_topic_summary_tasks, - original_input, + turn, ) stream_completed = False if emit_start: @@ -271,7 +271,7 @@ async def generate_agent_response( # pylint: disable=too-many-statements responses_params, turn_summary, background_topic_summary_tasks, - original_input, + turn, ) yield serialize_event( TokenStreamPayload.create( @@ -290,9 +290,10 @@ async def generate_agent_response( # pylint: disable=too-many-statements root_span.end() return - await _persist_compacted_turn( - context, responses_params, turn_summary, original_input, persist_guard - ) + await _persist_compacted_turn(context, turn, turn_summary, persist_guard) + # The stream ran to its end. Fail if its turn had to be stored here and + # neither this path nor an interrupt tried to. + turn.ensure_settled() should_generate_topic_summary = ( context.query_request.conversation_id is None diff --git a/src/utils/conversation_compaction.py b/src/utils/conversation_compaction.py index 1b9f7c480..17a994e80 100644 --- a/src/utils/conversation_compaction.py +++ b/src/utils/conversation_compaction.py @@ -69,7 +69,6 @@ summarize_chunk, ) from utils.conversations import ( - append_turn_items_to_conversation, build_add_items_request, get_all_conversation_items, ) @@ -187,9 +186,9 @@ class CompactionResult: original_input: The new user query exactly as it arrived (before the explicit-input rewrite). Populated only in compacted mode (where ``compacted`` is True); ``None`` otherwise. In compacted mode the - caller must append this plus the LLM output to the conversation - items itself, since the ``conversation`` parameter is no longer - passed to OGX. + caller hands it to ``utils.pending_turn.PendingTurn``, which + stores this plus the LLM output in the conversation, since the + ``conversation`` parameter is no longer passed to OGX. """ params: ResponsesApiParams @@ -889,21 +888,3 @@ async def needs_compaction_path( estimated += estimate_conversation_tokens(items, encoding_name=encoding_name) estimated += _estimate_response_input_tokens(params.input, encoding_name) return _should_compact(estimated, context_window, compaction_config) - - -async def store_compacted_turn( - client: AsyncOgxClient, - conversation_id: str, - original_input: ResponseInput, - output_items: Sequence[Any], -) -> None: - """Append a completed turn to the conversation when in compacted mode. - - In compacted mode the ``conversation`` parameter is not sent to inference, - so OGX does not auto-store the turn. lightspeed-stack appends the - user query and the LLM output to the conversation items itself, keeping the - full history (and the recent-turn buffer for the next request) intact. - """ - await append_turn_items_to_conversation( - client, conversation_id, original_input, output_items - ) diff --git a/src/utils/pending_turn.py b/src/utils/pending_turn.py new file mode 100644 index 000000000..36ddc1036 --- /dev/null +++ b/src/utils/pending_turn.py @@ -0,0 +1,285 @@ +"""The owner of the turns lightspeed-stack has to store itself (LCORE-3908). + +OGX appends a turn to the conversation only when it is handed the +``conversation`` parameter and runs the inference. In every other case the turn +is lightspeed-stack's to store: + +* the conversation is served in compacted mode, where the ``conversation`` + parameter is dropped in favor of explicit input (LCORE-1572), +* the request continues from a ``previous_response_id``, +* a shield blocked the request, so OGX was never called, +* the client interrupted the stream, so the OGX call was cancelled. + +When that write goes missing nothing fails: the conversation just stops +growing, and the model loses the turn on the next request (LCORE-3883). So the +write has one owner, :class:`PendingTurn`. It decides whether lightspeed-stack +has to store the turn and which input is stored, and it stores a turn once, +regardless of how many callers ask. + +A request served in compacted mode that reaches the end of its handler without +anyone having tried to store its turn fails: +:meth:`PendingTurn.ensure_settled` raises :class:`TurnNotStoredError`. Three +endings are not covered by that check and store nothing: a stream the client +stops reading, a ``/v1/responses`` stream without a final response, and a write +that fails where the failure is logged. + +The shield capabilities (``pydantic_ai_lightspeed.capabilities``) are the one +writer outside this module. They store the turn they rejected from inside the +agent run, and only when the model was handed the conversation, so never in +compacted mode. +""" + +from collections.abc import AsyncIterator, Sequence +from contextlib import asynccontextmanager +from dataclasses import dataclass, field +from typing import Optional + +from ogx_api import OpenAIResponseOutput +from ogx_api.openai_responses import OpenAIResponseMessage +from ogx_client import AsyncOgxClient + +from models.common.responses.responses_api_params import ResponsesApiParams +from models.common.responses.types import ResponseInput +from utils.conversations import append_turn_items_to_conversation + + +class TurnNotStoredError(Exception): + """Nobody tried to store a turn lightspeed-stack has to store. + + Deliberately not a ``RuntimeError``: the endpoints catch that one around + the model call and report it as an inference failure, which this is not. + """ + + +@dataclass +class PendingTurn: + """The turn of the request being served, until it is settled. + + A turn is settled by the first of :meth:`store_completed`, + :meth:`store_blocked` and :meth:`store_interrupted` that is called. Every + later call does nothing, so the write happens once when several paths of a + request want it (the end of a stream, the cancellation handler, the + interrupt callback). + + Attributes: + client: OGX client used for the write. + conversation_id: Conversation the turn belongs to (OGX format). + user_input: The input as it arrived. This is what is stored as the + user's side of the turn; in compacted mode it differs from + ``params.input``, which holds the explicit rewrite. + store: Whether the request wants its turn stored at all. + left_to_ogx: Whether OGX stores a completed turn itself, which it does + when it gets the ``conversation`` parameter. + outcome: How the turn was settled; ``None`` while it is pending. A + write that failed is recorded as ``failed: ``. + """ + + client: AsyncOgxClient + conversation_id: str + user_input: ResponseInput + store: bool = True + left_to_ogx: bool = False + outcome: Optional[str] = field(default=None, init=False) + + @classmethod + def for_request( + cls, + client: AsyncOgxClient, + params: ResponsesApiParams, + original_input: Optional[ResponseInput] = None, + ) -> "PendingTurn": + """Create the pending turn of a request from its prepared parameters. + + Parameters: + client: OGX client used for the write. + params: The parameters the request is sent with, after compaction + was applied. + original_input: The input before the explicit-input rewrite, as + the compaction result carries it. Required in compacted mode. + + Returns: + The pending turn of the request. + + Raises: + ValueError: When the request is compacted and its original input + is missing; the turn could then only be stored against the + explicit rewrite, which would write the summaries and the + replayed history into the conversation. + """ + if not params.omit_conversation: + return cls( + client=client, + conversation_id=params.conversation, + user_input=params.input, + store=params.store, + left_to_ogx=not params.previous_response_id, + ) + if original_input is None: + raise ValueError( + "a request served in compacted mode needs its original input " + "to store the turn" + ) + return cls( + client=client, + conversation_id=params.conversation, + user_input=original_input, + store=params.store, + ) + + @property + def settled(self) -> bool: + """Whether the turn has been taken care of. + + That is the case once a write was attempted, whether or not it + succeeded, or the turn was left to OGX. + """ + return self.outcome is not None + + @property + def ours(self) -> bool: + """Whether lightspeed-stack has to store the turn when it completes.""" + return self.store and not self.left_to_ogx + + async def store_completed( + self, output_items: Sequence[OpenAIResponseOutput] + ) -> bool: + """Store a completed turn, unless OGX stores it. + + Parameters: + output_items: The output items of the turn, as OGX returned them. + + Returns: + Whether this call stored the turn. + + Raises: + HTTPException: When the write fails. + """ + if self.left_to_ogx: + self._settle("left to OGX") + return False + return await self._store("completed", output_items) + + async def store_blocked(self, refusal: OpenAIResponseMessage) -> bool: + """Store the turn of a request a shield blocked. + + Parameters: + refusal: The refusal message returned in place of an answer. + + Returns: + Whether this call stored the turn. + + Raises: + HTTPException: When the write fails. + """ + return await self._store("blocked", [refusal]) + + async def store_interrupted(self, partial_response: str) -> bool: + """Store the turn of a stream the client interrupted. + + Parameters: + partial_response: The part of the answer that was received, with + the interruption notice. + + Returns: + Whether this call stored the turn. + + Raises: + HTTPException: When the write fails. + """ + return await self._store( + "interrupted", + [OpenAIResponseMessage(role="assistant", content=partial_response)], + ) + + def ensure_settled(self) -> None: + """Fail when lightspeed-stack has to store the turn and nobody tried to. + + A turn OGX stores needs no check, and neither does one the request + does not want stored. A write that was tried and failed passes the + check: the failure was raised or logged where it happened. + + Raises: + TurnNotStoredError: When the turn has to be stored by + lightspeed-stack and is still pending. + """ + if self.ours and not self.settled: + raise TurnNotStoredError( + f"the turn on conversation {self.conversation_id} was served " + "without the conversation parameter, and nobody tried to store " + "it; the next request would not see it" + ) + + def _settle(self, outcome: str) -> bool: + """Settle the turn; only the first caller gets ``True``. + + There is no ``await`` between the check and the assignment, so of + several tasks on the event loop exactly one settles the turn. + """ + if self.outcome is not None: + return False + self.outcome = outcome + return True + + async def _store( + self, outcome: str, output_items: Sequence[OpenAIResponseOutput] + ) -> bool: + """Settle the turn and append it to the conversation. + + The turn is settled before the write starts. A write that fails is + therefore not repeated by another path of the same request: the turn + is lost, as it was before this class existed, but never doubled. + + Parameters: + outcome: How the turn ended; recorded when this call settles it. + output_items: The assistant's side of the turn. + + Returns: + Whether this call wrote the turn. ``False`` when the turn was + settled before, or the request does not want its turn stored. + + Raises: + HTTPException: When the write fails. The outcome is then recorded + as ``failed: ``. + """ + if not self._settle(outcome) or not self.store: + return False + try: + await append_turn_items_to_conversation( + self.client, self.conversation_id, self.user_input, output_items + ) + except BaseException: + self.outcome = f"failed: {outcome}" + raise + return True + + +@asynccontextmanager +async def pending_turn( + client: AsyncOgxClient, + params: ResponsesApiParams, + original_input: Optional[ResponseInput] = None, +) -> AsyncIterator[PendingTurn]: + """Serve a request inside a scope that does not let its turn get lost. + + The scope hands out the pending turn of the request. When it is left + without an error, the turn has to be settled. + + Parameters: + client: OGX client used for the write. + params: The parameters the request is sent with, after compaction was + applied. + original_input: The input before the explicit-input rewrite; required + in compacted mode. + + Yields: + The pending turn of the request. + + Raises: + TurnNotStoredError: When the scope is left normally with a turn of + ours still pending. + ValueError: When the request is compacted and its original input is + missing. + """ + turn = PendingTurn.for_request(client, params, original_input) + yield turn + turn.ensure_settled() diff --git a/src/utils/stream_interrupts.py b/src/utils/stream_interrupts.py index 4c56f0410..99b807bf8 100644 --- a/src/utils/stream_interrupts.py +++ b/src/utils/stream_interrupts.py @@ -6,9 +6,7 @@ from dataclasses import dataclass, field from enum import StrEnum from threading import Lock -from typing import Any, Optional, cast - -from ogx_api import OpenAIResponseMessage +from typing import Any, Optional from constants import ( INTERRUPTED_RESPONSE_MESSAGE, @@ -17,13 +15,9 @@ from log import get_logger from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams -from models.common.responses.types import ResponseInput from models.common.turn_summary import TurnSummary -from utils.conversations import ( - append_turn_items_to_conversation, - append_turn_to_conversation, -) from utils.markdown_repair import close_open_markdown +from utils.pending_turn import PendingTurn from utils.query import store_query_results, update_conversation_topic_summary from utils.responses import get_topic_summary from utils.types import Singleton @@ -254,7 +248,7 @@ async def persist_interrupted_turn( responses_params: ResponsesApiParams, turn_summary: TurnSummary, background_topic_summary_tasks: list[asyncio.Task[None]], - original_input: Optional[ResponseInput] = None, + turn: Optional[PendingTurn] = None, ) -> None: """Persist the user query and an interrupted response into the conversation. @@ -271,31 +265,18 @@ async def persist_interrupted_turn( interrupted message. background_topic_summary_tasks: Mutable list tracking fire-and-forget topic summary tasks for graceful shutdown. - original_input: In compacted mode, the original user input before the - explicit-input rewrite. When set, the turn is persisted against it - (the ``conversation`` parameter was dropped, and - ``responses_params.input`` is the explicit rewrite); ``None`` - otherwise (LCORE-1572). + turn: The pending turn of the request, which stores the turn against + the input as it arrived and does so once (LCORE-3908). Required in + compacted mode; created from the parameters otherwise. + + Raises: + ------ + ValueError: When the request is compacted and no pending turn is given. """ + if turn is None: + turn = PendingTurn.for_request(context.client, responses_params) try: - if original_input is not None: - await append_turn_items_to_conversation( - context.client, - responses_params.conversation, - original_input, - [ - OpenAIResponseMessage( - role="assistant", content=turn_summary.llm_response - ) - ], - ) - else: - await append_turn_to_conversation( - context.client, - responses_params.conversation, - cast("str", responses_params.input), - turn_summary.llm_response, - ) + await turn.store_interrupted(turn_summary.llm_response) except Exception: # pylint: disable=broad-except logger.exception( "Failed to append interrupted turn to conversation for request %s", @@ -342,7 +323,7 @@ def register_interrupt_callback( responses_params: ResponsesApiParams, turn_summary: TurnSummary, background_topic_summary_tasks: list[asyncio.Task[None]], - original_input: Optional[ResponseInput] = None, + turn: Optional[PendingTurn] = None, ) -> list[bool]: """Build an interrupt callback and register the stream for cancellation. @@ -361,8 +342,7 @@ def register_interrupt_callback( turn_summary: TurnSummary populated during streaming. background_topic_summary_tasks: Mutable list tracking fire-and-forget topic summary tasks for graceful shutdown. - original_input: In compacted mode, the original user input before the - explicit-input rewrite; ``None`` otherwise. + turn: The pending turn of the request; required in compacted mode. Returns: ------- @@ -383,7 +363,7 @@ async def _on_interrupt() -> None: responses_params, turn_summary, background_topic_summary_tasks, - original_input, + turn, ) current_task = asyncio.current_task() diff --git a/tests/integration/endpoints/_compaction_helpers.py b/tests/integration/endpoints/_compaction_helpers.py index f5fea46c5..d977c7b72 100644 --- a/tests/integration/endpoints/_compaction_helpers.py +++ b/tests/integration/endpoints/_compaction_helpers.py @@ -3,9 +3,19 @@ # pylint: disable=import-outside-toplevel import asyncio +import json +import uuid +from collections.abc import AsyncIterator from typing import Any, cast +from a2a.types import ( + AgentCapabilities, + AgentCard, + AgentProvider, +) +from fastapi import Request from ogx_api.openai_responses import OpenAIResponseMessage +from pydantic_ai import AgentRunResultEvent from pytest_mock import AsyncMockType, MockerFixture from sqlalchemy.orm import Session @@ -278,3 +288,89 @@ def create_existing_conversation( ) db_session.add(conv) db_session.commit() + + +FAKE_AGENT_CARD = AgentCard( + name="Test Agent", + description="Test", + version="0.0.1", + url="http://localhost:8080/a2a", + provider=AgentProvider(organization="test", url="http://test"), + skills=[], + default_input_modes=["text/plain"], + default_output_modes=["text/plain"], + capabilities=AgentCapabilities(streaming=False), + protocol_version="0.3.0", +) + + +def build_a2a_request(user_input: str) -> Request: + """Build a FastAPI Request with a JSON-RPC ``message/send`` body. + + Args: + user_input: The user message text to include in the A2A request. + + Returns: + A FastAPI Request object with the JSON-RPC body ready for consumption. + """ + body_dict = { + "jsonrpc": "2.0", + "id": str(uuid.uuid4()), + "method": "message/send", + "params": { + "message": { + "role": "user", + "parts": [{"type": "text", "text": user_input}], + "messageId": f"msg-{uuid.uuid4()}", + "contextId": f"ctx-{uuid.uuid4()}", + } + }, + } + body_bytes = json.dumps(body_dict).encode() + + async def receive() -> dict[str, Any]: + """Return the pre-built body as an ASGI receive event.""" + return {"type": "http.request", "body": body_bytes, "more_body": False} + + return Request( + scope={ + "type": "http", + "method": "POST", + "path": "/a2a", + "root_path": "", + "query_string": b"", + "headers": [ + (b"content-type", b"application/json"), + ], + }, + receive=receive, + ) + + +def mock_a2a_agent(mocker: MockerFixture) -> Any: + """Build a mock pydantic-ai agent that yields a single result event. + + Args: + mocker: pytest-mock fixture. + + Returns: + A mock agent whose ``run_stream_events`` returns a single result event. + """ + mock_run_result = mocker.MagicMock() + mock_run_result.response.text = "Test A2A response" + result_event = mocker.MagicMock(spec=AgentRunResultEvent) + result_event.result = mock_run_result + + async def _event_stream() -> AsyncIterator[Any]: + """Yield a single agent run result event.""" + yield result_event + + mock_stream_ctx = mocker.AsyncMock() + mock_stream_ctx.__aenter__ = mocker.AsyncMock(return_value=_event_stream()) + mock_stream_ctx.__aexit__ = mocker.AsyncMock(return_value=False) + mock_agent = mocker.MagicMock() + mock_agent.run_stream_events.return_value = mock_stream_ctx + mock_agent.model.last_output_items = [ + OpenAIResponseMessage(role="assistant", content=DEFAULT_MODEL_RESPONSE) + ] + return mock_agent diff --git a/tests/integration/endpoints/test_compaction_a2a.py b/tests/integration/endpoints/test_compaction_a2a.py index a108b47ea..044ddfc07 100644 --- a/tests/integration/endpoints/test_compaction_a2a.py +++ b/tests/integration/endpoints/test_compaction_a2a.py @@ -4,20 +4,9 @@ # pylint: disable=too-many-positional-arguments import asyncio -import json -import uuid -from collections.abc import AsyncIterator from typing import Any import pytest -from a2a.types import ( - AgentCapabilities, - AgentCard, - AgentProvider, -) -from fastapi import Request -from ogx_api.openai_responses import OpenAIResponseMessage -from pydantic_ai import AgentRunResultEvent from pytest_mock import AsyncMockType, MockerFixture from app.endpoints.a2a import handle_a2a_jsonrpc_post @@ -30,102 +19,20 @@ CONV_ID_LLAMA, DEFAULT_MODEL_RESPONSE, DEFAULT_SUMMARY_TEXT, + FAKE_AGENT_CARD, TEST_MODEL, assert_marker_count, await_lock_contention, + build_a2a_request, collect_items, enable_compaction, marker, + mock_a2a_agent, msg, patch_get_all_conversation_items, verify_store_content, ) -_FAKE_AGENT_CARD = AgentCard( - name="Test Agent", - description="Test", - version="0.0.1", - url="http://localhost:8080/a2a", - provider=AgentProvider(organization="test", url="http://test"), - skills=[], - default_input_modes=["text/plain"], - default_output_modes=["text/plain"], - capabilities=AgentCapabilities(streaming=False), - protocol_version="0.3.0", -) - - -def _build_a2a_request(user_input: str) -> Request: - """Build a FastAPI Request with a JSON-RPC ``message/send`` body. - - Args: - user_input: The user message text to include in the A2A request. - - Returns: - A FastAPI Request object with the JSON-RPC body ready for consumption. - """ - body_dict = { - "jsonrpc": "2.0", - "id": str(uuid.uuid4()), - "method": "message/send", - "params": { - "message": { - "role": "user", - "parts": [{"type": "text", "text": user_input}], - "messageId": f"msg-{uuid.uuid4()}", - "contextId": f"ctx-{uuid.uuid4()}", - } - }, - } - body_bytes = json.dumps(body_dict).encode() - - async def receive() -> dict[str, Any]: - """Return the pre-built body as an ASGI receive event.""" - return {"type": "http.request", "body": body_bytes, "more_body": False} - - return Request( - scope={ - "type": "http", - "method": "POST", - "path": "/a2a", - "root_path": "", - "query_string": b"", - "headers": [ - (b"content-type", b"application/json"), - ], - }, - receive=receive, - ) - - -def _mock_a2a_agent(mocker: MockerFixture) -> Any: - """Build a mock pydantic-ai agent that yields a single result event. - - Args: - mocker: pytest-mock fixture. - - Returns: - A mock agent whose ``run_stream_events`` returns a single result event. - """ - mock_run_result = mocker.MagicMock() - mock_run_result.response.text = "Test A2A response" - result_event = mocker.MagicMock(spec=AgentRunResultEvent) - result_event.result = mock_run_result - - async def _event_stream() -> AsyncIterator[Any]: - """Yield a single agent run result event.""" - yield result_event - - mock_stream_ctx = mocker.AsyncMock() - mock_stream_ctx.__aenter__ = mocker.AsyncMock(return_value=_event_stream()) - mock_stream_ctx.__aexit__ = mocker.AsyncMock(return_value=False) - mock_agent = mocker.MagicMock() - mock_agent.run_stream_events.return_value = mock_stream_ctx - mock_agent.model.last_output_items = [ - OpenAIResponseMessage(role="assistant", content=DEFAULT_MODEL_RESPONSE) - ] - return mock_agent - def _setup_a2a_compaction_mocks( mocker: MockerFixture, @@ -148,7 +55,7 @@ def _setup_a2a_compaction_mocks( """ mocker.patch( "app.endpoints.a2a.get_lightspeed_agent_card", - return_value=_FAKE_AGENT_CARD, + return_value=FAKE_AGENT_CARD, ) async def _fake_prepare(client, query_request, *args, **kwargs): @@ -167,7 +74,7 @@ async def _fake_prepare(client, query_request, *args, **kwargs): side_effect=_fake_prepare, ) - mock_agent = _mock_a2a_agent(mocker) + mock_agent = mock_a2a_agent(mocker) mock_build_agent = mocker.patch( "app.endpoints.a2a.build_agent", return_value=mock_agent, @@ -222,7 +129,7 @@ async def test_a2a_compaction_triggers_summarization( mock_summarize, mock_build_agent = _setup_a2a_compaction_mocks(mocker, items) - request = _build_a2a_request("What else can you help with?") + request = build_a2a_request("What else can you help with?") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) mock_summarize.assert_awaited_once() @@ -287,7 +194,7 @@ async def test_a2a_compaction_partition( mock_summarize, mock_build_agent = _setup_a2a_compaction_mocks(mocker, items) - request = _build_a2a_request("What else can you help with?") + request = build_a2a_request("What else can you help with?") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) mock_summarize.assert_awaited_once() @@ -349,7 +256,7 @@ async def test_a2a_compaction_existing_marker_no_new_summarization( mock_summarize, mock_build_agent = _setup_a2a_compaction_mocks(mocker, items) - request = _build_a2a_request("Any updates?") + request = build_a2a_request("Any updates?") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) mock_summarize.assert_not_called() @@ -406,7 +313,7 @@ async def test_a2a_compaction_small_conversation_no_compaction( mock_summarize, mock_build_agent = _setup_a2a_compaction_mocks(mocker, items) - request = _build_a2a_request("short question") + request = build_a2a_request("short question") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) mock_summarize.assert_not_called() @@ -434,7 +341,7 @@ async def test_a2a_compaction_disabled_passes_through( _, mock_build_agent = _setup_a2a_compaction_mocks(mocker, []) - request = _build_a2a_request("What is Ansible?") + request = build_a2a_request("What is Ansible?") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) agent_params = mock_build_agent.call_args[0][1] @@ -472,7 +379,7 @@ async def test_a2a_compaction_additive_summarization( mock_summarize, mock_build_agent = _setup_a2a_compaction_mocks(mocker, items) # --- Round 1 --- - request = _build_a2a_request("What else can you help with?") + request = build_a2a_request("What else can you help with?") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) mock_summarize.assert_awaited_once() @@ -511,7 +418,7 @@ async def test_a2a_compaction_additive_summarization( mock_summarize.reset_mock() - request = _build_a2a_request("Follow-up question") + request = build_a2a_request("Follow-up question") await handle_a2a_jsonrpc_post(request=request, auth=test_auth, mcp_headers={}) mock_summarize.assert_awaited_once() @@ -565,13 +472,13 @@ async def test_a2a_compaction_blocking_concurrent_request_with_same_id( entered, release, task2_entered = patch_get_all_conversation_items(mocker) - request1 = _build_a2a_request("What is Ansible?") + request1 = build_a2a_request("What is Ansible?") task1 = asyncio.create_task( handle_a2a_jsonrpc_post(request=request1, auth=test_auth, mcp_headers={}) ) await entered.wait() - request2 = _build_a2a_request("What is RHEL?") + request2 = build_a2a_request("What is RHEL?") task2 = asyncio.create_task( handle_a2a_jsonrpc_post(request=request2, auth=test_auth, mcp_headers={}) ) diff --git a/tests/integration/endpoints/test_turn_persistence.py b/tests/integration/endpoints/test_turn_persistence.py new file mode 100644 index 000000000..2e6cd81ec --- /dev/null +++ b/tests/integration/endpoints/test_turn_persistence.py @@ -0,0 +1,1358 @@ +"""Integration tests for the turns lightspeed-stack stores itself (LCORE-3908). + +OGX stores a turn only when it is handed the ``conversation`` parameter and +runs the inference. In every other case lightspeed-stack appends the turn to +the conversation: a conversation served in compacted mode, a request a shield +blocked, a stream the client interrupted, a continuation from +``previous_response_id``. + +LCORE-3883 showed what happens when that write goes missing: nothing fails, +the conversation just stops growing. So every test here asserts on what the +conversation store holds after the request, item by item. That catches a write +that is missing, a write that happens twice, and a write of the wrong input. +""" + +# pylint: disable=too-many-arguments +# pylint: disable=too-many-positional-arguments +# pylint: disable=too-many-lines + +import asyncio +import json +from collections.abc import AsyncIterator +from datetime import UTC, datetime +from types import SimpleNamespace +from typing import Any, Optional + +import pytest +from fastapi import HTTPException, Request +from fastapi.responses import StreamingResponse +from ogx_api.openai_responses import OpenAIResponseMessage +from ogx_client import ApiException +from ogx_client.models.open_ai_response_object import OpenAIResponseObject +from ogx_client.models.open_ai_response_object_stream_response_completed import ( + OpenAIResponseObjectStreamResponseCompleted, +) +from ogx_client.models.open_ai_response_object_stream_response_created import ( + OpenAIResponseObjectStreamResponseCreated, +) +from ogx_client.models.open_ai_response_object_stream_response_failed import ( + OpenAIResponseObjectStreamResponseFailed, +) +from ogx_client.models.open_ai_response_object_stream_response_incomplete import ( + OpenAIResponseObjectStreamResponseIncomplete, +) +from pydantic_ai import AgentRunResultEvent +from pydantic_ai.messages import ModelResponse, PartStartEvent, TextPart +from pytest_mock import AsyncMockType, MockerFixture +from sqlalchemy.orm import Session + +from app.endpoints.a2a import handle_a2a_jsonrpc_post +from app.endpoints.query import query_endpoint_handler +from app.endpoints.responses import responses_endpoint_handler +from app.endpoints.streaming_query import streaming_query_endpoint_handler +from authentication.interface import AuthTuple +from configuration import AppConfig +from models.api.requests import QueryRequest, ResponsesRequest +from models.common.moderation import ShieldModerationBlocked +from models.common.responses.contexts import ResponsesContext +from models.common.responses.responses_api_params import ResponsesApiParams +from models.common.turn_summary import TurnSummary +from models.database.conversations import UserConversation, UserTurn +from tests.integration.conftest import ( + InMemoryConversationStore, + create_agent_run_result, + make_openai_response_object, + mock_agent_run_stream, +) +from tests.integration.endpoints._compaction_helpers import ( + CONV_ID_LLAMA, + DEFAULT_MODEL_RESPONSE, + EXISTING_CONV_ID, + FAKE_AGENT_CARD, + TEST_MODEL, + build_a2a_request, + collect_items, + create_existing_conversation, + enable_compaction, + marker, + mock_a2a_agent, + msg, +) +from utils.pending_turn import TurnNotStoredError +from utils.stream_interrupts import ( + CancelStreamResult, + build_interrupted_response, + get_stream_interrupt_registry, +) +from utils.token_estimator import extract_message_text + +NEW_QUERY = "What else can you help with?" +REFUSAL = "Content blocked by safety shield" +PARTIAL_ANSWER = "Ansible is" +PREVIOUS_RESPONSE_ID = "resp_previous_turn" +LARGE_WINDOW = 100_000 +"""Context window no test conversation comes near, so nothing is summarized.""" + + +def _compacted_conversation() -> list[OpenAIResponseMessage]: + """Return a stored conversation that has been compacted before. + + It holds a summary marker, which is what puts every later request on it in + compacted mode: the ``conversation`` parameter is dropped and OGX does not + store the turn. + """ + return [ + msg("user", "earlier question"), + msg("assistant", "earlier answer"), + marker("summary of the earlier turn"), + ] + + +def _plain_conversation() -> list[OpenAIResponseMessage]: + """Return a stored conversation that has never been compacted.""" + return [ + msg("user", "earlier question"), + msg("assistant", "earlier answer"), + ] + + +def _turn(answer: str, query: str = NEW_QUERY) -> list[OpenAIResponseMessage]: + """Return the two items one stored turn consists of.""" + return [msg("user", query), msg("assistant", answer)] + + +async def _seed( + test_config: AppConfig, + store: InMemoryConversationStore, + items: list[OpenAIResponseMessage], +) -> None: + """Enable compaction and store the conversation the request continues.""" + enable_compaction(test_config, context_window=LARGE_WINDOW) + await store.create(conversation_id=CONV_ID_LLAMA, items=items) + + +def _role_and_text(item: Any) -> tuple[str, str]: + """Reduce a stored item to who said it and what was said. + + The text is compared, not the content object: OGX returns the text of an + answer as a list of content parts, a request carries it as a string. + """ + return str(getattr(item, "role", "")), extract_message_text(item) + + +async def _assert_stored( + store: InMemoryConversationStore, expected: list[OpenAIResponseMessage] +) -> None: + """Assert the conversation holds exactly the expected items, in order.""" + stored = await collect_items(store, CONV_ID_LLAMA) + assert [_role_and_text(item) for item in stored] == [ + _role_and_text(item) for item in expected + ] + + +def _blocked() -> ShieldModerationBlocked: + """Return the verdict of a shield that blocked the request.""" + return ShieldModerationBlocked(message=REFUSAL, moderation_id="modr_blocked_1") + + +def _agent_answer(agent: Any, text: str = DEFAULT_MODEL_RESPONSE) -> None: + """Set the output items the agent's model captured for the turn.""" + agent.model.last_output_items = [ + OpenAIResponseMessage(role="assistant", content=text) + ] + + +def _run_cut_short(mocker: MockerFixture) -> Any: + """Return the result of an agent run the model ended for length, not success.""" + return create_agent_run_result( + mocker, + model_response=ModelResponse( + parts=[TextPart(PARTIAL_ANSWER)], + finish_reason="length", + provider_response_id="response-cut-short", + ), + ) + + +async def _drain(response: Any) -> list[str]: + """Read a streaming response to its end and return the chunks.""" + assert isinstance(response, StreamingResponse) + return [str(chunk) async for chunk in response.body_iterator] + + +# ========================================== +# /v1/query +# ========================================== + + +async def _send_query(test_request: Request, test_auth: AuthTuple) -> Any: + """Send the new query on the stored conversation.""" + return await query_endpoint_handler( + request=test_request, + query_request=QueryRequest(query=NEW_QUERY, conversation_id=EXISTING_CONV_ID), + auth=test_auth, + mcp_headers={}, + ) + + +class TestQueryTurnPersistence: + """What /v1/query leaves in the conversation.""" + + @pytest.mark.asyncio + async def test_completed_turn_in_compacted_mode_is_stored_once( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """The conversation gains the query as it arrived and the answer, once.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_query_agent) + + await _send_query(test_request, test_auth) + + params = mock_query_agent.build_agent_mock.call_args[0][1] + assert params.omit_conversation is True + await _assert_stored( + mock_conversation_store, + _compacted_conversation() + _turn(DEFAULT_MODEL_RESPONSE), + ) + + @pytest.mark.asyncio + async def test_two_requests_store_two_turns( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """Each request adds its own turn and nothing else.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_query_agent) + + await _send_query(test_request, test_auth) + await _send_query(test_request, test_auth) + + await _assert_stored( + mock_conversation_store, + _compacted_conversation() + + _turn(DEFAULT_MODEL_RESPONSE) + + _turn(DEFAULT_MODEL_RESPONSE), + ) + + @pytest.mark.asyncio + async def test_completed_turn_outside_compacted_mode_is_left_to_ogx( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """With the conversation parameter sent, OGX stores the turn, not we.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + _agent_answer(mock_query_agent) + + await _send_query(test_request, test_auth) + + params = mock_query_agent.build_agent_mock.call_args[0][1] + assert params.omit_conversation is False + await _assert_stored(mock_conversation_store, _plain_conversation()) + + @pytest.mark.asyncio + async def test_failed_turn_in_compacted_mode_is_not_stored( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """A request that fails in the model call stores nothing.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_query_agent) + mock_query_agent.run.side_effect = RuntimeError("the model call failed") + + with pytest.raises(HTTPException): + await _send_query(test_request, test_auth) + + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + @pytest.mark.asyncio + async def test_turn_the_model_cut_short_is_not_stored( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """A run that did not finish with success is an error here, and stores nothing. + + /v1/streaming_query treats the same run differently: see + ``test_turn_the_model_cut_short_is_stored_once`` there. + """ + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_query_agent) + mock_query_agent.run.return_value = _run_cut_short(mocker) + + with pytest.raises(HTTPException): + await _send_query(test_request, test_auth) + + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + +# ========================================== +# /v1/streaming_query +# ========================================== + + +def _stream_that_stalls(release: asyncio.Event) -> Any: + """Build an agent stream that sends one token and then waits. + + The wait ends only when the consuming task is cancelled, which is what an + interrupt does. + """ + + async def _events() -> AsyncIterator[Any]: + yield PartStartEvent(index=0, part=TextPart(content=PARTIAL_ANSWER)) + await release.wait() + + class _RunStreamCtx: + """Async context manager matching ``agent.run_stream_events``.""" + + async def __aenter__(self) -> AsyncIterator[Any]: + return _events() + + async def __aexit__(self, *_args: object) -> None: + return None + + return _RunStreamCtx() + + +def _request_id_of(chunks: list[str]) -> str: + """Return the request id the stream announced in its start event.""" + for chunk in chunks: + for line in chunk.splitlines(): + if not line.startswith("data: "): + continue + event = json.loads(line[len("data: ") :]) + if event.get("event") == "start": + return event["data"]["request_id"] + raise AssertionError(f"no start event in {chunks!r}") + + +async def _send_streaming_query(test_request: Request, test_auth: AuthTuple) -> Any: + """Send the new query on the stored conversation.""" + return await streaming_query_endpoint_handler( + request=test_request, + query_request=QueryRequest(query=NEW_QUERY, conversation_id=EXISTING_CONV_ID), + auth=test_auth, + mcp_headers={}, + ) + + +async def _interrupt_stream(test_request: Request, test_auth: AuthTuple) -> list[str]: + """Start a stream, interrupt it after its first token, read it to the end.""" + response = await _send_streaming_query(test_request, test_auth) + assert isinstance(response, StreamingResponse) + chunks: list[str] = [] + first_token = asyncio.Event() + + async def _consume() -> None: + async for chunk in response.body_iterator: + chunks.append(str(chunk)) + if '"event": "token"' in str(chunk): + first_token.set() + + consumer = asyncio.create_task(_consume()) + await asyncio.wait_for(first_token.wait(), timeout=5) + result = get_stream_interrupt_registry().cancel_stream( + _request_id_of(chunks), test_auth[0] + ) + assert result == CancelStreamResult.CANCELLED + await asyncio.wait_for(consumer, timeout=5) + # The interrupt callback runs as a task of its own; let it finish. + pending = [ + task + for task in asyncio.all_tasks() + if task is not asyncio.current_task() and not task.done() + ] + if pending: + await asyncio.wait(pending, timeout=5) + return chunks + + +class TestStreamingQueryTurnPersistence: + """What /v1/streaming_query leaves in the conversation.""" + + @pytest.mark.asyncio + async def test_completed_turn_in_compacted_mode_is_stored_once( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """The conversation gains the query as it arrived and the answer, once.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_streaming_query_agent) + + await _drain(await _send_streaming_query(test_request, test_auth)) + + params = mock_streaming_query_agent.build_agent_mock.call_args[0][1] + assert params.omit_conversation is True + await _assert_stored( + mock_conversation_store, + _compacted_conversation() + _turn(DEFAULT_MODEL_RESPONSE), + ) + + @pytest.mark.asyncio + async def test_completed_turn_outside_compacted_mode_is_left_to_ogx( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """With the conversation parameter sent, OGX stores the turn, not we.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + _agent_answer(mock_streaming_query_agent) + + await _drain(await _send_streaming_query(test_request, test_auth)) + + params = mock_streaming_query_agent.build_agent_mock.call_args[0][1] + assert params.omit_conversation is False + await _assert_stored(mock_conversation_store, _plain_conversation()) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "stored_conversation", + [_compacted_conversation, _plain_conversation], + ids=["compacted", "not-compacted"], + ) + async def test_interrupted_turn_is_stored_once( + self, + stored_conversation: Any, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """An interrupt stores what was received so far, once. + + Two things react to an interrupt: the generator, where the + cancellation lands, and the callback the interrupt endpoint schedules. + Both want to store the turn; only one of them may. + """ + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, stored_conversation()) + _agent_answer(mock_streaming_query_agent) + mock_streaming_query_agent.run_stream_events.return_value = _stream_that_stalls( + asyncio.Event() + ) + + chunks = await _interrupt_stream(test_request, test_auth) + + assert any('"event": "interrupted"' in chunk for chunk in chunks) + interrupted_answer, _ = build_interrupted_response([PARTIAL_ANSWER]) + await _assert_stored( + mock_conversation_store, + stored_conversation() + _turn(interrupted_answer), + ) + recorded_turns = ( + patch_db_session.query(UserTurn) + .filter_by(conversation_id=EXISTING_CONV_ID) + .count() + ) + assert recorded_turns == 1 + + @pytest.mark.asyncio + async def test_turn_the_model_cut_short_is_stored_once( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """A run that did not finish with success still ran to its end, and is stored. + + The stream reports the error as an event and ends normally, so the + turn is stored with the output received. /v1/query stores nothing + for the same run. + """ + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_streaming_query_agent, PARTIAL_ANSWER) + mock_streaming_query_agent.run_stream_events.return_value = ( + mock_agent_run_stream( + [ + PartStartEvent(index=0, part=TextPart(content=PARTIAL_ANSWER)), + AgentRunResultEvent(result=_run_cut_short(mocker)), + ] + ) + ) + + chunks = await _drain(await _send_streaming_query(test_request, test_auth)) + + assert any('"event": "error"' in chunk for chunk in chunks) + await _assert_stored( + mock_conversation_store, _compacted_conversation() + _turn(PARTIAL_ANSWER) + ) + + @pytest.mark.asyncio + async def test_stream_the_client_stopped_reading_stores_nothing( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """A stream closed before its end leaves no turn behind. + + The turn is stored after the last event of the stream. A client that + goes away before that closes the generator, and the write never runs. + This is the behaviour as it is, recorded here, not a goal. + """ + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_streaming_query_agent) + + response = await _send_streaming_query(test_request, test_auth) + assert isinstance(response, StreamingResponse) + body: Any = response.body_iterator + async for chunk in body: + if '"event": "token"' in str(chunk): + break + await body.aclose() + + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + +# ========================================== +# /v1/responses +# ========================================== + + +def _link_previous_response(db_session: Session) -> None: + """Make the stored conversation end with a response a request can continue.""" + conversation = ( + db_session.query(UserConversation).filter_by(id=EXISTING_CONV_ID).one() + ) + conversation.last_response_id = PREVIOUS_RESPONSE_ID + now = datetime.now(UTC) + db_session.add( + UserTurn( + conversation_id=EXISTING_CONV_ID, + turn_number=1, + started_at=now, + completed_at=now, + provider="test-provider", + model="test-model", + response_id=PREVIOUS_RESPONSE_ID, + ) + ) + db_session.commit() + + +def _response_that_ended_with(terminal_event: Optional[str]) -> Any: + """Build the response object a stream carries in its terminal event. + + A completed response holds the answer, an incomplete one the part of it + that was produced, a failed one an error and no output. + """ + if terminal_event == "response.failed": + failed = OpenAIResponseObject.from_dict( + { + "id": "response-failed", + "object": "response", + "created_at": 1_700_000_000, + "status": "failed", + "model": TEST_MODEL, + "store": False, + "output": [], + "error": {"code": "server_error", "message": "the model failed"}, + } + ) + assert failed is not None + return failed + if terminal_event == "response.incomplete": + return make_openai_response_object(content=PARTIAL_ANSWER) + return make_openai_response_object(content=DEFAULT_MODEL_RESPONSE) + + +TERMINAL_EVENTS: dict[str, Any] = { + "response.completed": OpenAIResponseObjectStreamResponseCompleted, + "response.incomplete": OpenAIResponseObjectStreamResponseIncomplete, + "response.failed": OpenAIResponseObjectStreamResponseFailed, +} + + +async def _one_chunk_stream( + response_object: Any, + terminal_event: Optional[str] = "response.completed", +) -> AsyncIterator[Any]: + """Yield the events of a streamed response, ending with the terminal one. + + Without a terminal event the stream consists of the opening event alone, + which is what a response that never finished looks like. + """ + yield OpenAIResponseObjectStreamResponseCreated( + response=response_object, + sequence_number=0, + type="response.created", + ) + if terminal_event is not None: + yield TERMINAL_EVENTS[terminal_event]( + response=response_object, + sequence_number=1, + type=terminal_event, + ) + + +@pytest.fixture(name="ogx_stream") +def ogx_stream_fixture( + mock_ogx_client: AsyncMockType, + mocker: MockerFixture, +) -> SimpleNamespace: + """Let the real handlers run against the mocked OGX client. + + Returns: + The settings of the mocked stream; ``terminal_event`` is the event + a streamed response ends with, ``None`` for a stream without one. + """ + ogx_stream = SimpleNamespace(terminal_event="response.completed") + original_context = ResponsesContext + + def _skip_validation(**kwargs: Any) -> ResponsesContext: + """Build the context without validating the mocked client.""" + return original_context.model_construct(**kwargs) + + mocker.patch( + "app.endpoints.responses.ResponsesContext", side_effect=_skip_validation + ) + mocker.patch( + "app.endpoints.responses.maybe_get_topic_summary", + new=mocker.AsyncMock(return_value=None), + ) + + async def _create(**kwargs: Any) -> Any: + """Answer like OGX: a response object, or a stream of events.""" + if kwargs.get("stream"): + return _one_chunk_stream( + _response_that_ended_with(ogx_stream.terminal_event), + ogx_stream.terminal_event, + ) + return make_openai_response_object(content=DEFAULT_MODEL_RESPONSE) + + mock_ogx_client.responses.create = mocker.AsyncMock(side_effect=_create) + return ogx_stream + + +async def _send_response_request( + test_request: Request, + test_auth: AuthTuple, + stream: bool, + store: bool = True, + previous_response_id: Optional[str] = None, +) -> Any: + """Send the new input on the stored conversation and read the answer.""" + response = await responses_endpoint_handler( + request=test_request, + responses_request=ResponsesRequest( + input=NEW_QUERY, + model=TEST_MODEL, + conversation=None if previous_response_id else EXISTING_CONV_ID, + previous_response_id=previous_response_id, + stream=stream, + store=store, + generate_topic_summary=False, + ), + auth=test_auth, + mcp_headers={}, + ) + if stream: + return await _drain(response) + return response + + +@pytest.mark.usefixtures("ogx_stream") +class TestResponsesTurnPersistence: + """What /v1/responses leaves in the conversation.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + async def test_completed_turn_in_compacted_mode_is_stored_once( + self, + stream: bool, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """The conversation gains the input as it arrived and the output, once.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + + await _send_response_request(test_request, test_auth, stream) + + sent = mock_ogx_client.responses.create.await_args.kwargs + assert "conversation" not in sent + await _assert_stored( + mock_conversation_store, + _compacted_conversation() + _turn(DEFAULT_MODEL_RESPONSE), + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + async def test_completed_turn_outside_compacted_mode_is_left_to_ogx( + self, + stream: bool, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """With the conversation parameter sent, OGX stores the turn, not we.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + + await _send_response_request(test_request, test_auth, stream) + + sent = mock_ogx_client.responses.create.await_args.kwargs + assert sent["conversation"] == CONV_ID_LLAMA + await _assert_stored(mock_conversation_store, _plain_conversation()) + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + @pytest.mark.parametrize( + "stored_conversation", + [_compacted_conversation, _plain_conversation], + ids=["compacted", "not-compacted"], + ) + async def test_blocked_turn_is_stored_once( + self, + stored_conversation: Any, + stream: bool, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """The refusal turn is stored once, against the input as it arrived.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, stored_conversation()) + mocker.patch( + "app.endpoints.responses.run_shield_moderation_v2", + return_value=_blocked(), + ) + + await _send_response_request(test_request, test_auth, stream) + + mock_ogx_client.responses.create.assert_not_awaited() + await _assert_stored( + mock_conversation_store, stored_conversation() + _turn(REFUSAL) + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + @pytest.mark.parametrize( + "stored_conversation", + [_compacted_conversation, _plain_conversation], + ids=["compacted", "not-compacted"], + ) + async def test_continuation_from_a_previous_response_is_stored_once( + self, + stored_conversation: Any, + stream: bool, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """OGX does not store a turn that continues from a previous response. + + Such a request is never compacted, so the turn is stored against its + input whichever state the conversation is in. + """ + create_existing_conversation(patch_db_session, test_auth[0]) + _link_previous_response(patch_db_session) + await _seed(test_config, mock_conversation_store, stored_conversation()) + + await _send_response_request( + test_request, test_auth, stream, previous_response_id=PREVIOUS_RESPONSE_ID + ) + + sent = mock_ogx_client.responses.create.await_args.kwargs + assert sent["previous_response_id"] == PREVIOUS_RESPONSE_ID + assert "conversation" not in sent + await _assert_stored( + mock_conversation_store, + stored_conversation() + _turn(DEFAULT_MODEL_RESPONSE), + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("terminal_event", "stored_output"), + [ + ("response.incomplete", [msg("assistant", PARTIAL_ANSWER)]), + ("response.failed", []), + ], + ids=["incomplete", "failed"], + ) + async def test_stream_that_did_not_complete_is_stored_once( + self, + terminal_event: str, + stored_output: list[OpenAIResponseMessage], + test_config: AppConfig, + ogx_stream: SimpleNamespace, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """A stream that ends incomplete or failed is stored with the output it has. + + That is the part of the answer an incomplete response produced, and + nothing for a failed one: the conversation gains the input alone. + """ + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + ogx_stream.terminal_event = terminal_event + + chunks = await _send_response_request(test_request, test_auth, stream=True) + + assert any(f"event: {terminal_event}" in chunk for chunk in chunks) + await _assert_stored( + mock_conversation_store, + _compacted_conversation() + [msg("user", NEW_QUERY)] + stored_output, + ) + + @pytest.mark.asyncio + async def test_stream_without_a_terminal_event_stores_nothing( + self, + test_config: AppConfig, + ogx_stream: SimpleNamespace, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + ) -> None: + """A stream that ends without a final response has no output to store.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + ogx_stream.terminal_event = None + + chunks = await _send_response_request(test_request, test_auth, stream=True) + + assert chunks[-1] == "data: [DONE]\n\n" + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + @pytest.mark.parametrize("turn_of_ours", ["blocked", "continuation"]) + async def test_turn_is_not_stored_when_the_request_says_so( + self, + turn_of_ours: str, + stream: bool, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """A request with ``store`` off leaves the conversation as it was. + + Both turns would be ours to store: a request a shield blocked, and a + continuation from a previous response. (A request with ``store`` off + is never compacted, so there is no third case.) + """ + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + previous_response_id = None + if turn_of_ours == "blocked": + mocker.patch( + "app.endpoints.responses.run_shield_moderation_v2", + return_value=_blocked(), + ) + else: + _link_previous_response(patch_db_session) + previous_response_id = PREVIOUS_RESPONSE_ID + + await _send_response_request( + test_request, + test_auth, + stream, + store=False, + previous_response_id=previous_response_id, + ) + + await _assert_stored(mock_conversation_store, _plain_conversation()) + + +# ========================================== +# /a2a +# ========================================== + + +def _put_a2a_on_the_conversation(mocker: MockerFixture) -> Any: + """Put the A2A endpoint on the stored conversation with a mocked agent.""" + mocker.patch( + "app.endpoints.a2a.get_lightspeed_agent_card", + return_value=FAKE_AGENT_CARD, + ) + + async def _prepare( + client: Any, query_request: Any, *args: Any, **kwargs: Any + ) -> ResponsesApiParams: + """Return params that carry the query as it arrived.""" + _ = client, args, kwargs + return ResponsesApiParams( + input=query_request.query, + model=TEST_MODEL, + conversation=CONV_ID_LLAMA, + store=True, + stream=True, + ) + + mocker.patch("app.endpoints.a2a.prepare_responses_params", side_effect=_prepare) + agent = mock_a2a_agent(mocker) + mocker.patch("app.endpoints.a2a.build_agent", return_value=agent) + return agent + + +class TestA2ATurnPersistence: + """What the A2A endpoint leaves in the conversation.""" + + @pytest.mark.asyncio + async def test_completed_turn_in_compacted_mode_is_stored_once( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_auth: AuthTuple, + mocker: MockerFixture, + ) -> None: + """The conversation gains the message as it arrived and the answer, once.""" + _ = mock_ogx_client + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _put_a2a_on_the_conversation(mocker) + + await handle_a2a_jsonrpc_post( + request=build_a2a_request(NEW_QUERY), auth=test_auth, mcp_headers={} + ) + + await _assert_stored( + mock_conversation_store, + _compacted_conversation() + _turn(DEFAULT_MODEL_RESPONSE), + ) + + @pytest.mark.asyncio + async def test_completed_turn_outside_compacted_mode_is_left_to_ogx( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_auth: AuthTuple, + mocker: MockerFixture, + ) -> None: + """With the conversation parameter sent, OGX stores the turn, not we.""" + _ = mock_ogx_client + await _seed(test_config, mock_conversation_store, _plain_conversation()) + _put_a2a_on_the_conversation(mocker) + + await handle_a2a_jsonrpc_post( + request=build_a2a_request(NEW_QUERY), auth=test_auth, mcp_headers={} + ) + + await _assert_stored(mock_conversation_store, _plain_conversation()) + + @pytest.mark.asyncio + async def test_failed_turn_in_compacted_mode_is_not_stored( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_auth: AuthTuple, + mocker: MockerFixture, + ) -> None: + """A turn that fails in the model call stores nothing.""" + _ = mock_ogx_client + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + agent = _put_a2a_on_the_conversation(mocker) + agent.run_stream_events.side_effect = RuntimeError("the model call failed") + + await handle_a2a_jsonrpc_post( + request=build_a2a_request(NEW_QUERY), auth=test_auth, mcp_headers={} + ) + + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + +# ========================================== +# A write that fails +# ========================================== + + +def _break_the_store(mock_ogx_client: AsyncMockType, mocker: MockerFixture) -> None: + """Make every write to the conversation fail; reads keep working.""" + mock_ogx_client.items.create = mocker.AsyncMock( + side_effect=ApiException(status=500, reason="the store is down") + ) + + +class TestFailedWrite: + """What a failed write does to the request, which differs by endpoint. + + /v1/query and /v1/responses answer with an error: the turn is part of + what they deliver. The streaming paths and A2A have delivered the answer + by the time they store the turn, so they log the failure and go on. + In every case the write is tried once. + """ + + @pytest.mark.asyncio + async def test_query_fails( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """/v1/query answers with an error when the turn cannot be stored.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_query_agent) + _break_the_store(mock_ogx_client, mocker) + + with pytest.raises(HTTPException) as error: + await _send_query(test_request, test_auth) + + assert error.value.status_code == 500 + assert mock_ogx_client.items.create.await_count == 1 + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + @pytest.mark.asyncio + async def test_streaming_query_delivers_the_answer( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """/v1/streaming_query ends the stream normally and logs the failure.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_streaming_query_agent) + _break_the_store(mock_ogx_client, mocker) + + chunks = await _drain(await _send_streaming_query(test_request, test_auth)) + + assert any('"event": "end"' in chunk for chunk in chunks) + assert not any('"event": "error"' in chunk for chunk in chunks) + assert mock_ogx_client.items.create.await_count == 1 + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + @pytest.mark.asyncio + async def test_interrupted_streaming_query_still_records_the_turn( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """An interrupt logs the failed write and records the turn in the database.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_streaming_query_agent) + mock_streaming_query_agent.run_stream_events.return_value = _stream_that_stalls( + asyncio.Event() + ) + _break_the_store(mock_ogx_client, mocker) + + chunks = await _interrupt_stream(test_request, test_auth) + + assert any('"event": "interrupted"' in chunk for chunk in chunks) + assert mock_ogx_client.items.create.await_count == 1 + await _assert_stored(mock_conversation_store, _compacted_conversation()) + recorded_turns = ( + patch_db_session.query(UserTurn) + .filter_by(conversation_id=EXISTING_CONV_ID) + .count() + ) + assert recorded_turns == 1 + + @pytest.mark.asyncio + async def test_a2a_delivers_the_answer( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_auth: AuthTuple, + mocker: MockerFixture, + ) -> None: + """A2A completes the task and logs the failure.""" + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _put_a2a_on_the_conversation(mocker) + _break_the_store(mock_ogx_client, mocker) + + response = await handle_a2a_jsonrpc_post( + request=build_a2a_request(NEW_QUERY), auth=test_auth, mcp_headers={} + ) + + assert response.status_code == 200 + assert mock_ogx_client.items.create.await_count == 1 + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + @pytest.mark.asyncio + @pytest.mark.usefixtures("ogx_stream") + async def test_responses_request_fails( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """The request answers with an error when the turn cannot be stored.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _break_the_store(mock_ogx_client, mocker) + + with pytest.raises(HTTPException) as error: + await _send_response_request(test_request, test_auth, stream=False) + + assert error.value.status_code == 500 + assert mock_ogx_client.items.create.await_count == 1 + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + @pytest.mark.asyncio + @pytest.mark.usefixtures("ogx_stream") + async def test_responses_stream_ends_before_done( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """The stream has sent its terminal event and breaks off before ``[DONE]``.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _break_the_store(mock_ogx_client, mocker) + + response = await responses_endpoint_handler( + request=test_request, + responses_request=ResponsesRequest( + input=NEW_QUERY, + model=TEST_MODEL, + conversation=EXISTING_CONV_ID, + stream=True, + store=True, + generate_topic_summary=False, + ), + auth=test_auth, + mcp_headers={}, + ) + assert isinstance(response, StreamingResponse) + chunks: list[str] = [] + with pytest.raises(HTTPException) as error: + async for chunk in response.body_iterator: + chunks.append(str(chunk)) + + assert error.value.status_code == 500 + assert any("event: response.completed" in chunk for chunk in chunks) + assert "data: [DONE]\n\n" not in chunks + assert mock_ogx_client.items.create.await_count == 1 + await _assert_stored(mock_conversation_store, _compacted_conversation()) + + +# ========================================== +# A turn nobody stored +# ========================================== + + +class TestTurnNobodyStored: + """What happens when an endpoint is changed and stops storing the turn. + + This is LCORE-3883 replayed: the step that stores the turn is taken out of + each endpoint, the way a cleanup would. The request then does not end as + if nothing had happened. + """ + + @pytest.mark.asyncio + async def test_query_fails( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """The scope around the model call reports the turn nobody stored.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + mocker.patch( + "app.endpoints.query.retrieve_agent_response", + new=mocker.AsyncMock(return_value=TurnSummary(llm_response="An answer")), + ) + + with pytest.raises(TurnNotStoredError, match=CONV_ID_LLAMA): + await _send_query(test_request, test_auth) + + @pytest.mark.asyncio + async def test_streaming_query_fails( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_streaming_query_agent: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """The stream breaks off before its end event.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _agent_answer(mock_streaming_query_agent) + mocker.patch( + "utils.agents.streaming._persist_compacted_turn", new=mocker.AsyncMock() + ) + + response = await _send_streaming_query(test_request, test_auth) + assert isinstance(response, StreamingResponse) + chunks: list[str] = [] + with pytest.raises(TurnNotStoredError, match=CONV_ID_LLAMA): + async for chunk in response.body_iterator: + chunks.append(str(chunk)) + + assert not any('"event": "end"' in chunk for chunk in chunks) + + @pytest.mark.asyncio + async def test_a2a_fails_the_task( + self, + test_config: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_auth: AuthTuple, + mocker: MockerFixture, + ) -> None: + """The task ends as failed, with the reason in its status message.""" + _ = mock_ogx_client + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + _put_a2a_on_the_conversation(mocker) + mocker.patch( + "app.endpoints.a2a._persist_compacted_a2a_turn", new=mocker.AsyncMock() + ) + + response = await handle_a2a_jsonrpc_post( + request=build_a2a_request(NEW_QUERY), auth=test_auth, mcp_headers={} + ) + + status = json.loads(bytes(response.body))["result"]["status"] + assert status["state"] == "failed" + assert CONV_ID_LLAMA in status["message"]["parts"][0]["text"] + assert "nobody tried to store it" in status["message"]["parts"][0]["text"] + + @pytest.mark.asyncio + @pytest.mark.usefixtures("ogx_stream") + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + async def test_responses_request_fails( + self, + stream: bool, + test_config: AppConfig, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, + ) -> None: + """The handler reports the turn nobody stored, in both modes.""" + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _compacted_conversation()) + mocker.patch( + "app.endpoints.responses._append_previous_response_turn", + new=mocker.AsyncMock(), + ) + + with pytest.raises(TurnNotStoredError, match=CONV_ID_LLAMA): + await _send_response_request(test_request, test_auth, stream) diff --git a/tests/unit/app/endpoints/test_a2a.py b/tests/unit/app/endpoints/test_a2a.py index 683dd904e..e87afa25f 100644 --- a/tests/unit/app/endpoints/test_a2a.py +++ b/tests/unit/app/endpoints/test_a2a.py @@ -790,6 +790,10 @@ async def test_process_task_streaming_handles_api_connection_error( # pylint: d mock_responses_params = mocker.Mock() mock_responses_params.model = "test-model" mock_responses_params.conversation = "conv_x" + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.a2a.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -873,6 +877,10 @@ async def test_process_task_streaming_handles_agent_run_error( # pylint: disabl mock_responses_params = mocker.Mock() mock_responses_params.model = "test-model" mock_responses_params.conversation = "conv_x" + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.a2a.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -950,6 +958,10 @@ async def test_process_task_streaming_applies_compaction( # pylint: disable=too mock_params = mocker.Mock() mock_params.model = "test-model" mock_params.conversation = "conv_x" + mock_params.input = "What is OpenShift?" + mock_params.store = True + mock_params.previous_response_id = None + mock_params.omit_conversation = False mock_params.skills = None mocker.patch( "app.endpoints.a2a.prepare_responses_params", @@ -1277,6 +1289,10 @@ async def test_execute_span_success_attributes( # pylint: disable=too-many-loca mock_responses_params = mocker.Mock() mock_responses_params.model = "watsonx/granite-3.1" mock_responses_params.conversation = "conv_x" + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.a2a.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -1377,6 +1393,10 @@ async def test_execute_span_tool_calls( # pylint: disable=too-many-locals,too-m mock_responses_params = mocker.Mock() mock_responses_params.model = "openai/gpt-4" mock_responses_params.conversation = "conv_x" + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.a2a.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -1384,6 +1404,8 @@ async def test_execute_span_tool_calls( # pylint: disable=too-many-locals,too-m compaction_result = mocker.Mock() compaction_result.params = mock_responses_params + compaction_result.compacted = False + compaction_result.original_input = None mocker.patch( "app.endpoints.a2a.apply_compaction_blocking", new=mocker.AsyncMock(return_value=compaction_result), @@ -1489,6 +1511,10 @@ async def test_execute_span_inference_completed_event( # pylint: disable=too-ma mock_responses_params = mocker.Mock() mock_responses_params.model = "test-model" mock_responses_params.conversation = "conv_x" + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.a2a.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -1496,6 +1522,8 @@ async def test_execute_span_inference_completed_event( # pylint: disable=too-ma compaction_result = mocker.Mock() compaction_result.params = mock_responses_params + compaction_result.compacted = False + compaction_result.original_input = None mocker.patch( "app.endpoints.a2a.apply_compaction_blocking", new=mocker.AsyncMock(return_value=compaction_result), @@ -1707,6 +1735,10 @@ async def test_execute_span_no_tool_calls( # pylint: disable=too-many-locals,to mock_responses_params = mocker.Mock() mock_responses_params.model = "test-model" mock_responses_params.conversation = "conv_x" + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.a2a.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -1714,6 +1746,8 @@ async def test_execute_span_no_tool_calls( # pylint: disable=too-many-locals,to compaction_result = mocker.Mock() compaction_result.params = mock_responses_params + compaction_result.compacted = False + compaction_result.original_input = None mocker.patch( "app.endpoints.a2a.apply_compaction_blocking", new=mocker.AsyncMock(return_value=compaction_result), diff --git a/tests/unit/app/endpoints/test_query.py b/tests/unit/app/endpoints/test_query.py index c24480d49..277f1f8f1 100644 --- a/tests/unit/app/endpoints/test_query.py +++ b/tests/unit/app/endpoints/test_query.py @@ -130,6 +130,10 @@ async def test_successful_query_no_conversation( mock_responses_params.model = "provider1/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mock_responses_params.model_dump.return_value = { "input": "test", "model": "provider1/model1", @@ -212,6 +216,10 @@ async def test_query_reports_context_status( mock_responses_params.model = "provider1/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mocker.patch( "app.endpoints.query.prepare_responses_params", new=mocker.AsyncMock(return_value=mock_responses_params), @@ -297,6 +305,10 @@ async def test_query_merges_inline_and_tool_rag_chunks_and_documents( mock_responses_params.model = "provider1/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mock_responses_params.model_dump.return_value = { "input": "test", "model": "provider1/model1", @@ -372,6 +384,10 @@ async def test_successful_query_with_conversation( mock_responses_params.model = "provider1/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mock_responses_params.model_dump.return_value = { "input": "test", "model": "provider1/model1", @@ -445,6 +461,10 @@ async def test_query_with_attachments( mock_responses_params.model = "provider1/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mock_responses_params.model_dump.return_value = { "input": "test", "model": "provider1/model1", @@ -508,6 +528,10 @@ async def test_query_with_topic_summary( mock_responses_params.model = "provider1/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mock_responses_params.model_dump.return_value = { "input": "test", "model": "provider1/model1", @@ -578,6 +602,10 @@ async def test_query_azure_token_refresh( mock_responses_params.model = "azure/model1" mock_responses_params.conversation = "conv_123" mock_responses_params.tools = None + mock_responses_params.input = "What is OpenShift?" + mock_responses_params.store = True + mock_responses_params.previous_response_id = None + mock_responses_params.omit_conversation = False mock_responses_params.model_dump.return_value = { "input": "test", "model": "azure/model1", diff --git a/tests/unit/app/endpoints/test_query_otel.py b/tests/unit/app/endpoints/test_query_otel.py index 032f9fac1..bcde6cfa8 100644 --- a/tests/unit/app/endpoints/test_query_otel.py +++ b/tests/unit/app/endpoints/test_query_otel.py @@ -85,6 +85,10 @@ def _patch_query_success(mocker: MockerFixture) -> None: mock_params.model = "provider1/model1" mock_params.conversation = "conv_123" mock_params.tools = None + mock_params.input = "What is OpenShift?" + mock_params.store = True + mock_params.previous_response_id = None + mock_params.omit_conversation = False mock_params.model_dump.return_value = {"input": "test", "model": "provider1/model1"} mocker.patch( f"{MODULE}.prepare_responses_params", diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index 27da0c1a1..7f86a7d90 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -58,6 +58,8 @@ VALID_CONV_ID_NORMALIZED = "e6afd7aaa97b49ce8f4f96a801b07893d9cb784d72e53e3c" MODULE = "app.endpoints.responses" ENDPOINTS_MODULE = "utils.endpoints" +# The one function that appends a turn to a conversation, where its owner calls it. +TURN_WRITER = "utils.pending_turn.append_turn_items_to_conversation" UTILS_RESPONSES_MODULE = "utils.responses" MODEL = "google-vertex/publishers/google/models/gemini-2.5-flash" SERVER_INSTRUCTIONS = "Server instructions" @@ -646,7 +648,7 @@ async def test_responses_blocked_with_conversation_appends_refusal( mock_moderation.message = "Blocked" mock_moderation.moderation_id = "resp_blocked_123" mock_append = mocker.patch( - f"{MODULE}.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) mocker.patch(f"{MODULE}.store_query_results") @@ -659,10 +661,10 @@ async def test_responses_blocked_with_conversation_appends_refusal( ) mock_append.assert_awaited_once_with( - client=mock_client, - conversation_id=VALID_CONV_ID, - user_input=responses_request.input, - llm_output=[mock_moderation.refusal_response], + mock_client, + VALID_CONV_ID, + responses_request.input, + [mock_moderation.refusal_response], ) assert isinstance(response, ResponsesResponse) payload = response.model_dump() @@ -786,7 +788,7 @@ async def test_handle_non_streaming_blocked_returns_refusal( _patch_handle_non_streaming_common(mocker, minimal_config) mocker.patch( - f"{MODULE}.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) mock_client.items.create = mocker.AsyncMock() @@ -972,7 +974,7 @@ async def test_handle_non_streaming_with_previous_response_id_appends_turn( return_value=VALID_CONV_ID_NORMALIZED, ) mock_append = mocker.patch( - f"{MODULE}.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) @@ -1547,7 +1549,7 @@ async def mock_stream() -> Any: return_value=VALID_CONV_ID_NORMALIZED, ) mock_append = mocker.patch( - f"{MODULE}.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) mock_holder = mocker.Mock() @@ -2830,7 +2832,7 @@ async def test_append_previous_response_turn_compacted(mocker: MockerFixture) -> rewritten explicit input on api_params. """ append = mocker.patch( - "app.endpoints.responses.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) api_params = mocker.Mock( @@ -2838,6 +2840,7 @@ async def test_append_previous_response_turn_compacted(mocker: MockerFixture) -> conversation="conv_x", previous_response_id=None, input=["rewritten explicit input"], + omit_conversation=True, ) context = mocker.Mock( client=mocker.AsyncMock(), @@ -2857,11 +2860,14 @@ async def test_append_previous_response_turn_not_stored_when_store_false( ) -> None: """No append happens when store is disabled, even in compacted mode.""" append = mocker.patch( - "app.endpoints.responses.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) api_params = mocker.Mock( - store=False, conversation="conv_x", previous_response_id=None + store=False, + conversation="conv_x", + previous_response_id=None, + omit_conversation=True, ) context = mocker.Mock(client=mocker.AsyncMock(), compacted_original_input="q") @@ -2879,12 +2885,15 @@ async def test_persist_blocked_response_turn_compacted(mocker: MockerFixture) -> against the original user input carried on the context (LCORE-1572). """ append = mocker.patch( - "app.endpoints.responses.append_turn_items_to_conversation", + TURN_WRITER, new=mocker.AsyncMock(), ) refusal = mocker.Mock() api_params = mocker.Mock( - store=True, conversation="conv_x", input=["rewritten explicit input"] + store=True, + conversation="conv_x", + input=["rewritten explicit input"], + omit_conversation=True, ) context = mocker.Mock( client=mocker.AsyncMock(), @@ -2895,8 +2904,5 @@ async def test_persist_blocked_response_turn_compacted(mocker: MockerFixture) -> await _persist_blocked_response_turn(api_params, context) append.assert_awaited_once_with( - client=context.client, - conversation_id="conv_x", - user_input="the original query", - llm_output=[refusal], + context.client, "conv_x", "the original query", [refusal] ) diff --git a/tests/unit/utils/agents/test_query.py b/tests/unit/utils/agents/test_query.py index 038609ec5..0f314fd48 100644 --- a/tests/unit/utils/agents/test_query.py +++ b/tests/unit/utils/agents/test_query.py @@ -34,6 +34,7 @@ get_agent_finish_reason, retrieve_agent_response, ) +from utils.pending_turn import PendingTurn from utils.token_counter import TokenCounter @@ -351,6 +352,20 @@ def test_raises_http_exception_on_missing_finish_reason( assert exc_info.value.status_code == 500 +def client_capturing_writes(mocker: MockerFixture) -> tuple[Any, list[Any]]: + """Return a client whose conversation writes are recorded.""" + stored: list[Any] = [] + + async def _create( + _conversation_id: str, *, add_items_request: Any = None, **_kwargs: Any + ) -> None: + stored.extend(getattr(add_items_request, "items", add_items_request) or []) + + client = mocker.AsyncMock() + client.items.create = _create + return client, stored + + class TestRetrieveAgentResponse: """Tests for retrieve_agent_response.""" @@ -406,16 +421,47 @@ async def test_compacted_input_runs_agent_with_prompt_text( mock_agent = mocker.AsyncMock() mock_agent.run = mocker.AsyncMock(return_value=run_result) mocker.patch("utils.agents.query.build_agent", return_value=mock_agent) + client = mocker.AsyncMock() summary = await retrieve_agent_response( - client=mocker.AsyncMock(), + client=client, responses_params=params, endpoint_path=ENDPOINT_PATH_QUERY, + turn=PendingTurn.for_request(client, params, "new question"), ) mock_agent.run.assert_awaited_once_with("new question") assert summary.llm_response == "Answer" + @pytest.mark.asyncio + async def test_compacted_request_without_its_turn_is_refused( + self, + mocker: MockerFixture, + make_responses_params: Callable[..., ResponsesApiParams], + ) -> None: + """Test a compacted request cannot run without the owner of its turn. + + This is how a caller that lost the hand-over looks (LCORE-3883): the + parameters are compacted, the input as it arrived is gone. The turn + could not be stored, so the request is refused before the model call. + """ + params = make_responses_params().model_copy( + update={ + "input": [OpenAIResponseMessage(role="user", content="q")], + "omit_conversation": True, + } + ) + mock_build_agent = mocker.patch("utils.agents.query.build_agent") + + with pytest.raises(ValueError, match="original input"): + await retrieve_agent_response( + client=mocker.AsyncMock(), + responses_params=params, + endpoint_path=ENDPOINT_PATH_QUERY, + ) + + mock_build_agent.assert_not_called() + @pytest.mark.asyncio async def test_inference_error_is_logged( self, @@ -594,21 +640,6 @@ class TestQueryCompactedTurnPersistence: arguments of a mocked helper, so a break in the write path is caught. """ - @staticmethod - def _client_capturing_writes(mocker: MockerFixture) -> tuple[Any, list[Any]]: - """Return a client whose conversation writes are recorded.""" - stored: list[Any] = [] - - async def _create( - conversation_id: str, *, add_items_request: Any = None, **_: Any - ) -> None: - _ = conversation_id - stored.extend(getattr(add_items_request, "items", add_items_request) or []) - - client = mocker.AsyncMock() - client.items.create = _create - return client, stored - @pytest.mark.asyncio async def test_compacted_success_appends_turn_with_captured_output( self, @@ -634,13 +665,13 @@ async def test_compacted_success_appends_turn_with_captured_output( OpenAIResponseMessage(role="assistant", content="Answer") ] mocker.patch("utils.agents.query.build_agent", return_value=mock_agent) - client, stored = self._client_capturing_writes(mocker) + client, stored = client_capturing_writes(mocker) await retrieve_agent_response( client=client, responses_params=params, endpoint_path=ENDPOINT_PATH_QUERY, - original_input="new question", + turn=PendingTurn.for_request(client, params, "new question"), ) texts = [str(item) for item in stored] @@ -665,7 +696,7 @@ async def test_non_compacted_success_does_not_append_turn( OpenAIResponseMessage(role="assistant", content="Answer") ] mocker.patch("utils.agents.query.build_agent", return_value=mock_agent) - client, stored = self._client_capturing_writes(mocker) + client, stored = client_capturing_writes(mocker) await retrieve_agent_response( client=client, diff --git a/tests/unit/utils/agents/test_streaming.py b/tests/unit/utils/agents/test_streaming.py index f614301df..2dfd7301c 100644 --- a/tests/unit/utils/agents/test_streaming.py +++ b/tests/unit/utils/agents/test_streaming.py @@ -67,6 +67,7 @@ serialize_event, ) from utils.otel_tracing import SpanAttributes, SpanEvents +from utils.pending_turn import PendingTurn from utils.token_counter import TokenCounter INTERRUPTED_INDICATOR = f"\n\n*{INTERRUPTED_RESPONSE_MESSAGE}*" @@ -489,6 +490,19 @@ def test_part_end_native_tool_call_returns_none_when_skipped( assert not turn_state.turn_summary.tool_calls +def capture_conversation_writes(context: Any) -> list[Any]: + """Wire a stateful fake onto the client and return the captured items.""" + stored: list[Any] = [] + + async def _create( + _conversation_id: str, *, add_items_request: Any = None, **_kwargs: Any + ) -> None: + stored.extend(getattr(add_items_request, "items", add_items_request) or []) + + context.client.items.create = _create + return stored + + @pytest.mark.usefixtures("patch_streaming_configuration") class TestRetrieveAgentResponseGenerator: """Tests for retrieve_agent_response_generator.""" @@ -1577,20 +1591,6 @@ class TestCompactedTurnPersistence: the arguments of a mocked helper. """ - @staticmethod - def _capture_conversation_writes(context: Any) -> list[Any]: - """Wire a stateful fake onto the client and return the captured items.""" - stored: list[Any] = [] - - async def _create( - conversation_id: str, *, add_items_request: Any = None, **_: Any - ) -> None: - _ = conversation_id - stored.extend(getattr(add_items_request, "items", add_items_request) or []) - - context.client.items.create = _create - return stored - @staticmethod def _patch_finalizers(mocker: MockerFixture) -> None: """Stub the post-stream finalization the persistence test does not exercise.""" @@ -1612,11 +1612,12 @@ async def test_compacted_success_appends_turn_to_conversation( self, mocker: MockerFixture, make_generator_context: Callable[..., ResponseGeneratorContext], - responses_params: ResponsesApiParams, + make_responses_params: Callable[..., ResponsesApiParams], ) -> None: """A completed compacted stream writes the user turn and the LLM output.""" context = make_generator_context() - stored = self._capture_conversation_writes(context) + compacted_params = make_responses_params(omit_conversation=True) + stored = capture_conversation_writes(context) self._patch_finalizers(mocker) turn_summary = TurnSummary() @@ -1636,10 +1637,12 @@ async def inner() -> AsyncIterator[str]: async for event in generate_agent_response( inner(), context, - responses_params, + compacted_params, turn_summary, [], - original_input="the original question", + turn=PendingTurn.for_request( + context.client, compacted_params, "the original question" + ), root_span=_dummy_root_span(), ) ] @@ -1659,7 +1662,7 @@ async def test_non_compacted_success_does_not_append_turn( ) -> None: """Without compaction OGX stores the turn, so we must not duplicate it.""" context = make_generator_context() - stored = self._capture_conversation_writes(context) + stored = capture_conversation_writes(context) self._patch_finalizers(mocker) turn_summary = TurnSummary() @@ -1689,13 +1692,22 @@ async def test_interrupt_guard_prevents_double_persistence( self, mocker: MockerFixture, make_generator_context: Callable[..., ResponseGeneratorContext], - responses_params: ResponsesApiParams, + make_responses_params: Callable[..., ResponsesApiParams], ) -> None: - """An interrupt that already persisted the turn blocks the success path.""" + """An interrupt that already persisted the turn blocks the success path. + + Both protections agree here: the guard shared with the interrupt path + is set, and the turn is settled. The next test opens the guard. + """ context = make_generator_context() - stored = self._capture_conversation_writes(context) + compacted_params = make_responses_params(omit_conversation=True) + stored = capture_conversation_writes(context) self._patch_finalizers(mocker) - # Simulate the interrupt path having already persisted this turn. + # The interrupt path has already persisted this turn. + turn = PendingTurn.for_request( + context.client, compacted_params, "the original question" + ) + await turn.store_interrupted("The ans") mocker.patch( "utils.agents.streaming.register_interrupt_callback", return_value=[True] ) @@ -1716,25 +1728,84 @@ async def inner() -> AsyncIterator[str]: async for event in generate_agent_response( inner(), context, - responses_params, + compacted_params, turn_summary, [], - original_input="the original question", + turn=turn, root_span=_dummy_root_span(), ) ] - assert stored == [] + texts = [str(item) for item in stored] + assert len(stored) == 2, f"expected the interrupted turn only, got {texts}" + assert "the original question" in texts[0] + assert "The ans" in texts[1] + assert "The answer." not in texts[1] + + @pytest.mark.asyncio + async def test_settled_turn_is_not_stored_again_at_the_end_of_the_stream( + self, + mocker: MockerFixture, + make_generator_context: Callable[..., ResponseGeneratorContext], + make_responses_params: Callable[..., ResponsesApiParams], + ) -> None: + """The owner alone keeps the end of the stream from storing a second turn. + + The guard shared with the interrupt path is open here, so the end of + the stream does try to store the turn. The turn has been stored by + then, and its owner stores it once. + """ + context = make_generator_context() + compacted_params = make_responses_params(omit_conversation=True) + stored = capture_conversation_writes(context) + self._patch_finalizers(mocker) + turn = PendingTurn.for_request( + context.client, compacted_params, "the original question" + ) + await turn.store_interrupted("The ans") + mocker.patch( + "utils.agents.streaming.register_interrupt_callback", return_value=[False] + ) + + turn_summary = TurnSummary() + turn_summary.token_usage = TokenCounter(input_tokens=1, output_tokens=1) + turn_summary.output_items = [ + OpenAIResponseMessage(role="assistant", content="The answer.") + ] + + async def inner() -> AsyncIterator[str]: + yield serialize_event( + TokenStreamPayload.create(chunk_id=0, token="x"), MEDIA_TYPE_JSON + ) + + _ = [ + event + async for event in generate_agent_response( + inner(), + context, + compacted_params, + turn_summary, + [], + turn=turn, + root_span=_dummy_root_span(), + ) + ] + + texts = [str(item) for item in stored] + assert len(stored) == 2, f"expected the interrupted turn only, got {texts}" + assert "The ans" in texts[1] + assert "The answer." not in texts[1] @pytest.mark.asyncio async def test_persistence_failure_does_not_fail_the_stream( self, mocker: MockerFixture, make_generator_context: Callable[..., ResponseGeneratorContext], - responses_params: ResponsesApiParams, + make_responses_params: Callable[..., ResponsesApiParams], ) -> None: """The client keeps its answer even when the conversation write fails.""" context = make_generator_context() + compacted_params = make_responses_params(omit_conversation=True) self._patch_finalizers(mocker) async def _boom(_conversation_id: str, **_: Any) -> None: @@ -1758,10 +1829,10 @@ async def inner() -> AsyncIterator[str]: async for event in generate_agent_response( inner(), context, - responses_params, + compacted_params, turn_summary, [], - original_input="q", + turn=PendingTurn.for_request(context.client, compacted_params, "q"), root_span=_dummy_root_span(), ) ] diff --git a/tests/unit/utils/test_conversation_compaction.py b/tests/unit/utils/test_conversation_compaction.py index b5bea40f7..6a0350e7a 100644 --- a/tests/unit/utils/test_conversation_compaction.py +++ b/tests/unit/utils/test_conversation_compaction.py @@ -404,17 +404,6 @@ async def test_streaming_emits_event_before_summarizing(mocker: MockerFixture) - assert yielded[-1].compacted is True -@pytest.mark.asyncio -async def test_store_compacted_turn_appends(mocker: MockerFixture) -> None: - """store_compacted_turn delegates to append_turn_items_to_conversation.""" - append = mocker.patch.object( - cc, "append_turn_items_to_conversation", mocker.AsyncMock() - ) - client = mocker.AsyncMock() - await cc.store_compacted_turn(client, CONV, "the query", ["out"]) - append.assert_awaited_once_with(client, CONV, "the query", ["out"]) - - # --- needs_compaction_path (the tight gate protecting non-compacting requests) --- diff --git a/tests/unit/utils/test_pending_turn.py b/tests/unit/utils/test_pending_turn.py new file mode 100644 index 000000000..c92e3b769 --- /dev/null +++ b/tests/unit/utils/test_pending_turn.py @@ -0,0 +1,408 @@ +"""Unit tests for the owner of the turns lightspeed-stack stores itself (LCORE-3908). + +Every test asserts on the items that reach the conversation, captured by a +client that records its writes, not on a helper having been called. +""" + +from pathlib import Path +from typing import Any, Optional + +import pytest +from fastapi import HTTPException +from ogx_api.openai_responses import OpenAIResponseMessage +from ogx_client import ApiException +from pytest_mock import MockerFixture + +from models.common.responses.responses_api_params import ResponsesApiParams +from models.common.responses.types import ResponseInput +from models.config import CompactionConfiguration, InferenceConfiguration +from utils.conversation_compaction import MARKER_SENTINEL, apply_compaction_blocking +from utils.pending_turn import PendingTurn, TurnNotStoredError, pending_turn +from utils.token_estimator import extract_message_text + +CONVERSATION = "conv_abc123" +MODEL = "openai/gpt-4o-mini" +QUERY = "new question" +ANSWER = OpenAIResponseMessage(role="assistant", content="the answer") +REFUSAL = OpenAIResponseMessage(role="assistant", content="blocked by a shield") + + +class RecordingClient: # pylint: disable=too-few-public-methods + """Stand-in for the OGX client that keeps what is written to a conversation.""" + + def __init__(self, fail_with: Optional[Exception] = None) -> None: + """Create the client, optionally one whose writes fail.""" + self.stored: list[tuple[str, str, str]] = [] + self.writes = 0 + self._fail_with = fail_with + self.items = self + + async def create( + self, conversation_id: str, *, add_items_request: Any = None, **_: Any + ) -> None: + """Record the items of one write as (conversation, role, text).""" + self.writes += 1 + if self._fail_with is not None: + raise self._fail_with + self.stored.extend( + (conversation_id, str(item.role), extract_message_text(item)) + for item in add_items_request.items + ) + + +def _params( + compacted: bool = False, + store: bool = True, + previous_response_id: Optional[str] = None, +) -> ResponsesApiParams: + """Build request params the way the endpoints hand them over.""" + explicit: ResponseInput = [ + OpenAIResponseMessage(role="user", content="Summary of earlier turns"), + OpenAIResponseMessage(role="user", content=QUERY), + ] + return ResponsesApiParams( + input=explicit if compacted else QUERY, + model=MODEL, + conversation=CONVERSATION, + previous_response_id=previous_response_id, + store=store, + stream=False, + omit_conversation=compacted, + ) + + +def _turn(client: RecordingClient, **kwargs: Any) -> PendingTurn: + """Build the pending turn of a request; a compacted one gets its original input.""" + params = _params(**kwargs) + original = QUERY if params.omit_conversation else None + return PendingTurn.for_request(client, params, original) # type: ignore[arg-type] + + +USER = (CONVERSATION, "user", QUERY) + + +# --- who stores a completed turn --- + + +@pytest.mark.asyncio +async def test_completed_turn_is_left_to_ogx_when_it_got_the_conversation() -> None: + """With the conversation parameter sent, a completed turn is not ours.""" + client = RecordingClient() + turn = _turn(client) + + assert await turn.store_completed([ANSWER]) is False + + assert not client.stored + assert turn.settled + + +@pytest.mark.asyncio +async def test_completed_turn_in_compacted_mode_is_stored() -> None: + """In compacted mode the turn is stored against the input as it arrived.""" + client = RecordingClient() + turn = _turn(client, compacted=True) + + assert await turn.store_completed([ANSWER]) is True + + assert client.stored == [USER, (CONVERSATION, "assistant", "the answer")] + assert client.writes == 1 + + +@pytest.mark.asyncio +async def test_completed_turn_continuing_a_previous_response_is_stored() -> None: + """OGX does not store a turn that continues from a previous response.""" + client = RecordingClient() + turn = _turn(client, previous_response_id="resp_1") + + assert await turn.store_completed([ANSWER]) is True + + assert client.stored == [USER, (CONVERSATION, "assistant", "the answer")] + + +@pytest.mark.asyncio +async def test_input_given_as_items_is_stored_as_given() -> None: + """An input that arrived as a list of items is stored item by item.""" + client = RecordingClient() + original: ResponseInput = [ + OpenAIResponseMessage(role="user", content="first part"), + OpenAIResponseMessage(role="user", content="second part"), + ] + turn = PendingTurn.for_request( + client, _params(compacted=True), original # type: ignore[arg-type] + ) + + await turn.store_completed([ANSWER]) + + assert client.stored == [ + (CONVERSATION, "user", "first part"), + (CONVERSATION, "user", "second part"), + (CONVERSATION, "assistant", "the answer"), + ] + + +def test_compacted_request_needs_its_original_input() -> None: + """Without the original input a compacted turn cannot be stored correctly.""" + with pytest.raises(ValueError, match="original input"): + PendingTurn.for_request( + RecordingClient(), _params(compacted=True) # type: ignore[arg-type] + ) + + +# --- blocked and interrupted turns --- + + +@pytest.mark.asyncio +@pytest.mark.parametrize("compacted", [False, True]) +async def test_blocked_turn_is_stored(compacted: bool) -> None: + """A blocked request never reaches OGX, so the refusal turn is always ours.""" + client = RecordingClient() + turn = _turn(client, compacted=compacted) + + assert await turn.store_blocked(REFUSAL) is True + + assert client.stored == [USER, (CONVERSATION, "assistant", "blocked by a shield")] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("compacted", [False, True]) +async def test_interrupted_turn_is_stored(compacted: bool) -> None: + """An interrupted stream stores the part of the answer that was received.""" + client = RecordingClient() + turn = _turn(client, compacted=compacted) + + assert await turn.store_interrupted("half an ans") is True + + assert client.stored == [USER, (CONVERSATION, "assistant", "half an ans")] + + +@pytest.mark.asyncio +async def test_nothing_is_stored_when_the_request_says_so() -> None: + """A request with ``store`` off leaves the conversation alone, whatever happens.""" + client = RecordingClient() + + assert ( + await _turn(client, compacted=True, store=False).store_completed([ANSWER]) + is False + ) + assert await _turn(client, store=False).store_blocked(REFUSAL) is False + assert await _turn(client, store=False).store_interrupted("half") is False + + assert client.writes == 0 + + +# --- exactly once --- + + +@pytest.mark.asyncio +async def test_a_turn_is_stored_once() -> None: + """A turn that is stored is not stored again, however the second caller ends it.""" + client = RecordingClient() + turn = _turn(client, compacted=True) + + assert await turn.store_completed([ANSWER]) is True + assert await turn.store_completed([ANSWER]) is False + assert await turn.store_interrupted("half an ans") is False + assert await turn.store_blocked(REFUSAL) is False + + assert client.stored == [USER, (CONVERSATION, "assistant", "the answer")] + assert client.writes == 1 + + +@pytest.mark.asyncio +async def test_a_blocked_turn_is_not_stored_again_by_an_interrupt() -> None: + """A turn settled as blocked is the turn; a later interrupt report adds none.""" + client = RecordingClient() + turn = _turn(client) + + assert await turn.store_blocked(REFUSAL) is True + assert await turn.store_interrupted("block") is False + + assert client.stored == [USER, (CONVERSATION, "assistant", "blocked by a shield")] + + +@pytest.mark.asyncio +async def test_a_failed_write_is_not_repeated() -> None: + """A write that failed is reported to the caller and never tried again.""" + client = RecordingClient(fail_with=ApiException(status=500, reason="boom")) + turn = _turn(client, compacted=True) + + with pytest.raises(HTTPException): + await turn.store_completed([ANSWER]) + + assert turn.settled + assert await turn.store_completed([ANSWER]) is False + assert client.writes == 1 + + +# --- a turn nobody settled --- + + +def test_unsettled_turn_of_ours_is_an_error() -> None: + """A compacted turn that nobody tried to store is a lost turn.""" + turn = _turn(RecordingClient(), compacted=True) + + with pytest.raises(TurnNotStoredError, match=CONVERSATION): + turn.ensure_settled() + + +@pytest.mark.asyncio +async def test_settled_turn_passes_the_check() -> None: + """Stored, left to OGX and not wanted are all fine.""" + stored = _turn(RecordingClient(), compacted=True) + await stored.store_completed([ANSWER]) + stored.ensure_settled() + + _turn(RecordingClient()).ensure_settled() + _turn(RecordingClient(), compacted=True, store=False).ensure_settled() + + +@pytest.mark.asyncio +async def test_scope_reports_a_turn_nobody_settled() -> None: + """Leaving the scope with a compacted turn unsettled raises.""" + with pytest.raises(TurnNotStoredError): + async with pending_turn( + RecordingClient(), _params(compacted=True), QUERY # type: ignore[arg-type] + ): + pass + + +@pytest.mark.asyncio +async def test_scope_lets_a_settled_turn_through() -> None: + """The scope hands out the turn and is silent once it is stored.""" + client = RecordingClient() + + async with pending_turn( + client, _params(compacted=True), QUERY # type: ignore[arg-type] + ) as turn: + await turn.store_completed([ANSWER]) + + assert client.stored == [USER, (CONVERSATION, "assistant", "the answer")] + + +@pytest.mark.asyncio +async def test_scope_does_not_mask_a_failure() -> None: + """A request that fails inside the scope fails with its own error.""" + with pytest.raises(RuntimeError, match="the model call failed"): + async with pending_turn( + RecordingClient(), _params(compacted=True), QUERY # type: ignore[arg-type] + ): + raise RuntimeError("the model call failed") + + +# --- what the compaction seam hands over --- + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("stored_items", "compacted"), + [ + ([], False), + ([OpenAIResponseMessage(role="user", content="a question")], False), + ( + [ + OpenAIResponseMessage(role="user", content="a question"), + OpenAIResponseMessage( + role="user", content=f"{MARKER_SENTINEL} an earlier summary" + ), + ], + True, + ), + ], + ids=["new conversation", "never compacted", "compacted before"], +) +async def test_the_seam_hands_over_what_the_owner_needs( + mocker: MockerFixture, stored_items: list[Any], compacted: bool +) -> None: + """Dropping the conversation parameter and handing over the input go together. + + The owner takes ``omit_conversation`` as the sign that the turn is ours + and needs the input as it arrived to store it. The seam sets both or + neither; were that to change, a turn would be lost or stored twice. + """ + mocker.patch( + "utils.conversation_compaction.get_all_conversation_items", + new=mocker.AsyncMock(return_value=stored_items), + ) + client = RecordingClient() + + result = await apply_compaction_blocking( + client, # type: ignore[arg-type] + _params(), + InferenceConfiguration(context_windows={MODEL: 100_000}), + CompactionConfiguration(enabled=True), + ) + + assert result.compacted is compacted + assert result.params.omit_conversation is compacted + assert (result.original_input is not None) is compacted + turn = PendingTurn.for_request( + client, result.params, result.original_input # type: ignore[arg-type] + ) + assert turn.ours is compacted + assert turn.user_input == QUERY + + +# --- who may write to a conversation --- + + +def _modules_mentioning(name: str) -> list[str]: + """Return the modules under ``src`` whose source mentions *name*.""" + src = Path(__file__).resolve().parents[3] / "src" + return sorted( + str(path.relative_to(src)) + for path in src.rglob("*.py") + if name in path.read_text(encoding="utf-8") + ) + + +WRITERS = { + # the function that appends a turn, and its one caller, the owner + "append_turn_items_to_conversation": [ + "utils/conversations.py", + "utils/pending_turn.py", + ], + # the shield capabilities store the turn they rejected, from inside the + # agent run, and only when the model was handed the conversation; in + # compacted mode it is not, so they write nothing there + "append_turn_to_conversation": [ + "pydantic_ai_lightspeed/capabilities/granite_guardian/_capability.py", + "pydantic_ai_lightspeed/capabilities/question_validity/_capability.py", + "utils/conversations.py", + ], + # the call underneath both, and the write of a compaction summary marker + "items.create(": [ + "utils/conversation_compaction.py", + "utils/conversations.py", + ], + "build_add_items_request": [ + "utils/conversation_compaction.py", + "utils/conversations.py", + ], + # Granite Guardian replaces the answer OGX stored when it rejects an + # output or a tool result: the last assistant message is deleted and the + # violation message appended, again only when the model was handed the + # conversation. conversations_v1.py only names the function in a comment + "replace_last_assistant_message": [ + "app/endpoints/conversations_v1.py", + "pydantic_ai_lightspeed/capabilities/granite_guardian/_capability.py", + "utils/conversations.py", + ], +} + + +@pytest.mark.parametrize("writer", sorted(WRITERS)) +def test_conversations_are_written_by_the_known_writers_only(writer: str) -> None: + """No module other than the listed ones appends to a conversation. + + LCORE-3883 happened because every endpoint made the write itself and two + cleanups removed it. An endpoint that starts writing on its own again, + through one of the functions that write, shows up here. + + The check reads the source as text. It therefore also fails when a module + only mentions one of the names, in a comment for example. + """ + assert _modules_mentioning(writer) == WRITERS[writer], ( + f"the modules that mention {writer!r} changed; a turn has to be stored " + "through utils.pending_turn.PendingTurn, and a writer that is meant to " + "be new has to be added to WRITERS in this test" + ) diff --git a/tests/unit/utils/test_stream_interrupts.py b/tests/unit/utils/test_stream_interrupts.py index f612ef54d..62707b922 100644 --- a/tests/unit/utils/test_stream_interrupts.py +++ b/tests/unit/utils/test_stream_interrupts.py @@ -1,25 +1,62 @@ """Unit tests for stream interrupt registry and persistence utilities.""" import asyncio +from typing import Any import pytest +from ogx_api.openai_responses import OpenAIResponseMessage from pytest_mock import MockerFixture from constants import INTERRUPTED_RESPONSE_MESSAGE from models.api.requests import QueryRequest from models.common.responses.contexts import ResponseGeneratorContext from models.common.responses.responses_api_params import ResponsesApiParams +from models.common.responses.types import ResponseInput from models.common.turn_summary import TurnSummary +from utils.pending_turn import PendingTurn from utils.stream_interrupts import ( StreamInterruptRegistry, build_interrupted_response, persist_interrupted_turn, register_interrupt_callback, ) +from utils.token_estimator import extract_message_text INTERRUPTED_INDICATOR = f"\n\n*{INTERRUPTED_RESPONSE_MESSAGE}*" +def _params( + conversation: str, + input_items: ResponseInput, + omit_conversation: bool = False, +) -> ResponsesApiParams: + """Build the parameters of a streaming request.""" + return ResponsesApiParams( + input=input_items, + model="provider1/model1", + conversation=conversation, + store=True, + stream=True, + omit_conversation=omit_conversation, + ) + + +def _capture_conversation_writes(context: Any) -> list[tuple[str, str]]: + """Wire a stateful fake onto the client and return what it is asked to store.""" + stored: list[tuple[str, str]] = [] + + async def _create( + _conversation_id: str, *, add_items_request: Any = None, **_kwargs: Any + ) -> None: + stored.extend( + (str(item.role), extract_message_text(item)) + for item in add_items_request.items + ) + + context.client.items.create = _create + return stored + + @pytest.mark.asyncio async def test_persist_interrupted_turn_compacted_uses_original_input( mocker: MockerFixture, @@ -40,22 +77,17 @@ async def test_persist_interrupted_turn_compacted_uses_original_input( query="hi", conversation_id=conv ) # pyright: ignore[reportCallIssue] - responses_params = mocker.Mock(spec=ResponsesApiParams) - responses_params.conversation = conv - responses_params.model = "provider1/model1" - responses_params.input = ["explicit rewrite"] + stored = _capture_conversation_writes(context) + + responses_params = _params( + conversation=conv, + input_items=[OpenAIResponseMessage(role="user", content="explicit rewrite")], + omit_conversation=True, + ) turn_summary = TurnSummary() turn_summary.llm_response = f"partial content{INTERRUPTED_INDICATOR}" background_tasks: list[asyncio.Task[None]] = [] - items = mocker.patch( - "utils.stream_interrupts.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) - strs = mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new=mocker.AsyncMock(), - ) mocker.patch("utils.stream_interrupts.store_query_results") await persist_interrupted_turn( @@ -63,14 +95,52 @@ async def test_persist_interrupted_turn_compacted_uses_original_input( responses_params, turn_summary, background_tasks, - original_input="the original query", + turn=PendingTurn.for_request( + context.client, responses_params, "the original query" + ), ) - items.assert_awaited_once() - assert items.call_args.args[2] == "the original query" - call_output = items.call_args.args[3] - assert call_output[0].content == f"partial content{INTERRUPTED_INDICATOR}" - strs.assert_not_awaited() + assert stored == [ + ("user", "the original query"), + ("assistant", f"partial content{INTERRUPTED_INDICATOR}"), + ] + + +@pytest.mark.asyncio +async def test_persist_interrupted_turn_stores_the_turn_once( + mocker: MockerFixture, +) -> None: + """Both reactions to an interrupt persist; the conversation gains one turn. + + The cancellation handler and the interrupt callback are told apart by a + guard of their own. Should both get through, the turn still is stored + once, because they share its owner (LCORE-3908). + """ + context = mocker.Mock(spec=ResponseGeneratorContext) + context.client = mocker.AsyncMock() + context.request_id = "req-1" + context.user_id = "user_1" + context.conversation_id = "conv_1" + context.started_at = "2024-01-01T00:00:00Z" + context.skip_userid_check = False + context.query_request = QueryRequest( + query="hi", conversation_id=None + ) # pyright: ignore[reportCallIssue] + stored = _capture_conversation_writes(context) + responses_params = _params(conversation="conv_1", input_items="hi") + turn = PendingTurn.for_request(context.client, responses_params) + + turn_summary = TurnSummary() + turn_summary.llm_response = f"partial{INTERRUPTED_INDICATOR}" + mocker.patch("utils.stream_interrupts.store_query_results") + + await persist_interrupted_turn(context, responses_params, turn_summary, [], turn) + await persist_interrupted_turn(context, responses_params, turn_summary, [], turn) + + assert stored == [ + ("user", "hi"), + ("assistant", f"partial{INTERRUPTED_INDICATOR}"), + ] @pytest.mark.asyncio @@ -91,19 +161,12 @@ async def test_persist_interrupted_turn_schedules_background_topic_summary( generate_topic_summary=True, ) # pyright: ignore[reportCallIssue] - responses_params = mocker.Mock(spec=ResponsesApiParams) - responses_params.conversation = "conv_new" - responses_params.model = "provider1/model1" - responses_params.input = "hello" + responses_params = _params(conversation="conv_new", input_items="hello") turn_summary = TurnSummary() turn_summary.llm_response = INTERRUPTED_INDICATOR background_tasks: list[asyncio.Task[None]] = [] - mocker.patch( - "utils.stream_interrupts.append_turn_to_conversation", - new=mocker.AsyncMock(), - ) mocker.patch("utils.stream_interrupts.store_query_results") background_mock = mocker.patch( "utils.stream_interrupts.background_update_topic_summary",