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",