From 2af20de580332eab58f6479f690ce94da8cf681b Mon Sep 17 00:00:00 2001 From: Maxim Svistunov Date: Mon, 28 Sep 2026 11:49:04 +0200 Subject: [PATCH 1/2] LCORE-3908: give the turns we store ourselves one owner OGX appends a turn to the conversation only when it is handed the conversation parameter and runs the inference. In compacted mode that parameter is dropped, so lightspeed-stack appends the turn itself. Every endpoint made that write on its own, and nothing enforced it: two unrelated cleanups removed the calls, and for six weeks conversations stopped growing after a compaction without anything failing (LCORE-3883). The write now has one owner, PendingTurn in utils/pending_turn.py. An endpoint creates it from the parameters the request is sent with and the original input of the compaction result, and reports how the turn ended: store_completed, store_blocked, store_interrupted or drop. The owner decides whether the turn is ours, which input is stored, and it stores a turn once: the first report settles the turn, every later one does nothing. It covers all nine places that stored a turn, not only the compacted ones, because they are the same duty: a turn OGX does not store. That is a conversation in compacted mode, a request a shield blocked, a stream the client interrupted, a continuation from previous_response_id. store_compacted_turn is removed; no endpoint appends a turn except through the owner. The shield capabilities still store the turn they rejected themselves, from inside the agent run, and only when the model was handed the conversation, so never in compacted mode. A compacted request that loses its turn through a change in the code fails: - creating the owner for compacted parameters without the original input raises ValueError, which is how a caller that lost the hand-over looks; - ensure_settled raises TurnNotStoredError when a compacted request reaches the end of its handler and nobody tried to store its turn or dropped it on purpose. /v1/query runs in a scope that makes the check when it is left; the streaming paths, /v1/responses and A2A make it after their write, in the function that calls the write, not in the one that performs it. The error is not a RuntimeError, because the endpoints report that one as an inference failure. The check does not cover a stream the client stops reading, a /v1/responses stream without a final response, and a failed write that is logged. They store nothing, as before. Every condition of the old call sites is kept: where the write happens, what is stored for a failed or incomplete turn, what a failed write does to the request (it fails /v1/query and /v1/responses, it is logged on /v1/streaming_query, A2A and the interrupt path), and the order relative to quota consumption. /v1/query still stores nothing for a request blocked on a compacted conversation; that is now an explicit drop, and LCORE-3788 settles what all endpoints should do. The guard that lets only one of stream end, cancellation handler and interrupt callback finish a turn is unchanged. One behaviour changes, on purpose. A blocked request on a conversation that is not compacted has its refusal turn stored before the stream starts. An interrupt while the refusal was streamed stored the turn a second time, with the interruption notice for an answer. It is stored once now; the interrupt still records the turn in the database. The path cannot be reached today, because run_shield_moderation always passes. Tests: - integration, tests/integration/endpoints/test_turn_persistence.py: for /v1/query, /v1/streaming_query, /v1/responses (blocking and streaming) and A2A, every way a turn can end (completed, blocked, interrupted, failed, cut short by the model, abandoned by the client, continued from a previous response, store off), compacted and not; what a failed write does to the request on each endpoint; and what happens when the step that stores the turn is taken out of an endpoint. Each test runs the real handler and compares the conversation item by item afterwards, so a missing write, a second write and a write of the wrong input all fail it. 44 of the 50 tests pass on main unchanged. The six that fail there are the double write above and the five that take the write out of an endpoint, which main does not notice. - unit: the owner, the hand-over from the compaction seam (omit_conversation and the original input are set together), and a scan of src that fails when a module other than the known writers mentions one of the functions that append to a conversation. - existing unit tests follow the new parameter. Those that asserted on the arguments of a patched helper in the agent and interrupt paths now assert on what is stored. Mocked request parameters got the fields the owner reads. The A2A helpers of the compaction tests moved to _compaction_helpers.py, where both test modules import them. The tests were checked by mutation: with the write removed from any of the four endpoints, a check removed, the write doubled, the settled state ignored, the interrupt guard ignored, the explicit rewrite stored in place of the input, a turn OGX stores stored by us as well, or the store flag ignored, tests fail. Docs: the design document describes the owner and has a table of what is stored per endpoint and per way a turn can end. The endpoint guide names the exception for an interrupted stream. --- .../conversation-compaction.md | 62 + docs/devel_doc/ARCHITECTURE.md | 2 +- docs/devel_doc/query_endpoint.md | 2 +- src/app/endpoints/a2a.py | 78 +- src/app/endpoints/query.py | 28 +- src/app/endpoints/responses.py | 89 +- src/app/endpoints/streaming_query.py | 16 +- src/utils/agents/query.py | 40 +- src/utils/agents/streaming.py | 70 +- src/utils/conversation_compaction.py | 25 +- src/utils/pending_turn.py | 301 ++++ src/utils/stream_interrupts.py | 52 +- .../endpoints/_compaction_helpers.py | 96 ++ .../endpoints/test_compaction_a2a.py | 121 +- .../endpoints/test_turn_persistence.py | 1534 +++++++++++++++++ tests/unit/app/endpoints/test_a2a.py | 36 + tests/unit/app/endpoints/test_query.py | 28 + tests/unit/app/endpoints/test_query_otel.py | 4 + tests/unit/app/endpoints/test_responses.py | 40 +- .../app/endpoints/test_streaming_query.py | 24 + tests/unit/utils/agents/test_query.py | 102 +- tests/unit/utils/agents/test_streaming.py | 146 +- .../utils/test_conversation_compaction.py | 11 - tests/unit/utils/test_pending_turn.py | 416 +++++ tests/unit/utils/test_stream_interrupts.py | 115 +- 25 files changed, 2997 insertions(+), 441 deletions(-) create mode 100644 src/utils/pending_turn.py create mode 100644 tests/integration/endpoints/test_turn_persistence.py create mode 100644 tests/unit/utils/test_pending_turn.py diff --git a/docs/design/conversation-compaction/conversation-compaction.md b/docs/design/conversation-compaction/conversation-compaction.md index 6c70380be..7a5cb77cd 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) | @@ -346,6 +347,67 @@ 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`, `store_interrupted` or `drop`. +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, 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 or dropped it on purpose. `/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 | +| | blocked by a shield | stored | not stored (LCORE-3788) | 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 | +| | blocked by a shield | stored before the stream starts | stored when the stream ends | request fails / logged | +| | interrupted by the client | stored, with the answer so far | stored, with the answer so far | logged | +| | blocked by a shield, then interrupted | the refusal, stored once before the stream; the interrupt adds nothing | 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, for a blocked and +for an interrupted stream. + +The row "blocked by a shield, then interrupted" is the one place where this +change alters what is stored. The interrupt used to store the turn a second +time, with the interruption notice for an answer. It still records the turn in +the database. + ## 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/docs/devel_doc/query_endpoint.md b/docs/devel_doc/query_endpoint.md index bc431c969..066ca910d 100644 --- a/docs/devel_doc/query_endpoint.md +++ b/docs/devel_doc/query_endpoint.md @@ -308,7 +308,7 @@ Cancels an in-progress streaming query. | `interrupted` | boolean | Whether an active stream was interrupted (`false` if already completed) | | `message` | string | Human-readable status message | -When a stream is interrupted, any partial response is persisted to conversation history and token consumption is skipped. +When a stream is interrupted, any partial response is persisted to conversation history and token consumption is skipped. A request a shield had blocked is the exception: its refusal turn is stored before the stream starts, so an interrupt adds nothing to the conversation. --- diff --git a/src/app/endpoints/a2a.py b/src/app/endpoints/a2a.py index 793f69ac2..287c124c5 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 @@ -206,10 +203,39 @@ def _record_execution_span( span.set_attribute(SpanAttributes.OUTPUT, output_text) -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: @@ -221,22 +247,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. @@ -471,19 +488,9 @@ async def _process_task_streaming( # pylint: disable=too-many-locals 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( @@ -574,9 +581,8 @@ async def _process_task_streaming( # pylint: disable=too-many-locals ) 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) diff --git a/src/app/endpoints/query.py b/src/app/endpoints/query.py index 64e7b56e3..1204dd19f 100644 --- a/src/app/endpoints/query.py +++ b/src/app/endpoints/query.py @@ -46,6 +46,7 @@ anonymize_value, set_span_attributes, ) +from utils.pending_turn import pending_turn from utils.query import ( consume_query_tokens, prepare_input, @@ -265,17 +266,22 @@ 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, - moderation_result, - 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 the + # turn or dropped it on purpose. + async with pending_turn( + client, responses_params, compaction.original_input + ) as turn: + turn_summary = await retrieve_agent_response( + client, + responses_params, + moderation_result, + endpoint_path, + turn, + shield_ids=query_request.shield_ids, + no_tools=bool(query_request.no_tools), + image_attachments=image_attachments, + ) if moderation_result.decision == "passed": # Combine inline RAG results (BYOK + Solr) with tool-based RAG results for the transcript diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index ca154705b..58335ea0a 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, @@ -82,6 +81,7 @@ record_exception, set_span_attributes, ) +from utils.pending_turn import PendingTurn from utils.prompts import get_system_prompt from utils.query import ( consume_query_tokens, @@ -368,38 +368,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. @@ -414,23 +428,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( @@ -803,7 +805,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, @@ -1251,13 +1255,17 @@ async def response_generator( turn_summary, ) - # 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( """ root_span = context.root_span user_id = context.auth[0] + 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() @@ -1384,11 +1393,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 ( @@ -1406,6 +1417,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() + # Get available quotas logger.info("Getting available quotas") available_quotas = get_available_quotas( diff --git a/src/app/endpoints/streaming_query.py b/src/app/endpoints/streaming_query.py index 2c000ae04..151e3643f 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, @@ -352,12 +352,16 @@ async def _handle_streaming_query_with_tracing( # pylint: disable=too-many-loca media_type=response_media_type, ) + # This request is not compacted, so OGX stores its turn when it completes. + # A turn that is blocked or interrupted is still ours to store. + turn = PendingTurn.for_request(context.client, responses_params) generator, turn_summary = await retrieve_agent_response_generator( responses_params=responses_params, context=context, endpoint_path=endpoint_path, no_tools=bool(query_request.no_tools), image_attachments=image_attachments, + turn=turn, ) # Combine inline RAG results (BYOK + Solr) with tool-based results @@ -374,6 +378,7 @@ async def _handle_streaming_query_with_tracing( # pylint: disable=too-many-loca turn_summary=turn_summary, background_topic_summary_tasks=_background_topic_summary_tasks, root_span=root_span, + turn=turn, ), media_type=response_media_type, ) @@ -430,7 +435,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( @@ -447,7 +452,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( @@ -455,6 +462,7 @@ async def generate_response_with_compaction( context=context, endpoint_path=endpoint_path, image_attachments=image_attachments, + turn=turn, ) except HTTPException as e: yield http_exception_stream_event(e) @@ -502,7 +510,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 fb78c1dac..e7b6d036c 100644 --- a/src/utils/agents/query.py +++ b/src/utils/agents/query.py @@ -27,7 +27,6 @@ from models.common.moderation import ShieldModerationResult 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,15 +38,14 @@ from utils.conversation_compaction import ( agent_prompt_text, reject_image_attachments_in_compacted_mode, - store_compacted_turn, ) -from utils.conversations import append_turn_items_to_conversation from utils.otel_tracing import ( SpanAttributes, SpanEvents, add_span_event, 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, @@ -238,7 +236,7 @@ async def retrieve_agent_response( responses_params: ResponsesApiParams, moderation_result: ShieldModerationResult, 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, @@ -250,9 +248,10 @@ async def retrieve_agent_response( responses_params: Prepared Responses API parameters. moderation_result: Shield moderation outcome for the turn. 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 @@ -262,7 +261,10 @@ async def retrieve_agent_response( Raises: HTTPException: On moderation is not applicable; 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( @@ -279,13 +281,13 @@ async def retrieve_agent_response( ) if moderation_result.decision == "blocked": - if not responses_params.omit_conversation: - await append_turn_items_to_conversation( - client, - responses_params.conversation, - responses_params.input, - [moderation_result.refusal_response], - ) + if responses_params.omit_conversation: + # Kept as it was: on this endpoint the refusal turn of a + # compacted conversation is not stored. LCORE-3788 settles + # what all endpoints should do with it. + turn.drop("blocked by a shield in compacted mode") + else: + await turn.store_blocked(moderation_result.refusal_response) return TurnSummary( id=moderation_result.moderation_id, llm_response=moderation_result.message, @@ -345,14 +347,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 06192e27d..f8930d726 100644 --- a/src/utils/agents/streaming.py +++ b/src/utils/agents/streaming.py @@ -43,7 +43,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 @@ -63,15 +62,14 @@ from utils.conversation_compaction import ( agent_prompt_text, reject_image_attachments_in_compacted_mode, - store_compacted_turn, ) -from utils.conversations import append_turn_items_to_conversation from utils.otel_tracing import ( SpanAttributes, SpanEvents, add_span_event, 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, @@ -108,6 +106,7 @@ async def retrieve_agent_response_generator( endpoint_path: str, no_tools: bool = False, image_attachments: Optional[list[Attachment]] = None, + turn: Optional[PendingTurn] = None, ) -> tuple[AsyncIterator[str], TurnSummary]: """Return the SSE generator and mutable turn summary for an agent run. @@ -117,6 +116,9 @@ async def retrieve_agent_response_generator( endpoint_path: Endpoint path used for metric labeling. no_tools: Whether to skip tool processing. image_attachments: Image attachments for multimodal prompt construction. + turn: The pending turn of the request, shared with + :func:`generate_agent_response` so that the turn is stored once + (LCORE-3908). Created from the parameters when not given. Returns: Tuple of SSE async iterator and mutable turn summary. @@ -127,13 +129,13 @@ async def retrieve_agent_response_generator( turn_summary.llm_response = context.moderation_result.message turn_summary.id = context.moderation_result.moderation_id turn_summary.output_items = [context.moderation_result.refusal_response] + # Outside compacted mode the refusal turn is stored before the + # stream starts. In compacted mode it is stored when the stream + # ends, as the output of the turn. if not responses_params.omit_conversation: - await append_turn_items_to_conversation( - context.client, - responses_params.conversation, - responses_params.input, - [context.moderation_result.refusal_response], - ) + if turn is None: + turn = PendingTurn.for_request(context.client, responses_params) + await turn.store_blocked(context.moderation_result.refusal_response) media_type = context.query_request.media_type or MEDIA_TYPE_JSON return ( shield_violation_generator( @@ -169,9 +171,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). @@ -181,24 +182,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. @@ -216,7 +211,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. @@ -234,23 +229,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: @@ -290,7 +293,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( @@ -309,9 +312,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 f04ae63b4..66e9e8c39 100644 --- a/src/utils/conversation_compaction.py +++ b/src/utils/conversation_compaction.py @@ -67,7 +67,6 @@ summarize_chunk, ) from utils.conversations import ( - append_turn_items_to_conversation, build_add_items_request, get_all_conversation_items, ) @@ -175,9 +174,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 @@ -778,21 +777,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..e63ac1651 --- /dev/null +++ b/src/utils/pending_turn.py @@ -0,0 +1,301 @@ +"""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, or having dropped it on purpose, 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 log import get_logger +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 + +logger = get_logger(__name__) + + +class TurnNotStoredError(Exception): + """Nobody tried to store a turn lightspeed-stack has to store, or dropped it. + + 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`, :meth:`store_interrupted` and :meth:`drop` 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 or left out on purpose. + """ + 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 drop(self, reason: str) -> None: + """Leave the turn out of the conversation on purpose. + + Parameters: + reason: Why the turn is not stored; kept as the outcome and logged. + """ + if self._settle(f"dropped: {reason}") and self.ours: + logger.info( + "Turn on conversation %s is not stored: %s", + self.conversation_id, + reason, + ) + + 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 or dropped 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 286ec2dd6..114d452e6 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 @@ -255,3 +265,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..cd412a020 --- /dev/null +++ b/tests/integration/endpoints/test_turn_persistence.py @@ -0,0 +1,1534 @@ +"""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_blocked_turn_outside_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, + mocker: MockerFixture, + ) -> None: + """A blocked request never reaches OGX, so the refusal turn is ours.""" + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + mocker.patch( + "app.endpoints.query.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + + await _send_query(test_request, test_auth) + + mock_query_agent.run.assert_not_awaited() + await _assert_stored( + mock_conversation_store, _plain_conversation() + _turn(REFUSAL) + ) + + @pytest.mark.asyncio + async def test_blocked_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, + mocker: MockerFixture, + ) -> None: + """A blocked request on a compacted conversation leaves no turn behind. + + This is the behaviour as it is, recorded so that a change to it is a + decision: /v1/streaming_query and /v1/responses do store this turn, + and LCORE-3788 settles what all of them should do. + """ + _ = 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.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + + await _send_query(test_request, test_auth) + + mock_query_agent.run.assert_not_awaited() + await _assert_stored(mock_conversation_store, _compacted_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_blocked_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, + mocker: MockerFixture, + ) -> None: + """The refusal turn is stored once, whichever mode the conversation is in. + + Outside compacted mode it is written before the stream starts, in + compacted mode when the stream ends. Neither may also do the other's. + """ + _ = mock_ogx_client + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, stored_conversation()) + mocker.patch( + "app.endpoints.streaming_query.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + + await _drain(await _send_streaming_query(test_request, test_auth)) + + mock_streaming_query_agent.build_agent_mock.assert_not_called() + await _assert_stored( + mock_conversation_store, stored_conversation() + _turn(REFUSAL) + ) + + @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()) + + @pytest.mark.asyncio + async def test_blocked_turn_is_not_stored_again_by_an_interrupt( + 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 refusal that was stored is the turn; an interrupt adds no second one. + + Outside compacted mode the refusal turn is stored before the stream + starts. An interrupt while the refusal is streamed used to store the + turn again, with the interruption notice for an answer. The turn has + one owner now, so it is stored once; the interrupt still records the + turn in the database, as before. + """ + _ = mock_ogx_client + _ = mock_streaming_query_agent + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + mocker.patch( + "app.endpoints.streaming_query.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + + async def _refusal_that_stalls(*_args: Any, **_kwargs: Any) -> Any: + yield 'data: {"event": "token", "data": {"id": 0, "token": "Content"}}\n\n' + await asyncio.Event().wait() + + mocker.patch( + "utils.agents.streaming.shield_violation_generator", + side_effect=_refusal_that_stalls, + ) + + chunks = await _interrupt_stream(test_request, test_auth) + + assert any('"event": "interrupted"' in chunk for chunk in chunks) + await _assert_stored( + mock_conversation_store, _plain_conversation() + _turn(REFUSAL) + ) + recorded_turns = ( + patch_db_session.query(UserTurn) + .filter_by(conversation_id=EXISTING_CONV_ID) + .count() + ) + assert recorded_turns == 1 + + +# ========================================== +# /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_blocked_streaming_query_fails_before_the_stream( + 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: + """Outside compacted mode a refusal that cannot be stored is an error.""" + _ = mock_streaming_query_agent + create_existing_conversation(patch_db_session, test_auth[0]) + await _seed(test_config, mock_conversation_store, _plain_conversation()) + mocker.patch( + "app.endpoints.streaming_query.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + _break_the_store(mock_ogx_client, mocker) + + with pytest.raises(HTTPException) as error: + await _send_streaming_query(test_request, test_auth) + + assert error.value.status_code == 500 + assert mock_ogx_client.items.create.await_count == 1 + + @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 ac0bf936b..920398ae4 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), @@ -1284,6 +1300,8 @@ async def test_execute_span_success_attributes( # pylint: disable=too-many-loca 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), @@ -1373,6 +1391,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), @@ -1380,6 +1402,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), @@ -1483,6 +1507,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), @@ -1490,6 +1518,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), @@ -1701,6 +1731,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), @@ -1708,6 +1742,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 3710401ab..42c767b77 100644 --- a/tests/unit/app/endpoints/test_query.py +++ b/tests/unit/app/endpoints/test_query.py @@ -135,6 +135,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", @@ -221,6 +225,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), @@ -310,6 +318,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", @@ -385,6 +397,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", @@ -466,6 +482,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", @@ -533,6 +553,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", @@ -607,6 +631,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 84e325230..75bdb0a66 100644 --- a/tests/unit/app/endpoints/test_query_otel.py +++ b/tests/unit/app/endpoints/test_query_otel.py @@ -90,6 +90,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 e42eea2e8..357971379 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/app/endpoints/test_streaming_query.py b/tests/unit/app/endpoints/test_streaming_query.py index bdfa73f46..12f27b81e 100644 --- a/tests/unit/app/endpoints/test_streaming_query.py +++ b/tests/unit/app/endpoints/test_streaming_query.py @@ -148,6 +148,10 @@ async def test_successful_streaming_query( 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", @@ -235,6 +239,10 @@ async def test_streaming_query_text_media_type_header( 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", @@ -333,6 +341,10 @@ async def test_streaming_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", @@ -429,6 +441,10 @@ async def test_streaming_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", @@ -519,6 +535,10 @@ async def test_streaming_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", @@ -622,6 +642,10 @@ def _setup_common_mocks( 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", diff --git a/tests/unit/utils/agents/test_query.py b/tests/unit/utils/agents/test_query.py index 74ee49d0f..3b25dc33b 100644 --- a/tests/unit/utils/agents/test_query.py +++ b/tests/unit/utils/agents/test_query.py @@ -36,6 +36,7 @@ get_agent_finish_reason, retrieve_agent_response, ) +from utils.pending_turn import PendingTurn from utils.token_counter import TokenCounter @@ -362,6 +363,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.""" @@ -373,25 +388,19 @@ async def test_blocked_moderation_returns_refusal_summary( blocked_moderation: ShieldModerationBlocked, ) -> None: """Test blocked moderation persists refusal and returns a turn summary.""" - mock_client = mocker.AsyncMock() - mock_append = mocker.patch( - "utils.agents.query.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) + client, stored = client_capturing_writes(mocker) summary = await retrieve_agent_response( - client=mock_client, + client=client, responses_params=responses_params, moderation_result=blocked_moderation, endpoint_path=ENDPOINT_PATH_QUERY, ) - mock_append.assert_awaited_once_with( - mock_client, - responses_params.conversation, - responses_params.input, - [blocked_moderation.refusal_response], - ) + texts = [str(item) for item in stored] + assert len(stored) == 2, f"expected user turn + refusal, got {texts}" + assert str(responses_params.input) in texts[0] + assert blocked_moderation.message in texts[1] assert summary == TurnSummary( id="modr-test-456", llm_response="Content blocked by shield.", @@ -450,12 +459,14 @@ 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, moderation_result=ShieldModerationPassed(), endpoint_path=ENDPOINT_PATH_QUERY, + turn=PendingTurn.for_request(client, params, "new question"), ) mock_agent.run.assert_awaited_once_with("new question") @@ -468,28 +479,58 @@ async def test_blocked_moderation_compacted_skips_append( make_responses_params: Callable[..., ResponsesApiParams], blocked_moderation: ShieldModerationBlocked, ) -> None: - """Test blocked moderation does not append explicit input in compacted mode.""" + """Test blocked moderation stores nothing in compacted mode, on purpose.""" params = make_responses_params().model_copy( update={ "input": [OpenAIResponseMessage(role="user", content="q")], "omit_conversation": True, } ) - mock_append = mocker.patch( - "utils.agents.query.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) + client, stored = client_capturing_writes(mocker) + turn = PendingTurn.for_request(client, params, "q") summary = await retrieve_agent_response( - client=mocker.AsyncMock(), + client=client, responses_params=params, moderation_result=blocked_moderation, endpoint_path=ENDPOINT_PATH_QUERY, + turn=turn, ) - mock_append.assert_not_awaited() + assert stored == [] + assert turn.outcome == "dropped: blocked by a shield in compacted mode" assert summary.llm_response == "Content blocked by shield." + @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, + moderation_result=ShieldModerationPassed(), + endpoint_path=ENDPOINT_PATH_QUERY, + ) + + mock_build_agent.assert_not_called() + @pytest.mark.asyncio async def test_inference_error_is_logged( self, @@ -673,21 +714,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, @@ -713,14 +739,14 @@ 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, moderation_result=ShieldModerationPassed(), endpoint_path=ENDPOINT_PATH_QUERY, - original_input="new question", + turn=PendingTurn.for_request(client, params, "new question"), ) texts = [str(item) for item in stored] @@ -745,7 +771,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 bdf6b8863..3e5dbb14a 100644 --- a/tests/unit/utils/agents/test_streaming.py +++ b/tests/unit/utils/agents/test_streaming.py @@ -68,6 +68,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}*" @@ -503,6 +504,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.""" @@ -522,10 +536,7 @@ async def test_blocked_moderation_returns_shield_generator( "utils.agents.streaming.shield_violation_generator", return_value=_async_iter(["shield-event"]), ) - mock_append = mocker.patch( - "utils.agents.streaming.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) + stored = capture_conversation_writes(context) generator, turn_summary = await retrieve_agent_response_generator( responses_params, @@ -539,7 +550,10 @@ async def test_blocked_moderation_returns_shield_generator( blocked_moderation.message, MEDIA_TYPE_JSON, ) - mock_append.assert_awaited_once() + texts = [str(item) for item in stored] + assert len(stored) == 2, f"expected user turn + refusal, got {texts}" + assert str(responses_params.input) in texts[0] + assert blocked_moderation.message in texts[1] assert turn_summary.llm_response == blocked_moderation.message assert turn_summary.id == blocked_moderation.moderation_id @@ -558,10 +572,7 @@ async def test_blocked_moderation_skips_append_when_omit_conversation( "utils.agents.streaming.shield_violation_generator", return_value=_async_iter([]), ) - mock_append = mocker.patch( - "utils.agents.streaming.append_turn_items_to_conversation", - new=mocker.AsyncMock(), - ) + stored = capture_conversation_writes(context) await retrieve_agent_response_generator( responses_params, @@ -569,7 +580,7 @@ async def test_blocked_moderation_skips_append_when_omit_conversation( ENDPOINT_PATH_STREAMING_QUERY, ) - mock_append.assert_not_awaited() + assert stored == [] @pytest.mark.asyncio async def test_success_returns_agent_generator( @@ -1645,20 +1656,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.""" @@ -1680,11 +1677,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() @@ -1704,10 +1702,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(), ) ] @@ -1727,7 +1727,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() @@ -1757,13 +1757,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] ) @@ -1784,25 +1793,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: @@ -1826,10 +1894,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 af1e40b71..cb5d36128 100644 --- a/tests/unit/utils/test_conversation_compaction.py +++ b/tests/unit/utils/test_conversation_compaction.py @@ -321,17 +321,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..e5ef3fd14 --- /dev/null +++ b/tests/unit/utils/test_pending_turn.py @@ -0,0 +1,416 @@ +"""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: + """The refusal written before the stream started is the turn; an interrupt 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_dropped_turn_is_not_stored_later() -> None: + """A turn left out on purpose stays out.""" + client = RecordingClient() + turn = _turn(client, compacted=True) + + turn.drop("the request was blocked") + + assert turn.settled + assert await turn.store_completed([ANSWER]) is False + assert client.writes == 0 + + +@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 was neither stored nor dropped 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, dropped and left to OGX are all fine.""" + stored = _turn(RecordingClient(), compacted=True) + await stored.store_completed([ANSWER]) + stored.ensure_settled() + + dropped = _turn(RecordingClient(), compacted=True) + dropped.drop("the model call failed") + dropped.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", + ], +} + + +@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", From 5411cdfcf233686ae4b7f386801369eb19b1b52e Mon Sep 17 00:00:00 2001 From: Maxim Svistunov Date: Mon, 28 Sep 2026 14:46:21 +0200 Subject: [PATCH 2/2] LCORE-3910: count the summarization calls in quota, token counts and metrics Compaction summarizes older turns with an LLM call of its own, and folds the summaries with another. The provider bills both. Their usage was discarded: summarize_chunk() and recursively_resummarize() read the text of the response and nothing else. So the calls consumed no quota, were missing from the input_tokens / output_tokens the client is told, and were missing from the token and call metrics. A user could go over the quota without it showing, by more the longer the conversation was. The user guide said the opposite. What is decided (by the owner of the ticket): - the summarization tokens are charged to the user's quota; - they are part of the existing input_tokens / output_tokens of /v1/query and of the end event of /v1/streaming_query, no new fields; - /v1/responses passes the usage object of the response through unchanged, to stay a drop-in for clients of the OpenAI Responses API; the quota is charged and available_quotas shows it; - /a2a has no quota, so the calls show in the metrics only. How it works: - summarize_chunk() and recursively_resummarize() take a count_call callback and report each LLM call to it, with the model and the usage the provider reported. They do so as soon as the response arrived and before they look at it, so a call that returned no text is reported as well. - apply_compaction() takes the path of the endpoint and a charge callback, and counts the calls with SummarizationCalls (new module utils/compaction_usage.py). For each call it records ls_llm_token_sent_total, ls_llm_token_received_total and ls_llm_calls_total under that endpoint, calls charge, and adds the usage up. The sum is returned as CompactionResult.summarization_usage. Without an endpoint path no metrics are recorded, without charge nobody is charged. The counter has a module of its own because conversation_compaction.py is close to the 1000 lines pylint allows once the open pull requests on it are merged. - The endpoints with a quota pass consume_summarization_tokens(), bound to the user, as charge. A call is therefore charged when it returned, and not together with the turn. The summary is written before the model is asked for the answer and is kept whatever becomes of the turn, so the call is charged also when the turn is blocked by a shield, fails or is interrupted, and when compaction itself fails after the call (the marker cannot be written, the fold fails). The next request is served from the stored summary and does not pay again. - /v1/query and /v1/streaming_query add the summarization usage to the usage of the turn in what they report to the client and set on the request span. The turn summary itself is left as it is. TokenCounter gained __add__ for the sum; it returns a new counter. - /a2a passes its endpoint path (new constant ENDPOINT_PATH_A2A) and no charge. Known limits: - A call that fails, or is cancelled before its response arrived, has no usage to count, whatever the provider bills for it. - A call the provider reports no usage for is counted as a call and charges nothing. Tests: - tests/integration/endpoints/test_compaction_token_usage.py runs the four endpoints against a real quota limiter on SQLite and the real summarization path, answered by the mocked OGX client. It reads what the limiter holds after the request, what the client is told, and the Prometheus registry. It covers a turn that summarizes, one that also folds (against a real SQLite conversation cache), a turn served from a stored summary, a short conversation, and a turn that is blocked, fails or is interrupted after the summarization. 14 of the 19 tests fail without the endpoint changes; the other five are the controls that nothing extra is charged. - tests/unit/utils/test_summarization_usage.py covers what the two calls report, and what a request records, charges and is told, including a call that returned no text, a marker that cannot be written, a fold that fails and a fold that cannot be stored. - tests/unit/utils/test_query.py covers consume_summarization_tokens(). Documentation: the user guide gets a section on token usage and quota and the FAQ answers are made exact per endpoint; the design document describes who pays for the calls; the statements about failed and interrupted requests in ARCHITECTURE.md and query_endpoint.md now name the exception; the open question in the guardrails design document is updated. --- .../conversation-compaction.md | 54 +- .../prompt-guardrails/prompt-guardrails.md | 3 +- docs/devel_doc/ARCHITECTURE.md | 6 +- docs/devel_doc/query_endpoint.md | 5 +- docs/user_doc/conversation_compaction.md | 37 +- src/app/endpoints/a2a.py | 3 +- src/app/endpoints/query.py | 15 +- src/app/endpoints/responses.py | 7 + src/app/endpoints/streaming_query.py | 8 + src/constants.py | 1 + src/utils/agents/streaming.py | 20 +- src/utils/compaction.py | 47 +- src/utils/compaction_usage.py | 63 ++ src/utils/conversation_compaction.py | 63 +- src/utils/query.py | 27 + src/utils/token_counter.py | 16 + .../endpoints/test_compaction_token_usage.py | 853 ++++++++++++++++++ tests/unit/utils/test_query.py | 60 ++ tests/unit/utils/test_summarization_usage.py | 561 ++++++++++++ 19 files changed, 1816 insertions(+), 33 deletions(-) create mode 100644 src/utils/compaction_usage.py create mode 100644 tests/integration/endpoints/test_compaction_token_usage.py create mode 100644 tests/unit/utils/test_summarization_usage.py diff --git a/docs/design/conversation-compaction/conversation-compaction.md b/docs/design/conversation-compaction/conversation-compaction.md index 7a5cb77cd..2d430c3dd 100644 --- a/docs/design/conversation-compaction/conversation-compaction.md +++ b/docs/design/conversation-compaction/conversation-compaction.md @@ -311,7 +311,7 @@ Add `compaction` field to the root `Configuration` class. |----------------------------------------|----------------------------------------------------------------------------| | `pyproject.toml` | Add `tiktoken` dependency | | `src/utils/token_estimator.py` | New module: `estimate_tokens()`, `estimate_conversation_tokens()` | -| `src/utils/compaction.py` | New module: summarization logic, partitioning, additive summary management | +| `src/utils/compaction.py` | New module: summarization logic, partitioning, additive summary management; reports the usage of each LLM call it makes — LCORE-3910 | | `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 | @@ -332,9 +332,11 @@ reusable unit in `src/utils/conversation_compaction.py` that each endpoint calls after its params are prepared: - `apply_compaction_blocking(client, params, inference_config, compaction_config)` - returns a `CompactionResult` (possibly-rewritten params, a `summarized` - flag, and the `original_input`). Non-streaming `/v1/query`, A2A, and - `/v1/responses` use this. + returns a `CompactionResult` (possibly-rewritten params, a `compacted` + flag, the `original_input`, and the `summarization_usage`). Non-streaming + `/v1/query`, A2A, and `/v1/responses` use this. The endpoints also pass + `endpoint_path` and, where there is a quota, `charge`; see + [Who pays for the summarization calls](#who-pays-for-the-summarization-calls). - `apply_compaction(..., emit_events=True)` is the async-generator variant that yields a `CompactionStartedEvent` before the summarization LLM call; the native `/v1/streaming_query` SSE path uses it to satisfy R12. @@ -447,6 +449,50 @@ Example config files go in `examples/`. Compaction adds latency only on the trigger turn. In PoC testing, compaction turns took 14-40 seconds vs 9-20 seconds for normal turns (gpt-4o-mini). +## Who pays for the summarization calls + +The provider bills the summarization call and the fold call like any other, so +lightspeed-stack counts them (LCORE-3910). Before that, their usage was +discarded: a user could go over the quota without it ever showing, and by more +the longer the conversation was. + +`summarize_chunk()` and `recursively_resummarize()` report each call they make +to a `count_call` callback, with the model and the usage the provider reported. +They do it as soon as the response arrived and before they look at it, so a +call that returned no text is reported too. `apply_compaction()` takes the +path of the endpoint and a `charge` callback, and counts the calls with +`SummarizationCalls` (`src/utils/compaction_usage.py`): for each call it +records the token and call metrics under that endpoint, calls `charge`, and +adds the usage up. The sum is handed to the endpoint as +`CompactionResult.summarization_usage`. + +| Endpoint | Quota | Counts the client sees | +|------------------------|---------|----------------------------------------------------------| +| `/v1/query` | charged | `input_tokens` / `output_tokens` include summarization | +| `/v1/streaming_query` | charged | the same, in the `end` event | +| `/v1/responses` | charged | `usage` as the provider reported it, the answer alone | +| `/a2a` | none | none; the endpoint has no quota, the metrics are recorded | + +The endpoints with a quota pass `consume_summarization_tokens()` +(`src/utils/query.py`), bound to the user, as `charge`. A call is therefore +charged when it returned, and not together with the turn: the summary is +written before the model is asked for the answer and is kept whatever becomes +of the turn, so the call is a cost of the conversation. It is charged also +when the turn is blocked, fails or is interrupted, and when compaction itself +fails after the call (the marker cannot be written, the fold fails). `/a2a` +passes no `charge`. + +`/v1/query` and `/v1/streaming_query` add `summarization_usage` to the usage of +the turn in what they report to the client and set on the request span +(`llm.usage.input_tokens`, `llm.usage.output_tokens`). The sum is not stored +with the turn; the token usage history receives the charges one by one. +`/v1/responses` passes the `usage` object of the response through unchanged, to +stay a drop-in for clients of the OpenAI Responses API. + +Two limits. A call that fails, or is cancelled before its response arrived, +has no usage to count, whatever the provider bills for it. And a call the +provider reports no usage for is counted as a call and charges nothing. + # Open Questions for Future Work - **Compaction-proof instructions**: Allow "pinned" messages that always survive compaction (inspired by Claude Code's CLAUDE.md pattern). Not needed for v1. diff --git a/docs/design/prompt-guardrails/prompt-guardrails.md b/docs/design/prompt-guardrails/prompt-guardrails.md index 3cfee9181..886c5fc0e 100644 --- a/docs/design/prompt-guardrails/prompt-guardrails.md +++ b/docs/design/prompt-guardrails/prompt-guardrails.md @@ -442,7 +442,8 @@ product need justifies it. version-specific prompt. - **Guardian token usage:** whether guardian calls should count against user quota or be tracked as service overhead. Compaction's summarization calls - raise the same question. + are charged to the user's quota (LCORE-3910); the same choice is open for + guardian calls. - **Streaming checkpoint sizing:** defaults for LCORE-3391 (spike Decision T4, 70% confidence); tune with real latency data. - **Cheap classifier tier for `tool`:** Prompt Guard 2-class, and its diff --git a/docs/devel_doc/ARCHITECTURE.md b/docs/devel_doc/ARCHITECTURE.md index 1eeb76b1d..0b8b012b4 100644 --- a/docs/devel_doc/ARCHITECTURE.md +++ b/docs/devel_doc/ARCHITECTURE.md @@ -274,9 +274,9 @@ The system defines 30+ actions that can be authorized. Examples (see `docs/user_ - Update quota counters 3. **On Error:** - - If LLM call fails, no tokens are consumed - - Quota remains unchanged - - User can retry the request + - If LLM call fails, no tokens are consumed for that call + - Quota remains unchanged, with one exception: the calls conversation compaction made to summarize older turns for the request were charged when they were made + - User can retry the request; the summary is kept, so the retry does not pay for it again --- diff --git a/docs/devel_doc/query_endpoint.md b/docs/devel_doc/query_endpoint.md index 066ca910d..76e1aaa2d 100644 --- a/docs/devel_doc/query_endpoint.md +++ b/docs/devel_doc/query_endpoint.md @@ -308,7 +308,7 @@ Cancels an in-progress streaming query. | `interrupted` | boolean | Whether an active stream was interrupted (`false` if already completed) | | `message` | string | Human-readable status message | -When a stream is interrupted, any partial response is persisted to conversation history and token consumption is skipped. A request a shield had blocked is the exception: its refusal turn is stored before the stream starts, so an interrupt adds nothing to the conversation. +When a stream is interrupted, any partial response is persisted to conversation history and token consumption for the answer is skipped. The summarization calls conversation compaction made for the request were charged when they were made (see [Quota and Token Counting](#quota-and-token-counting)). A request a shield had blocked is the exception: its refusal turn is stored before the stream starts, so an interrupt adds nothing to the conversation. --- @@ -361,8 +361,9 @@ If the server configuration sets `disable_query_system_prompt` to `true`, reques - **Pre-request:** `check_tokens_available()` verifies the user/cluster has available quota (429 if not) - **Post-response:** `consume_query_tokens()` deducts `input_tokens` and `output_tokens` from configured quota limiters +- **Compaction:** each LLM call that summarizes older turns is charged by `consume_summarization_tokens()` as soon as it returned, so also when the turn is then blocked, fails or is interrupted. The `input_tokens` and `output_tokens` the response reports include these calls - **Available quotas:** Remaining balances per limiter are included in the response (`available_quotas` field in sync, `end` event in streaming) -- **Stream interruption:** Token consumption is skipped for interrupted streams +- **Stream interruption:** Token consumption is skipped for interrupted streams, except for the summarization calls, which were charged before the answer started --- diff --git a/docs/user_doc/conversation_compaction.md b/docs/user_doc/conversation_compaction.md index f29887a60..51b394a23 100644 --- a/docs/user_doc/conversation_compaction.md +++ b/docs/user_doc/conversation_compaction.md @@ -144,6 +144,37 @@ Compaction acquires a per-conversation lock to prevent concurrent requests on th Over very long conversations, multiple compaction summaries may accumulate. When the total size of cached summaries approaches the context window threshold, they are recursively folded into a single summary using a dedicated re-summarization prompt. This prevents summaries from themselves exceeding the context window. +### Token usage and quota + +Summarizing older turns takes an LLM call of its own, and so does folding the +summaries. The provider bills these calls, so they are counted: + +| Endpoint | Quota | Reported token counts | Metrics | +|---|---|---|---| +| `POST /v1/query` | Charged to the user. | `input_tokens` and `output_tokens` include the summarization calls. | Counted under `/v1/query`. | +| `POST /v1/streaming_query` | Charged to the user. | `input_tokens` and `output_tokens` of the `end` event include the summarization calls. | Counted under `/v1/streaming_query`. | +| `POST /v1/responses` | Charged to the user. | `usage` is what the provider reported for the response, so it covers the answer alone. `available_quotas` shows the quota left after both. | Counted under `/v1/responses`. | +| `POST /a2a` | Not charged. The endpoint does not use quotas. | None. | The summarization calls are counted under `/a2a`. The call that answers the request is not recorded in these metrics on this endpoint. | + +The metrics are `ls_llm_calls_total`, `ls_llm_token_sent_total` and +`ls_llm_token_received_total`, labelled with the endpoint. + +Only a request that makes a summarization call pays for it: one whose +estimated input crosses the threshold. The other requests are served from the +stored summary and cost what they cost without compaction. With the default +settings a summarization happens once every several turns; with a small +context window or a low `threshold_ratio` it can happen on every request. + +Each summarization call is charged as soon as it returned, before the model is +asked for the answer. It is therefore charged also when the request is blocked +by a shield, fails afterwards, or is interrupted by the client. The summary is +kept in these cases, so the next request on the conversation does not pay for +it again. A summarization call that itself fails, or is cut off before it +returned, reports no usage and is not charged. + +On `/v1/responses`, the difference between the quota consumed and the `usage` +of the response is the summarization. + ## When compaction is disabled When compaction is disabled (the default), requests that cause the conversation history to exceed the model's context window will fail with HTTP 413 (Prompt Too Long). Clients must manage conversation length themselves, for example by starting new conversations or deleting old ones. @@ -158,7 +189,11 @@ Compaction summarizes older turns, so fine-grained details from early in the con **Does compaction use extra tokens?** -Yes. The summarization step requires an additional LLM call, which consumes tokens. These tokens are counted against the user's quota. The trade-off is that the conversation can continue instead of failing with HTTP 413. +Yes. The summarization step requires an additional LLM call, which consumes tokens. On `/v1/query` and `/v1/streaming_query` these tokens are counted against the user's quota and are included in the `input_tokens` and `output_tokens` of the request that triggered the summarization. `/v1/responses` charges the quota and leaves `usage` as the provider reported it; `/a2a` has no quota. See [Token usage and quota](#token-usage-and-quota). The trade-off is that the conversation can continue instead of failing with HTTP 413. + +**Why did one request use many more tokens than the ones before it?** + +On `/v1/query` and `/v1/streaming_query`, that request triggered a summarization. Its `input_tokens` include the older turns that were sent to the LLM to be summarized, and its `output_tokens` include the summary. `context_status` is `"summarized"` on that request, and also on the requests after it that are served from the stored summary and do not pay for it again. **Can I use compaction with all LLM providers?** diff --git a/src/app/endpoints/a2a.py b/src/app/endpoints/a2a.py index 287c124c5..8126c285c 100644 --- a/src/app/endpoints/a2a.py +++ b/src/app/endpoints/a2a.py @@ -57,7 +57,7 @@ from authorization.middleware import authorize from client.ogx import AsyncOgxClientHolder from configuration import configuration -from constants import MEDIA_TYPE_EVENT_STREAM +from constants import ENDPOINT_PATH_A2A, MEDIA_TYPE_EVENT_STREAM from log import get_logger from models.api.requests import QueryRequest from models.common.responses.responses_api_params import ResponsesApiParams @@ -228,6 +228,7 @@ async def _compact_a2a_request( responses_params, configuration.inference, configuration.compaction, + endpoint_path=ENDPOINT_PATH_A2A, ) return compaction.params, PendingTurn.for_request( client, compaction.params, compaction.original_input diff --git a/src/app/endpoints/query.py b/src/app/endpoints/query.py index 1204dd19f..ff67d517b 100644 --- a/src/app/endpoints/query.py +++ b/src/app/endpoints/query.py @@ -1,6 +1,7 @@ """Handler for REST API call to provide answer to query using Response API.""" import datetime +from functools import partial from typing import Annotated from fastapi import APIRouter, Depends, Request @@ -49,6 +50,7 @@ from utils.pending_turn import pending_turn from utils.query import ( consume_query_tokens, + consume_summarization_tokens, prepare_input, store_query_results, validate_attachments_metadata, @@ -247,6 +249,8 @@ async def _handle_query_with_tracing( cache=configured_conversation_cache(), user_id=user_id, skip_user_id_check=_skip_userid_check, + endpoint_path=endpoint_path, + charge=partial(consume_summarization_tokens, user_id), ) responses_params = compaction.params @@ -314,6 +318,9 @@ async def _handle_query_with_tracing( model_id=responses_params.model, token_usage=turn_summary.token_usage, ) + # The turn reports the summarization calls made for it as part of its + # usage. They were charged when they were made. + reported_usage = turn_summary.token_usage + compaction.summarization_usage logger.info("Getting available quotas") available_quotas = get_available_quotas( @@ -346,8 +353,8 @@ async def _handle_query_with_tracing( root_span, { SpanAttributes.SESSION_ID: conversation_id, - SpanAttributes.LLM_USAGE_INPUT_TOKENS: turn_summary.token_usage.input_tokens, - SpanAttributes.LLM_USAGE_OUTPUT_TOKENS: turn_summary.token_usage.output_tokens, + SpanAttributes.LLM_USAGE_INPUT_TOKENS: reported_usage.input_tokens, + SpanAttributes.LLM_USAGE_OUTPUT_TOKENS: reported_usage.output_tokens, SpanAttributes.OUTPUT: turn_summary.llm_response, }, ) @@ -364,7 +371,7 @@ async def _handle_query_with_tracing( referenced_documents=turn_summary.referenced_documents, truncated=False, context_status=compaction.context_status, - input_tokens=turn_summary.token_usage.input_tokens, - output_tokens=turn_summary.token_usage.output_tokens, + input_tokens=reported_usage.input_tokens, + output_tokens=reported_usage.output_tokens, available_quotas=available_quotas, ) diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index 58335ea0a..122496427 100644 --- a/src/app/endpoints/responses.py +++ b/src/app/endpoints/responses.py @@ -6,6 +6,7 @@ import time from collections.abc import AsyncIterator, Sequence from datetime import UTC, datetime +from functools import partial from typing import Annotated, Any, Final, NoReturn, Optional, cast from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request @@ -85,6 +86,7 @@ from utils.prompts import get_system_prompt from utils.query import ( consume_query_tokens, + consume_summarization_tokens, extract_provider_and_model_from_model_id, handle_known_apistatus_errors, is_context_length_error, @@ -723,6 +725,11 @@ async def handle_responses_with_tracing( # pylint: disable=too-many-locals cache=configured_conversation_cache(), user_id=user_id, skip_user_id_check=skip_userid_check, + endpoint_path=endpoint_path, + # The summarization calls are charged to the user. The usage + # object of the response stays as OGX reported it, so it covers + # the answer alone. + charge=partial(consume_summarization_tokens, user_id), ) api_params = compaction.params if compaction.compacted: diff --git a/src/app/endpoints/streaming_query.py b/src/app/endpoints/streaming_query.py index 151e3643f..104d93b63 100644 --- a/src/app/endpoints/streaming_query.py +++ b/src/app/endpoints/streaming_query.py @@ -3,6 +3,7 @@ import asyncio import datetime from collections.abc import AsyncIterator +from functools import partial from typing import Annotated, Optional from fastapi import APIRouter, Depends, HTTPException, Request @@ -70,6 +71,7 @@ ) from utils.pending_turn import PendingTurn from utils.query import ( + consume_summarization_tokens, extract_provider_and_model_from_model_id, handle_known_apistatus_errors, is_context_length_error, @@ -94,6 +96,7 @@ stream_start_event, ) from utils.suid import get_suid, normalize_conversation_id +from utils.token_counter import TokenCounter from utils.types import Responses from utils.vector_search import build_rag_context @@ -437,6 +440,7 @@ async def generate_response_with_compaction( turn: Optional[PendingTurn] = None context_status: ContextStatus = "full" + summarization_usage = TokenCounter() try: async for item in apply_compaction( context.client, @@ -447,6 +451,8 @@ async def generate_response_with_compaction( cache=configured_conversation_cache(), user_id=context.user_id, skip_user_id_check=context.skip_userid_check, + endpoint_path=endpoint_path, + charge=partial(consume_summarization_tokens, context.user_id), ): if isinstance(item, CompactionStartedEvent): yield stream_compaction_event(context.conversation_id) @@ -456,6 +462,7 @@ async def generate_response_with_compaction( context.client, item.params, item.original_input ) context_status = item.context_status + summarization_usage = item.summarization_usage generator, turn_summary = await retrieve_agent_response_generator( responses_params=responses_params, @@ -512,6 +519,7 @@ async def generate_response_with_compaction( emit_start=False, turn=turn, context_status=context_status, + summarization_usage=summarization_usage, ): yield event finally: diff --git a/src/constants.py b/src/constants.py index 3c6c6355d..aca8ab72a 100644 --- a/src/constants.py +++ b/src/constants.py @@ -344,6 +344,7 @@ ENDPOINT_PATH_QUERY: Final[str] = "/v1/query" ENDPOINT_PATH_STREAMING_QUERY: Final[str] = "/v1/streaming_query" ENDPOINT_PATH_RESPONSES: Final[str] = "/v1/responses" +ENDPOINT_PATH_A2A: Final[str] = "/a2a" # Input size limits for API request validation # Maximum character length for the question field in /v1/infer requests (32 KiB) diff --git a/src/utils/agents/streaming.py b/src/utils/agents/streaming.py index f8930d726..2c5701775 100644 --- a/src/utils/agents/streaming.py +++ b/src/utils/agents/streaming.py @@ -89,6 +89,7 @@ register_interrupt_callback, ) from utils.streaming_sse import shield_violation_generator +from utils.token_counter import TokenCounter type AgentDispatchEvent = AgentStreamEvent | AgentRunResultEvent @@ -213,6 +214,7 @@ async def generate_agent_response( # pylint: disable=too-many-statements emit_start: bool = True, turn: Optional[PendingTurn] = None, context_status: ContextStatus = "full", + summarization_usage: Optional[TokenCounter] = None, ) -> AsyncIterator[str]: """Wrap an agent SSE generator with cleanup logic. @@ -236,6 +238,9 @@ async def generate_agent_response( # pylint: disable=too-many-statements 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. + summarization_usage: Usage of the summarization calls compaction made + for this request (LCORE-3910). They were charged when they were + made; here they are added to the usage the turn reports. Yields: SSE-formatted strings from the wrapped generator. @@ -354,6 +359,9 @@ async def generate_agent_response( # pylint: disable=too-many-statements model_id=responses_params.model, token_usage=turn_summary.token_usage, ) + # The turn reports the summarization calls made for it as part of its + # usage. They were charged when they were made. + reported_usage = turn_summary.token_usage + (summarization_usage or TokenCounter()) logger.info("Getting available quotas") available_quotas = get_available_quotas( quota_limiters=configuration.quota_limiters, @@ -362,8 +370,8 @@ async def generate_agent_response( # pylint: disable=too-many-statements end_payload = EndStreamPayload.create( referenced_documents=turn_summary.referenced_documents, context_status=context_status, - input_tokens=turn_summary.token_usage.input_tokens, - output_tokens=turn_summary.token_usage.output_tokens, + input_tokens=reported_usage.input_tokens, + output_tokens=reported_usage.output_tokens, available_quotas=available_quotas, ) yield serialize_event(end_payload, media_type) @@ -388,12 +396,8 @@ async def generate_agent_response( # pylint: disable=too-many-statements root_span, { SpanAttributes.SESSION_ID: context.conversation_id, - SpanAttributes.LLM_USAGE_INPUT_TOKENS: ( - turn_summary.token_usage.input_tokens - ), - SpanAttributes.LLM_USAGE_OUTPUT_TOKENS: ( - turn_summary.token_usage.output_tokens - ), + SpanAttributes.LLM_USAGE_INPUT_TOKENS: reported_usage.input_tokens, + SpanAttributes.LLM_USAGE_OUTPUT_TOKENS: reported_usage.output_tokens, SpanAttributes.OUTPUT: turn_summary.llm_response, }, ) diff --git a/src/utils/compaction.py b/src/utils/compaction.py index 47c2c9530..06d4c1cdb 100644 --- a/src/utils/compaction.py +++ b/src/utils/compaction.py @@ -27,14 +27,16 @@ to disentangle a tangle of side effects. """ +from collections.abc import Callable from datetime import UTC, datetime -from typing import Any +from typing import Any, Optional from ogx_client import AsyncOgxClient from log import get_logger from models.compaction import ConversationSummary from utils.query import normalize_vertex_ai_model_id +from utils.token_counter import TokenCounter from utils.token_estimator import ( estimate_conversation_tokens, estimate_tokens, @@ -44,6 +46,15 @@ logger = get_logger(__name__) +CallCounter = Callable[[str, TokenCounter], None] +"""Receives the model and the token usage of one LLM call made here. + +The provider bills the calls this module makes, so their usage must not be +lost (LCORE-3910). What is done with it (metrics, quota) is the caller's +business: this module only reports each call, as soon as its response +arrived and before the response is looked at. +""" + SUMMARIZATION_PROMPT = ( "Summarize this conversation history for an AI assistant that helps with\n" @@ -197,12 +208,32 @@ def _extract_response_text(response: Any) -> str: return "".join(parts) -async def summarize_chunk( +def reported_usage(response: Any) -> TokenCounter: + """Return the token usage the provider reported for one LLM call. + + Parameters: + response: The result of ``client.responses.create``. + + Returns: + The usage of that one call. The call is counted also when the + provider reported no usage for it; the token counts are 0 then. + """ + usage = getattr(response, "usage", None) + return TokenCounter( + input_tokens=int(getattr(usage, "input_tokens", 0) or 0), + output_tokens=int(getattr(usage, "output_tokens", 0) or 0), + llm_calls=1, + ) + + +async def summarize_chunk( # pylint: disable=too-many-arguments client: AsyncOgxClient, model: str, old_items: list[Any], summarized_through_turn: int, encoding_name: str, + *, + count_call: Optional[CallCounter] = None, ) -> ConversationSummary: """Summarize *old_items* via one LLM call and return a ConversationSummary. @@ -242,6 +273,9 @@ async def summarize_chunk( encoding_name: Tiktoken encoding name used to count tokens in the produced summary. Should match the encoding used to decide the compaction trigger. + count_call: Called with the model and the usage the provider + reported, once the LLM call returned. It is called also + when the call returned no text and this function raises. Returns: A populated ConversationSummary. @@ -275,6 +309,8 @@ async def summarize_chunk( stream=False, store=False, ) + if count_call is not None: + count_call(model, reported_usage(response)) summary_text = _extract_response_text(response).strip() if not summary_text: raise ValueError( @@ -314,6 +350,8 @@ async def recursively_resummarize( model: str, summaries: list[ConversationSummary], encoding_name: str, + *, + count_call: Optional[CallCounter] = None, ) -> ConversationSummary: """Collapse multiple ``ConversationSummary`` records into one. @@ -348,6 +386,9 @@ async def recursively_resummarize( short-circuit before invoking this function. encoding_name: Tiktoken encoding name used to count tokens in the produced fold. + count_call: Called with the model and the usage the provider + reported, once the LLM call returned. It is called also + when the call returned no text and this function raises. Returns: A single ConversationSummary representing the union of @@ -387,6 +428,8 @@ async def recursively_resummarize( stream=False, store=False, ) + if count_call is not None: + count_call(model, reported_usage(response)) folded_text = _extract_response_text(response).strip() if not folded_text: raise ValueError( diff --git a/src/utils/compaction_usage.py b/src/utils/compaction_usage.py new file mode 100644 index 000000000..3651f23a7 --- /dev/null +++ b/src/utils/compaction_usage.py @@ -0,0 +1,63 @@ +"""Accounting for the LLM calls conversation compaction makes (LCORE-3910). + +Compaction summarizes older turns with an LLM call of its own, and folds the +summaries with another. The provider bills both, so they are counted: in the +LLM metrics, in the quota of the user, and in the token counts the client is +told. :class:`SummarizationCalls` does the first two and adds the usage up for +the third. +""" + +from dataclasses import dataclass, field +from typing import Optional + +from metrics import recording +from utils.compaction import CallCounter +from utils.query import extract_provider_and_model_from_model_id +from utils.token_counter import TokenCounter + + +@dataclass +class SummarizationCalls: + """The LLM calls compaction makes for one request (LCORE-3910). + + The provider bills them. So each call is recorded in the LLM metrics under + the endpoint that triggered it, and handed to ``charge``. Both happen as + soon as the response of the call arrived, and do not depend on what + becomes of the request afterwards. + + Attributes: + endpoint_path: Path of the endpoint serving the request, the label of + the metrics. No metrics are recorded without it. + charge: Charges one call to whoever pays for the request. Nobody is + charged without it. + usage: The usage of the calls counted so far, added up. + """ + + endpoint_path: Optional[str] = None + charge: Optional[CallCounter] = None + usage: TokenCounter = field(default_factory=TokenCounter) + + def count(self, model: str, usage: TokenCounter) -> None: + """Record one call in the metrics, charge it, and add it to the total. + + Parameters: + model: Fully-qualified identifier of the model that was called. + usage: The token usage the provider reported for the call. + + Raises: + HTTPException: When charging the call fails. + """ + self.usage = self.usage + usage + if self.endpoint_path is not None: + provider_id, model_id = extract_provider_and_model_from_model_id(model) + if usage.input_tokens or usage.output_tokens: + recording.record_llm_token_usage( + provider_id, + model_id, + usage.input_tokens, + usage.output_tokens, + self.endpoint_path, + ) + recording.record_llm_call(provider_id, model_id, self.endpoint_path) + if self.charge is not None: + self.charge(model, usage) diff --git a/src/utils/conversation_compaction.py b/src/utils/conversation_compaction.py index 66e9e8c39..eb83865ac 100644 --- a/src/utils/conversation_compaction.py +++ b/src/utils/conversation_compaction.py @@ -44,7 +44,7 @@ import asyncio from collections.abc import AsyncIterator, Sequence from contextlib import asynccontextmanager -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Any, Optional, cast from fastapi import HTTPException @@ -62,14 +62,17 @@ from models.compaction import ConversationSummary from models.config import CompactionConfiguration, InferenceConfiguration from utils.compaction import ( + CallCounter, partition_conversation, recursively_resummarize, summarize_chunk, ) +from utils.compaction_usage import SummarizationCalls from utils.conversations import ( build_add_items_request, get_all_conversation_items, ) +from utils.token_counter import TokenCounter from utils.token_estimator import ( DEFAULT_ENCODING_NAME, estimate_conversation_tokens, @@ -177,11 +180,18 @@ class CompactionResult: 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. + summarization_usage: Token usage of the LLM calls compaction made for + this request: the summarization of older turns and the fold of + the summaries. Empty when the request made neither. The calls are + recorded and charged by the time the result is returned; the + usage is here for the endpoints that report token counts to the + client (LCORE-3910). """ params: ResponsesApiParams compacted: bool original_input: Optional[ResponseInput] = None + summarization_usage: TokenCounter = field(default_factory=TokenCounter) @property def context_status(self) -> ContextStatus: @@ -532,12 +542,15 @@ async def _maybe_persist_fold( # pylint: disable=too-many-arguments,too-many-po context_window: int, threshold_ratio: float, encoding_name: str, + count_call: Optional[CallCounter] = None, ) -> tuple[list[str], list[ConversationSummary]]: """Run the recursive fold (R3) if persisted summaries crossed the threshold. Requires a persisting cache (marker-only conversations keep additive chunks). Returns the (possibly updated) ``(summaries, cached_summaries)``; on a cache - write failure the unfolded values are returned unchanged. + write failure the unfolded values are returned unchanged. ``count_call`` is + told about the fold call once it returned, also when its result cannot be + stored afterwards (LCORE-3910). """ if cache is None or len(cached_summaries) < 2: return summaries, cached_summaries @@ -551,7 +564,7 @@ async def _maybe_persist_fold( # pylint: disable=too-many-arguments,too-many-po conversation_id, ) folded = await recursively_resummarize( - client, model, cached_summaries, encoding_name + client, model, cached_summaries, encoding_name, count_call=count_call ) try: cache.replace_summaries(user_id, conversation_id, folded, skip_user_id_check) @@ -566,14 +579,29 @@ def _compacted_result( summaries: list[str], recent_items: list[Any], original_input: ResponseInput, + usage: TokenCounter, ) -> CompactionResult: - """Build the CompactionResult for compacted mode (explicit input + omit_conversation).""" + """Build the CompactionResult for compacted mode (explicit input + omit_conversation). + + Parameters: + params: The prepared parameters of the request. + summaries: The summaries that stand for the older turns. + recent_items: The recent items sent verbatim. + original_input: The input as it arrived. + usage: Token usage of the LLM calls compaction made for the request. + + Returns: + The result for a request served in compacted mode. + """ explicit_input = _build_explicit_input(summaries, recent_items, original_input) compacted_params = params.model_copy( update={"input": explicit_input, "omit_conversation": True} ) return CompactionResult( - compacted_params, compacted=True, original_input=original_input + compacted_params, + compacted=True, + original_input=original_input, + summarization_usage=usage, ) @@ -587,6 +615,8 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit cache: Optional[Cache] = None, user_id: str = "", skip_user_id_check: bool = False, + endpoint_path: Optional[str] = None, + charge: Optional[CallCounter] = None, ) -> AsyncIterator[Any]: """Apply conversation compaction to a prepared request, yielding the result. @@ -613,9 +643,18 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit falls back to marker-only summaries with no folding. user_id: User identifier for cache reads/writes. skip_user_id_check: Whether to bypass the cache's user_id validation. + endpoint_path: Path of the endpoint serving the request. The LLM calls + compaction makes are recorded in the metrics under it; without it + they are not recorded. + charge: Called with the model and the token usage of each LLM call + compaction makes, as soon as the call returned. The endpoints + with a quota charge the user here (LCORE-3910). Yields: Zero or more CompactionStartedEvent, then exactly one CompactionResult. + + Raises: + HTTPException: When ``charge`` fails. """ if not compaction_config.enabled: # ``enabled: false`` is a full off-switch: the request passes through @@ -631,6 +670,7 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit conversation_id = params.conversation model = params.model original_input = params.input + calls = SummarizationCalls(endpoint_path, charge) async with _conversation_lock(conversation_id): items = await get_all_conversation_items(client, conversation_id) @@ -665,6 +705,7 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit old_items, summarized_through_turn=already + len(old_items), encoding_name=encoding_name, + count_call=calls.count, ) await _persist_new_summary_chunk( client, @@ -689,6 +730,7 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit context_window, compaction_config.threshold_ratio, encoding_name, + calls.count, ) if not summaries: @@ -699,7 +741,9 @@ async def apply_compaction( # pylint: disable=too-many-arguments,too-many-posit # Compacted mode: lightspeed owns the context. Build explicit input and # stop passing the conversation parameter to inference. - yield _compacted_result(params, summaries, recent_items, original_input) + yield _compacted_result( + params, summaries, recent_items, original_input, calls.usage + ) async def apply_compaction_blocking( # pylint: disable=too-many-arguments,too-many-positional-arguments @@ -711,12 +755,15 @@ async def apply_compaction_blocking( # pylint: disable=too-many-arguments,too-m cache: Optional[Cache] = None, user_id: str = "", skip_user_id_check: bool = False, + endpoint_path: Optional[str] = None, + charge: Optional[CallCounter] = None, ) -> CompactionResult: """Non-streaming wrapper around :func:`apply_compaction`. Drains the generator with event emission disabled and returns the final :class:`CompactionResult`. See :func:`apply_compaction` for the ``cache`` / - ``user_id`` / ``skip_user_id_check`` parameters. + ``user_id`` / ``skip_user_id_check`` / ``endpoint_path`` / ``charge`` + parameters. """ result: Optional[CompactionResult] = None async for item in apply_compaction( @@ -729,6 +776,8 @@ async def apply_compaction_blocking( # pylint: disable=too-many-arguments,too-m cache=cache, user_id=user_id, skip_user_id_check=skip_user_id_check, + endpoint_path=endpoint_path, + charge=charge, ): if isinstance(item, CompactionResult): result = item diff --git a/src/utils/query.py b/src/utils/query.py index 2cc911e66..7d5d4d595 100644 --- a/src/utils/query.py +++ b/src/utils/query.py @@ -310,6 +310,33 @@ def consume_query_tokens( raise HTTPException(**response.model_dump()) from e +def consume_summarization_tokens( + user_id: str, + model_id: str, + summarization_usage: TokenCounter, +) -> None: + """Consume the tokens of an LLM call compaction made for a request. + + Compaction summarizes older turns with LLM calls of its own, and the + provider bills them. Each is charged as soon as it returned and not with + the turn, so that it is charged also when the turn is blocked, fails or + is interrupted. Nothing is consumed for a call the provider reported no + usage for. + + Parameters: + user_id: The authenticated user ID. + model_id: The full model identifier in "provider/model" format. + summarization_usage: The token usage the provider reported. + + Raises: + HTTPException: On database errors during token consumption. + """ + if summarization_usage.input_tokens or summarization_usage.output_tokens: + consume_query_tokens( + user_id=user_id, model_id=model_id, token_usage=summarization_usage + ) + + def is_transcripts_enabled() -> bool: """Check if transcripts is enabled. diff --git a/src/utils/token_counter.py b/src/utils/token_counter.py index 94f0667d0..cbb78d893 100644 --- a/src/utils/token_counter.py +++ b/src/utils/token_counter.py @@ -23,6 +23,22 @@ class TokenCounter: input_tokens_counted: int = 0 llm_calls: int = 0 + def __add__(self, other: "TokenCounter") -> "TokenCounter": + """Return the usage of two sets of LLM calls taken together. + + Parameters: + other: The token counter to add. + + Returns: + A new counter holding the sums; neither operand is changed. + """ + return TokenCounter( + input_tokens=self.input_tokens + other.input_tokens, + output_tokens=self.output_tokens + other.output_tokens, + input_tokens_counted=self.input_tokens_counted + other.input_tokens_counted, + llm_calls=self.llm_calls + other.llm_calls, + ) + def __str__(self) -> str: """ Return a human-readable summary of the token usage stored in this TokenCounter. diff --git a/tests/integration/endpoints/test_compaction_token_usage.py b/tests/integration/endpoints/test_compaction_token_usage.py new file mode 100644 index 000000000..92a8e605e --- /dev/null +++ b/tests/integration/endpoints/test_compaction_token_usage.py @@ -0,0 +1,853 @@ +"""Integration tests for the token usage of the summarization calls (LCORE-3910). + +Compaction makes an LLM call of its own to summarize older turns. The provider +bills it, so it is charged to the user's quota and shows in the token counts +of the turn that triggered it. + +The summarization call is charged when it is made, so it is charged also when +the turn that triggered it is blocked, fails or is interrupted. + +The quota here is a real limiter on a SQLite database and the summarization +call is the real one, answered by the mocked OGX client. Every test reads what +the limiter holds after the request. +""" + +# pylint: disable=too-many-arguments +# pylint: disable=too-many-positional-arguments + +import asyncio +import json +from collections.abc import AsyncIterator, Callable, Generator +from dataclasses import dataclass +from pathlib import Path +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.models.open_ai_response_object_stream_response_completed import ( + OpenAIResponseObjectStreamResponseCompleted, +) +from prometheus_client import REGISTRY +from pydantic_ai.messages import 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 cache.sqlite_cache import SQLiteCache +from configuration import AppConfig +from constants import ( + ENDPOINT_PATH_A2A, + ENDPOINT_PATH_QUERY, + ENDPOINT_PATH_RESPONSES, + ENDPOINT_PATH_STREAMING_QUERY, +) +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.compaction import ConversationSummary +from models.config import ( + QuotaHandlersConfiguration, + QuotaLimiterConfiguration, + SQLiteDatabaseConfiguration, +) +from tests.integration.conftest import ( + TEST_MODEL_NAME, + TEST_PROVIDER, + InMemoryConversationStore, + make_openai_response_object, + set_query_agent_run, + set_streaming_query_agent_run, +) +from tests.integration.endpoints._compaction_helpers import ( + CONV_ID_LLAMA, + DEFAULT_MODEL_RESPONSE, + EXISTING_CONV_ID, + FAKE_AGENT_CARD, + TEST_MODEL, + assert_marker_count, + build_a2a_request, + create_existing_conversation, + enable_compaction, + marker, + mock_a2a_agent, + msg, +) +from utils.compaction import RECURSIVE_RESUMMARIZATION_PROMPT, SUMMARIZATION_PROMPT +from utils.stream_interrupts import CancelStreamResult, get_stream_interrupt_registry + +INITIAL_QUOTA = 100_000 +QUOTA = "UserQuotaLimiter" +NEW_QUERY = "What else can you help with?" + +# what the provider reports for the turn itself, for the summarization call and +# for the call that folds the summaries +TURN_INPUT, TURN_OUTPUT = 100, 50 +SUMMARY_INPUT, SUMMARY_OUTPUT = 640, 72 +FOLD_INPUT, FOLD_OUTPUT = 310, 45 +TURN = TURN_INPUT + TURN_OUTPUT +SUMMARY = SUMMARY_INPUT + SUMMARY_OUTPUT +FOLD = FOLD_INPUT + FOLD_OUTPUT + + +@pytest.fixture(name="quota") +def quota_fixture( + test_config: AppConfig, tmp_path: Path +) -> Generator[AppConfig, None, None]: + """Give every user a quota, kept by a real limiter on a SQLite database.""" + # pylint: disable=protected-access + assert test_config._configuration is not None + test_config._configuration.quota_handlers = QuotaHandlersConfiguration( + sqlite=SQLiteDatabaseConfiguration(db_path=str(tmp_path / "quota.db")), + limiters=[ + QuotaLimiterConfiguration( + type="user_limiter", + name="user quota", + initial_quota=INITIAL_QUOTA, + quota_increase=0, + period="1 day", + ) + ], + ) + test_config._quota_limiters = [] + yield test_config + for limiter in test_config._quota_limiters: + if limiter.connection is not None: + limiter.connection.close() + test_config._quota_limiters = [] + + +@pytest.fixture(name="conversation_cache") +def conversation_cache_fixture( + test_config: AppConfig, mocker: MockerFixture +) -> Generator[SQLiteCache, None, None]: + """Configure a conversation cache, where compaction keeps and folds its summaries.""" + test_config.conversation_cache_configuration.type = "sqlite" + cache = SQLiteCache(SQLiteDatabaseConfiguration(db_path=":memory:")) + cache.connect() + cache.initialize_cache() + mocker.patch.object( + type(test_config), + "conversation_cache", + new_callable=mocker.PropertyMock, + return_value=cache, + ) + yield cache + + +def _metrics_added(endpoint: str) -> Callable[[], tuple[float, ...]]: + """Start watching the LLM metrics of the test model under one endpoint. + + Parameters: + endpoint: The endpoint label of the metrics. + + Returns: + A function that reads what was added since: the tokens sent, the + tokens received and the number of LLM calls. + """ + labels = {"provider": TEST_PROVIDER, "model": TEST_MODEL_NAME, "endpoint": endpoint} + + def _read() -> tuple[float, ...]: + """Read the three metrics.""" + return tuple( + REGISTRY.get_sample_value(name, labels) or 0.0 + for name in ( + "ls_llm_token_sent_total", + "ls_llm_token_received_total", + "ls_llm_calls_total", + ) + ) + + before = _read() + return lambda: tuple(now - then for now, then in zip(_read(), before)) + + +def _long_conversation() -> list[OpenAIResponseMessage]: + """Two stored turns, long enough for the next request to summarize them.""" + return [ + msg("user", "question one " * 20), + msg("assistant", "answer one " * 20), + msg("user", "question two " * 20), + msg("assistant", "answer two " * 20), + ] + + +EARLIER_SUMMARIES = ("the first summary", "the second summary") + + +def _store_two_summaries(cache: SQLiteCache, user_id: str) -> None: + """Store two summaries that grow too large to keep once a third is added.""" + for number, text in enumerate(EARLIER_SUMMARIES, 1): + cache.store_summary( + user_id, + CONV_ID_LLAMA, + ConversationSummary( + summary_text=text, + summarized_through_turn=2 * number, + token_count=15, + created_at=f"2026-09-28T00:00:0{number}Z", + model_used=TEST_MODEL, + ), + ) + + +def _compacted_conversation() -> list[OpenAIResponseMessage]: + """A conversation that was summarized before and has nothing to summarize now.""" + return [marker("summary of the earlier turns")] + + +def _answer_like_ogx( + mock_ogx_client: AsyncMockType, + mocker: MockerFixture, + turn_error: Optional[Exception] = None, +) -> None: + """Let the mocked OGX answer the calls made through ``client.responses.create``. + + Those are the summarization call and the fold call, told apart by their + instructions, and on ``/v1/responses`` the turn itself. On ``/v1/query`` + and ``/v1/streaming_query`` the turn is answered by the mocked agent. With + *turn_error* the calls of compaction are answered and the turn fails. + """ + + async def _create(**kwargs: Any) -> Any: + """Answer one call.""" + if kwargs.get("instructions") == SUMMARIZATION_PROMPT: + return make_openai_response_object( + content="condensed earlier turns", + input_tokens=SUMMARY_INPUT, + output_tokens=SUMMARY_OUTPUT, + ) + if kwargs.get("instructions") == RECURSIVE_RESUMMARIZATION_PROMPT: + return make_openai_response_object( + content="all earlier turns in one summary", + input_tokens=FOLD_INPUT, + output_tokens=FOLD_OUTPUT, + ) + if turn_error is not None: + raise turn_error + response = make_openai_response_object( + content=DEFAULT_MODEL_RESPONSE, + input_tokens=TURN_INPUT, + output_tokens=TURN_OUTPUT, + ) + if kwargs.get("stream"): + return _one_chunk_stream(response) + return response + + mock_ogx_client.responses.create = mocker.AsyncMock(side_effect=_create) + + +async def _one_chunk_stream(response: Any) -> Any: + """Yield the terminal event of a streamed response.""" + yield OpenAIResponseObjectStreamResponseCompleted( + response=response, sequence_number=1, type="response.completed" + ) + + +def _events(chunks: list[str]) -> list[dict[str, Any]]: + """Parse the ``data:`` lines of a stream.""" + events = [] + for chunk in chunks: + for line in chunk.splitlines(): + if line.startswith("data: ") and line != "data: [DONE]": + events.append(json.loads(line[len("data: ") :])) + return events + + +async def _drain(response: Any) -> list[dict[str, Any]]: + """Read a streaming response to its end and return its events.""" + assert isinstance(response, StreamingResponse) + return _events([str(chunk) async for chunk in response.body_iterator]) + + +def _blocked() -> ShieldModerationBlocked: + """Return the verdict of a shield that blocked the request.""" + return ShieldModerationBlocked( + message="Content blocked by safety shield", moderation_id="modr_blocked_1" + ) + + +def _stream_that_stalls() -> Any: + """Build an agent stream that sends one token and then waits to be interrupted.""" + + async def _agent_events() -> AsyncIterator[Any]: + """Send one token and wait.""" + yield PartStartEvent(index=0, part=TextPart(content="Ansible is")) + await asyncio.Event().wait() + + class _RunStreamCtx: + """Async context manager matching ``agent.run_stream_events``.""" + + async def __aenter__(self) -> AsyncIterator[Any]: + return _agent_events() + + async def __aexit__(self, *_args: object) -> None: + return None + + return _RunStreamCtx() + + +async def _interrupt_stream(response: Any, user_id: str) -> list[dict[str, Any]]: + """Interrupt a stream after its first token and return its events.""" + assert isinstance(response, StreamingResponse) + chunks: list[str] = [] + first_token = asyncio.Event() + + async def _consume() -> None: + """Read the stream until it ends.""" + 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) + (start,) = [e for e in _events(chunks) if e.get("event") == "start"] + result = get_stream_interrupt_registry().cancel_stream( + start["data"]["request_id"], user_id + ) + 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 _events(chunks) + + +@dataclass +class Scene: + """A user with a quota, a stored conversation, and what a request on it needs. + + Attributes: + config: The configuration, with the quota limiter. + store: The conversation store behind the mocked OGX client. + ogx: The mocked OGX client, which answers the calls of compaction. + request: The request object the handlers get. + auth: The authenticated user. + """ + + config: AppConfig + store: InMemoryConversationStore + ogx: AsyncMockType + request: Request + auth: AuthTuple + + async def holds( + self, items: list[OpenAIResponseMessage], context_window: int = 200 + ) -> None: + """Enable compaction and store the conversation the request continues.""" + enable_compaction(self.config, context_window=context_window) + await self.store.create(conversation_id=CONV_ID_LLAMA, items=items) + + def quota_left(self) -> int: + """Read the quota the user has left, from the limiter.""" + (limiter,) = self.config.quota_limiters + return limiter.available_quota(self.auth[0]) + + def quotas(self, consumed: int) -> dict[str, int]: + """Return what the client is told after *consumed* tokens were charged.""" + return {QUOTA: INITIAL_QUOTA - consumed} + + def summarized_once(self) -> None: + """Check that the request stored one summary marker in the conversation.""" + assert_marker_count(self.store, CONV_ID_LLAMA, 1) + + +@pytest.fixture(name="scene") +def scene_fixture( + quota: AppConfig, + mock_ogx_client: AsyncMockType, + mock_conversation_store: InMemoryConversationStore, + test_request: Request, + test_auth: AuthTuple, + patch_db_session: Session, + mocker: MockerFixture, +) -> Scene: + """Set the scene: the user, the quota, and an OGX that answers compaction.""" + create_existing_conversation(patch_db_session, test_auth[0]) + _answer_like_ogx(mock_ogx_client, mocker) + return Scene( + quota, mock_conversation_store, mock_ogx_client, test_request, test_auth + ) + + +# ========================================== +# /v1/query +# ========================================== + + +@pytest.fixture(name="query_agent") +def query_agent_fixture( + mock_query_agent: AsyncMockType, mocker: MockerFixture +) -> AsyncMockType: + """Let the agent answer the turn, with the usage the provider reports for it.""" + set_query_agent_run( + mock_query_agent, mocker, input_tokens=TURN_INPUT, output_tokens=TURN_OUTPUT + ) + return mock_query_agent + + +async def _query(scene: Scene) -> Any: + """Send the new query on the stored conversation.""" + return await query_endpoint_handler( + request=scene.request, + query_request=QueryRequest(query=NEW_QUERY, conversation_id=EXISTING_CONV_ID), + auth=scene.auth, + mcp_headers={}, + ) + + +@pytest.mark.usefixtures("query_agent") +class TestQueryTokenUsage: + """Token counts and quota on /v1/query.""" + + @pytest.mark.asyncio + async def test_turn_that_summarizes_pays_for_the_summarization( + self, scene: Scene + ) -> None: + """The counts of the turn include the summarization call, and so does the quota.""" + await scene.holds(_long_conversation()) + added = _metrics_added(ENDPOINT_PATH_QUERY) + + response = await _query(scene) + + scene.summarized_once() + assert response.context_status == "summarized" + assert response.input_tokens == TURN_INPUT + SUMMARY_INPUT + assert response.output_tokens == TURN_OUTPUT + SUMMARY_OUTPUT + assert response.available_quotas == scene.quotas(TURN + SUMMARY) + assert scene.quota_left() == INITIAL_QUOTA - TURN - SUMMARY + assert added() == ( + TURN_INPUT + SUMMARY_INPUT, + TURN_OUTPUT + SUMMARY_OUTPUT, + 2, + ) + + @pytest.mark.asyncio + async def test_turn_that_summarizes_and_folds_pays_for_both_calls( + self, scene: Scene, conversation_cache: SQLiteCache + ) -> None: + """The summaries grow too large with the new one, so they are folded as well.""" + _store_two_summaries(conversation_cache, scene.auth[0]) + await scene.holds( + [marker(text) for text in EARLIER_SUMMARIES] + _long_conversation() + ) + added = _metrics_added(ENDPOINT_PATH_QUERY) + + response = await _query(scene) + + (folded,) = conversation_cache.get_summaries( + scene.auth[0], CONV_ID_LLAMA, False + ) + assert folded.summary_text == "all earlier turns in one summary" + assert response.input_tokens == TURN_INPUT + SUMMARY_INPUT + FOLD_INPUT + assert response.output_tokens == TURN_OUTPUT + SUMMARY_OUTPUT + FOLD_OUTPUT + assert scene.quota_left() == INITIAL_QUOTA - TURN - SUMMARY - FOLD + assert added()[2] == 3 + + @pytest.mark.asyncio + async def test_turn_served_from_a_summary_pays_for_itself_only( + self, scene: Scene + ) -> None: + """A compacted conversation costs nothing extra while nothing is summarized.""" + await scene.holds(_compacted_conversation(), context_window=100_000) + + response = await _query(scene) + + scene.ogx.responses.create.assert_not_awaited() + assert response.context_status == "summarized" + assert (response.input_tokens, response.output_tokens) == ( + TURN_INPUT, + TURN_OUTPUT, + ) + assert scene.quota_left() == INITIAL_QUOTA - TURN + + @pytest.mark.asyncio + async def test_turn_on_a_short_conversation_pays_for_itself_only( + self, scene: Scene + ) -> None: + """A conversation that is not compacted costs what it cost before.""" + await scene.holds( + [msg("user", "hi"), msg("assistant", "hello")], context_window=100_000 + ) + + response = await _query(scene) + + scene.ogx.responses.create.assert_not_awaited() + assert response.context_status == "full" + assert (response.input_tokens, response.output_tokens) == ( + TURN_INPUT, + TURN_OUTPUT, + ) + assert scene.quota_left() == INITIAL_QUOTA - TURN + + @pytest.mark.asyncio + async def test_blocked_turn_that_summarized_pays_for_the_summarization( + self, scene: Scene, query_agent: AsyncMockType, mocker: MockerFixture + ) -> None: + """Compaction runs before the model call, so its call is made and billed.""" + await scene.holds(_long_conversation()) + mocker.patch( + "app.endpoints.query.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + + response = await _query(scene) + + query_agent.run.assert_not_awaited() + assert (response.input_tokens, response.output_tokens) == ( + SUMMARY_INPUT, + SUMMARY_OUTPUT, + ) + assert response.available_quotas == scene.quotas(SUMMARY) + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + + @pytest.mark.asyncio + async def test_failed_turn_that_summarized_pays_for_the_summarization( + self, scene: Scene, query_agent: AsyncMockType + ) -> None: + """The summary is made and kept before the model call that then fails.""" + await scene.holds(_long_conversation()) + query_agent.run.side_effect = RuntimeError("the model is down") + + with pytest.raises(HTTPException): + await _query(scene) + + scene.summarized_once() + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + + +# ========================================== +# /v1/streaming_query +# ========================================== + + +@pytest.fixture(name="streaming_agent") +def streaming_agent_fixture( + mock_streaming_query_agent: AsyncMockType, mocker: MockerFixture +) -> AsyncMockType: + """Let the agent answer the turn, with the usage the provider reports for it.""" + set_streaming_query_agent_run( + mock_streaming_query_agent, + mocker, + input_tokens=TURN_INPUT, + output_tokens=TURN_OUTPUT, + ) + return mock_streaming_query_agent + + +async def _streaming_query(scene: Scene) -> Any: + """Send the new query on the stored conversation, to be answered in a stream.""" + return await streaming_query_endpoint_handler( + request=scene.request, + query_request=QueryRequest(query=NEW_QUERY, conversation_id=EXISTING_CONV_ID), + auth=scene.auth, + mcp_headers={}, + ) + + +def _only(events: list[dict[str, Any]], name: str) -> dict[str, Any]: + """Return the one event of the given name in a stream.""" + (event,) = [event for event in events if event.get("event") == name] + return event + + +class TestStreamingQueryTokenUsage: + """Token counts and quota on /v1/streaming_query.""" + + @pytest.mark.asyncio + async def test_turn_that_summarizes_pays_for_the_summarization( + self, scene: Scene, streaming_agent: AsyncMockType + ) -> None: + """The end event reports the turn with the summarization call included.""" + _ = streaming_agent + await scene.holds(_long_conversation()) + added = _metrics_added(ENDPOINT_PATH_STREAMING_QUERY) + + end = _only(await _drain(await _streaming_query(scene)), "end") + + assert end["data"]["context_status"] == "summarized" + assert end["data"]["input_tokens"] == TURN_INPUT + SUMMARY_INPUT + assert end["data"]["output_tokens"] == TURN_OUTPUT + SUMMARY_OUTPUT + assert end["available_quotas"] == scene.quotas(TURN + SUMMARY) + assert scene.quota_left() == INITIAL_QUOTA - TURN - SUMMARY + # The number of calls is left out: this endpoint counts the call of + # the turn twice, with or without compaction. + assert added()[:2] == ( + TURN_INPUT + SUMMARY_INPUT, + TURN_OUTPUT + SUMMARY_OUTPUT, + ) + + @pytest.mark.asyncio + async def test_blocked_stream_that_summarized_pays_for_the_summarization( + self, scene: Scene, streaming_agent: AsyncMockType, mocker: MockerFixture + ) -> None: + """Compaction runs before the model call, so its call is made and billed.""" + await scene.holds(_long_conversation()) + mocker.patch( + "app.endpoints.streaming_query.run_shield_moderation", + new=mocker.AsyncMock(return_value=_blocked()), + ) + + end = _only(await _drain(await _streaming_query(scene)), "end") + + streaming_agent.run_stream_events.assert_not_called() + assert end["data"]["input_tokens"] == SUMMARY_INPUT + assert end["data"]["output_tokens"] == SUMMARY_OUTPUT + assert end["available_quotas"] == scene.quotas(SUMMARY) + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + + @pytest.mark.asyncio + async def test_failed_stream_that_summarized_pays_for_the_summarization( + self, scene: Scene, streaming_agent: AsyncMockType + ) -> None: + """A stream that ends in an error reports no counts, and the call is charged.""" + await scene.holds(_long_conversation()) + streaming_agent.run_stream_events.side_effect = RuntimeError( + "the model is down" + ) + + events = await _drain(await _streaming_query(scene)) + + assert _only(events, "error") + assert not [event for event in events if event.get("event") == "end"] + scene.summarized_once() + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + + @pytest.mark.asyncio + async def test_interrupted_stream_that_summarized_pays_for_the_summarization( + self, scene: Scene, streaming_agent: AsyncMockType + ) -> None: + """An interrupted answer is not charged; the summary made before it is.""" + await scene.holds(_long_conversation()) + streaming_agent.run_stream_events.return_value = _stream_that_stalls() + + events = await _interrupt_stream(await _streaming_query(scene), scene.auth[0]) + + assert _only(events, "interrupted") + assert not [event for event in events if event.get("event") == "end"] + scene.summarized_once() + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + + +# ========================================== +# /v1/responses +# ========================================== + + +@pytest.fixture(name="responses_handlers") +def responses_handlers_fixture(mocker: MockerFixture) -> None: + """Let the real /v1/responses handlers run against the mocked OGX client.""" + 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 _respond(scene: Scene, stream: bool) -> Any: + """Send the new input on the stored conversation.""" + return await responses_endpoint_handler( + request=scene.request, + responses_request=ResponsesRequest( + input=NEW_QUERY, + model=TEST_MODEL, + conversation=EXISTING_CONV_ID, + stream=stream, + store=True, + generate_topic_summary=False, + ), + auth=scene.auth, + mcp_headers={}, + ) + + +def _block_responses(mocker: MockerFixture) -> None: + """Let the shield block the request.""" + mocker.patch( + "app.endpoints.responses.run_shield_moderation_v2", return_value=_blocked() + ) + + +@pytest.mark.usefixtures("responses_handlers") +class TestResponsesTokenUsage: + """Usage and quota on /v1/responses.""" + + @pytest.mark.asyncio + async def test_request_that_summarizes_pays_for_the_summarization( + self, scene: Scene + ) -> None: + """The quota is charged for both calls; the usage stays as the provider reported it. + + The endpoint passes the usage object of the response through, so that + it stays a drop-in for clients of the OpenAI Responses API. + """ + await scene.holds(_long_conversation()) + added = _metrics_added(ENDPOINT_PATH_RESPONSES) + + response = await _respond(scene, stream=False) + + scene.summarized_once() + assert response.usage is not None + assert response.usage.input_tokens == TURN_INPUT + assert response.usage.output_tokens == TURN_OUTPUT + assert response.available_quotas == scene.quotas(TURN + SUMMARY) + assert scene.quota_left() == INITIAL_QUOTA - TURN - SUMMARY + # At least, and not exactly: without a stream this endpoint records the + # tokens of the answer twice, with or without compaction. What the + # summarization adds is pinned by the failed request below. + assert added()[0] >= TURN_INPUT + SUMMARY_INPUT + + @pytest.mark.asyncio + async def test_stream_that_summarizes_pays_for_the_summarization( + self, scene: Scene + ) -> None: + """The terminal event reports the quota left after both calls.""" + await scene.holds(_long_conversation()) + added = _metrics_added(ENDPOINT_PATH_RESPONSES) + + events = await _drain(await _respond(scene, stream=True)) + + (completed,) = [e for e in events if e.get("type") == "response.completed"] + assert completed["response"]["usage"]["input_tokens"] == TURN_INPUT + assert completed["response"]["usage"]["output_tokens"] == TURN_OUTPUT + assert completed["response"]["available_quotas"] == scene.quotas(TURN + SUMMARY) + assert scene.quota_left() == INITIAL_QUOTA - TURN - SUMMARY + assert added() == ( + TURN_INPUT + SUMMARY_INPUT, + TURN_OUTPUT + SUMMARY_OUTPUT, + 2, + ) + + @pytest.mark.asyncio + async def test_failed_request_that_summarized_pays_for_the_summarization( + self, scene: Scene, mocker: MockerFixture + ) -> None: + """The summary is made and kept before the model call that then fails.""" + await scene.holds(_long_conversation()) + _answer_like_ogx( + scene.ogx, mocker, turn_error=RuntimeError("the model is down") + ) + added = _metrics_added(ENDPOINT_PATH_RESPONSES) + + with pytest.raises(RuntimeError, match="the model is down"): + await _respond(scene, stream=False) + + scene.summarized_once() + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + assert added() == (SUMMARY_INPUT, SUMMARY_OUTPUT, 1) + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + async def test_blocked_request_that_summarized_pays_for_the_summarization( + self, stream: bool, scene: Scene, mocker: MockerFixture + ) -> None: + """Compaction runs before the model call, so its call is made and billed.""" + await scene.holds(_long_conversation()) + _block_responses(mocker) + + response = await _respond(scene, stream) + if stream: + await _drain(response) + + scene.summarized_once() + assert scene.quota_left() == INITIAL_QUOTA - SUMMARY + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream", [False, True], ids=["blocking", "streaming"]) + async def test_blocked_request_that_did_not_summarize_pays_nothing( + self, stream: bool, scene: Scene, mocker: MockerFixture + ) -> None: + """A blocked request made no LLM call at all, as before.""" + await scene.holds(_compacted_conversation(), context_window=100_000) + _block_responses(mocker) + + response = await _respond(scene, stream) + if stream: + await _drain(response) + + assert scene.quota_left() == INITIAL_QUOTA + + +# ========================================== +# /a2a +# ========================================== + + +async def _send_a2a_message(scene: Scene, mocker: MockerFixture) -> None: + """Send the new query to the A2A endpoint, on the stored conversation.""" + 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) + mocker.patch("app.endpoints.a2a.build_agent", return_value=mock_a2a_agent(mocker)) + await handle_a2a_jsonrpc_post( + request=build_a2a_request(NEW_QUERY), auth=scene.auth, mcp_headers={} + ) + + +class TestA2ATokenUsage: + """The A2A endpoint has no quota; the summarization call shows in the metrics.""" + + @pytest.mark.asyncio + async def test_summarization_is_recorded_and_charged_to_nobody( + self, scene: Scene, mocker: MockerFixture + ) -> None: + """The tokens of the summarization call are counted under the endpoint.""" + await scene.holds(_long_conversation()) + # Reading the quota creates the user's row; a charge would change it. + assert scene.quota_left() == INITIAL_QUOTA + added = _metrics_added(ENDPOINT_PATH_A2A) + + await _send_a2a_message(scene, mocker) + + scene.summarized_once() + assert added() == (SUMMARY_INPUT, SUMMARY_OUTPUT, 1) + assert scene.quota_left() == INITIAL_QUOTA + + @pytest.mark.asyncio + async def test_message_served_from_a_summary_records_nothing( + self, scene: Scene, mocker: MockerFixture + ) -> None: + """Without a summarization call there is nothing to count.""" + await scene.holds(_compacted_conversation(), context_window=100_000) + added = _metrics_added(ENDPOINT_PATH_A2A) + + await _send_a2a_message(scene, mocker) + + scene.ogx.responses.create.assert_not_awaited() + assert added() == (0, 0, 0) diff --git a/tests/unit/utils/test_query.py b/tests/unit/utils/test_query.py index c311f0620..176370531 100644 --- a/tests/unit/utils/test_query.py +++ b/tests/unit/utils/test_query.py @@ -31,6 +31,7 @@ from utils.query import ( build_multimodal_input, consume_query_tokens, + consume_summarization_tokens, extract_provider_and_model_from_model_id, handle_known_apistatus_errors, is_transcripts_enabled, @@ -615,6 +616,65 @@ def test_consume_tokens_database_error(self, mocker: MockerFixture) -> None: assert exc_info.value.status_code == 500 +class TestConsumeSummarizationTokens: + """Tests for consume_summarization_tokens function.""" + + def test_summarization_calls_are_charged(self, mocker: MockerFixture) -> None: + """The tokens of the summarization calls are consumed for the user.""" + mock_consume = mocker.patch("utils.query.consume_tokens") + + consume_summarization_tokens( + "user1", + "provider1/model1", + TokenCounter(input_tokens=640, output_tokens=72, llm_calls=1), + ) + + mock_consume.assert_called_once() + charged = mock_consume.call_args.kwargs + assert charged["user_id"] == "user1" + assert (charged["input_tokens"], charged["output_tokens"]) == (640, 72) + assert (charged["provider_id"], charged["model_id"]) == ( + "provider1", + "model1", + ) + + def test_request_without_summarization_consumes_nothing( + self, mocker: MockerFixture + ) -> None: + """A request that made no summarization call does not touch the quota.""" + mock_consume = mocker.patch("utils.query.consume_tokens") + + consume_summarization_tokens("user1", "provider1/model1", TokenCounter()) + + mock_consume.assert_not_called() + + def test_call_without_reported_usage_consumes_nothing( + self, mocker: MockerFixture + ) -> None: + """A call the provider reported no usage for leaves nothing to consume.""" + mock_consume = mocker.patch("utils.query.consume_tokens") + + consume_summarization_tokens( + "user1", "provider1/model1", TokenCounter(llm_calls=1) + ) + + mock_consume.assert_not_called() + + def test_database_error(self, mocker: MockerFixture) -> None: + """A database error is reported the way it is for the turn itself.""" + mocker.patch( + "utils.query.consume_tokens", side_effect=sqlite3.Error("DB error") + ) + + with pytest.raises(HTTPException) as exc_info: + consume_summarization_tokens( + "user1", + "provider1/model1", + TokenCounter(input_tokens=640, output_tokens=72, llm_calls=1), + ) + assert exc_info.value.status_code == 500 + + class TestStoreQueryResults: """Tests for store_query_results function.""" diff --git a/tests/unit/utils/test_summarization_usage.py b/tests/unit/utils/test_summarization_usage.py new file mode 100644 index 000000000..64539923b --- /dev/null +++ b/tests/unit/utils/test_summarization_usage.py @@ -0,0 +1,561 @@ +"""Unit tests for the token usage of the summarization calls (LCORE-3910). + +Compaction makes LLM calls of its own: one to summarize older turns, and one +to fold the summaries when they grow too large. The provider bills them, so +their usage has to reach the metrics, the quota and the token counts of the +turn. Each call is counted as soon as its response arrived, so the count does +not depend on what becomes of the request afterwards. +""" + +from typing import Any, Optional, cast + +import pytest +from ogx_api.openai_responses import OpenAIResponseMessage +from pytest_mock import MockerFixture + +from cache.cache_error import CacheError +from models.common.responses.responses_api_params import ResponsesApiParams +from models.compaction import ConversationSummary +from models.config import CompactionConfiguration, InferenceConfiguration +from utils import conversation_compaction as cc +from utils.compaction import ( + RECURSIVE_RESUMMARIZATION_PROMPT, + SUMMARIZATION_PROMPT, + recursively_resummarize, + reported_usage, + summarize_chunk, +) +from utils.token_counter import TokenCounter +from utils.token_estimator import DEFAULT_ENCODING_NAME + +MODEL = "openai/gpt-4o-mini" +CONV = "conv_abc123" +ENDPOINT = "/v1/query" + +SUMMARIZATION = TokenCounter(input_tokens=640, output_tokens=72, llm_calls=1) +FOLD = TokenCounter(input_tokens=310, output_tokens=45, llm_calls=1) + + +def _response( + mocker: MockerFixture, + text: str, + input_tokens: Optional[int] = None, + output_tokens: Optional[int] = None, +) -> Any: + """Build the result of ``client.responses.create``, with or without usage.""" + usage = ( + None + if input_tokens is None + else mocker.Mock(input_tokens=input_tokens, output_tokens=output_tokens) + ) + content = [mocker.Mock(text=text)] if text else [] + return mocker.Mock(output=[mocker.Mock(content=content)], usage=usage) + + +def _summary(text: str, tokens: int) -> ConversationSummary: + """Build a summary chunk of the given size.""" + return ConversationSummary( + summary_text=text, + summarized_through_turn=2, + token_count=tokens, + created_at="2026-09-28T00:00:00Z", + model_used=MODEL, + ) + + +def _msg(role: str, text: str) -> OpenAIResponseMessage: + """Build a typed OGX message item.""" + return OpenAIResponseMessage(role=cast("Any", role), content=text) + + +def _long_conversation() -> list[OpenAIResponseMessage]: + """One stored turn, long enough for the next request to summarize it.""" + return [_msg("user", "q1 " * 50), _msg("assistant", "a1 " * 50)] + + +def _params() -> ResponsesApiParams: + """Build the parameters of a request on an existing conversation.""" + return ResponsesApiParams( + input="follow-up", + model=MODEL, + conversation=CONV, + instructions="system prompt", + store=True, + stream=False, + ) + + +# --- adding up usage --- + + +def test_token_counters_add_up() -> None: + """The sum counts the tokens and the calls of both, and changes neither.""" + turn = TokenCounter( + input_tokens=100, output_tokens=50, input_tokens_counted=90, llm_calls=1 + ) + summarization = TokenCounter( + input_tokens=700, output_tokens=80, input_tokens_counted=600, llm_calls=2 + ) + + total = turn + summarization + + assert total == TokenCounter( + input_tokens=800, output_tokens=130, input_tokens_counted=690, llm_calls=3 + ) + assert turn == TokenCounter( + input_tokens=100, output_tokens=50, input_tokens_counted=90, llm_calls=1 + ) + assert summarization == TokenCounter( + input_tokens=700, output_tokens=80, input_tokens_counted=600, llm_calls=2 + ) + + +def test_adding_nothing_changes_nothing() -> None: + """A request that did not summarize adds an empty counter.""" + turn = TokenCounter(input_tokens=100, output_tokens=50, llm_calls=1) + + assert turn + TokenCounter() == turn + + +# --- what the provider reported --- + + +def test_reported_usage(mocker: MockerFixture) -> None: + """The usage of a call is what the provider reported for it.""" + assert reported_usage(_response(mocker, "text", 640, 72)) == SUMMARIZATION + + +@pytest.mark.parametrize( + "usage", + [None, {"input_tokens": None, "output_tokens": None}, {}], + ids=["no usage", "empty counts", "no counts"], +) +def test_call_without_reported_usage_is_still_a_call( + mocker: MockerFixture, usage: Optional[dict[str, Any]] +) -> None: + """A provider that reports no usage leaves the counts at zero; the call counts.""" + response = mocker.Mock( + usage=None if usage is None else mocker.Mock(spec_set=list(usage), **usage) + ) + + assert reported_usage(response) == TokenCounter(llm_calls=1) + + +# --- what the two calls report --- + + +@pytest.mark.asyncio +async def test_summarize_chunk_reports_its_call(mocker: MockerFixture) -> None: + """The summarization call is reported with the usage the provider gave.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _response(mocker, "a summary", 640, 72) + count_call = mocker.Mock() + + await summarize_chunk( + client=client, + model=MODEL, + old_items=[_msg("user", "hi")], + summarized_through_turn=1, + encoding_name=DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + + count_call.assert_called_once_with(MODEL, SUMMARIZATION) + + +@pytest.mark.asyncio +async def test_summarize_chunk_reports_a_call_that_returned_no_text( + mocker: MockerFixture, +) -> None: + """The call was made and billed, even when there is no summary in the answer.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _response(mocker, "", 640, 72) + count_call = mocker.Mock() + + with pytest.raises(ValueError, match="no extractable text"): + await summarize_chunk( + client=client, + model=MODEL, + old_items=[_msg("user", "hi")], + summarized_through_turn=1, + encoding_name=DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + + count_call.assert_called_once_with(MODEL, SUMMARIZATION) + + +@pytest.mark.asyncio +async def test_summarize_chunk_reports_nothing_for_a_call_that_failed( + mocker: MockerFixture, +) -> None: + """A call that raised has no response, so there is no usage to report.""" + client = mocker.AsyncMock() + client.responses.create.side_effect = RuntimeError("the model is down") + count_call = mocker.Mock() + + with pytest.raises(RuntimeError): + await summarize_chunk( + client=client, + model=MODEL, + old_items=[_msg("user", "hi")], + summarized_through_turn=1, + encoding_name=DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + + count_call.assert_not_called() + + +@pytest.mark.asyncio +async def test_fold_reports_its_call(mocker: MockerFixture) -> None: + """The fold call is reported with the usage the provider gave.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _response(mocker, "folded", 310, 45) + count_call = mocker.Mock() + + await recursively_resummarize( + client, + MODEL, + [_summary("one", 5), _summary("two", 5)], + DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + + count_call.assert_called_once_with(MODEL, FOLD) + + +@pytest.mark.asyncio +async def test_fold_reports_a_call_that_returned_no_text( + mocker: MockerFixture, +) -> None: + """The fold call was made and billed, even when its answer is empty.""" + client = mocker.AsyncMock() + client.responses.create.return_value = _response(mocker, "", 310, 45) + count_call = mocker.Mock() + + with pytest.raises(ValueError, match="no extractable"): + await recursively_resummarize( + client, + MODEL, + [_summary("one", 5), _summary("two", 5)], + DEFAULT_ENCODING_NAME, + count_call=count_call, + ) + + count_call.assert_called_once_with(MODEL, FOLD) + + +# --- what a request records, charges and is told --- + + +@pytest.fixture(name="recorded") +def recorded_fixture(mocker: MockerFixture) -> Any: + """Replace the metric recorders and return them.""" + return mocker.patch("utils.compaction_usage.recording") + + +def _answering(mocker: MockerFixture, answers: dict[str, Any]) -> Any: + """Build an OGX client that answers each LLM call by its instructions. + + Parameters: + mocker: The pytest-mock fixture. + answers: The response, or the error to raise, per ``instructions``. + + Returns: + The mocked client. + """ + + async def _create(**kwargs: Any) -> Any: + """Answer one call.""" + answer = answers[kwargs["instructions"]] + if isinstance(answer, Exception): + raise answer + return answer + + client = mocker.AsyncMock() + client.responses.create = mocker.AsyncMock(side_effect=_create) + return client + + +async def _compact( # pylint: disable=too-many-arguments + mocker: MockerFixture, + client: Any, + items: list[Any], + *, + charge: Any = None, + cache: Any = None, + write_marker: Any = None, + endpoint_path: Optional[str] = ENDPOINT, + threshold_ratio: float = 0.1, +) -> cc.CompactionResult: + """Apply compaction to a follow-up request on the stored items.""" + mocker.patch.object( + cc, "get_all_conversation_items", mocker.AsyncMock(return_value=items) + ) + mocker.patch.object(cc, "_write_summary_marker", write_marker or mocker.AsyncMock()) + return await cc.apply_compaction_blocking( + client=client, + params=_params(), + inference_config=InferenceConfiguration(context_windows={MODEL: 50}), + compaction_config=CompactionConfiguration( + enabled=True, + threshold_ratio=threshold_ratio, + token_floor=0, + buffer_turns=0, + buffer_max_ratio=0.3, + ), + cache=cache, + user_id="u1", + endpoint_path=endpoint_path, + charge=charge, + ) + + +def _summarizing(mocker: MockerFixture) -> Any: + """Build a client whose summarization call succeeds.""" + return _answering( + mocker, {SUMMARIZATION_PROMPT: _response(mocker, "condensed", 640, 72)} + ) + + +@pytest.mark.asyncio +async def test_request_that_summarizes(mocker: MockerFixture, recorded: Any) -> None: + """The call is recorded in the metrics, charged, and the request is told.""" + charge = mocker.Mock() + + result = await _compact( + mocker, _summarizing(mocker), _long_conversation(), charge=charge + ) + + assert result.compacted is True + assert result.summarization_usage == SUMMARIZATION + charge.assert_called_once_with(MODEL, SUMMARIZATION) + recorded.record_llm_token_usage.assert_called_once_with( + "openai", "gpt-4o-mini", 640, 72, ENDPOINT + ) + recorded.record_llm_call.assert_called_once_with("openai", "gpt-4o-mini", ENDPOINT) + + +@pytest.mark.asyncio +async def test_request_served_from_an_earlier_summary( + mocker: MockerFixture, recorded: Any +) -> None: + """A request that makes no summarization call counts nothing.""" + client = _summarizing(mocker) + charge = mocker.Mock() + marker = _msg("user", f"{cc.MARKER_SENTINEL} an earlier summary") + + result = await _compact( + mocker, client, [marker], charge=charge, threshold_ratio=0.9 + ) + + client.responses.create.assert_not_awaited() + assert result.compacted is True + assert result.summarization_usage == TokenCounter() + charge.assert_not_called() + recorded.record_llm_token_usage.assert_not_called() + recorded.record_llm_call.assert_not_called() + + +@pytest.mark.asyncio +async def test_request_on_a_short_conversation( + mocker: MockerFixture, recorded: Any +) -> None: + """A request that is not compacted at all counts nothing.""" + charge = mocker.Mock() + + result = await _compact( + mocker, + _summarizing(mocker), + [_msg("user", "hi")], + charge=charge, + threshold_ratio=0.9, + ) + + assert result.compacted is False + assert result.summarization_usage == TokenCounter() + charge.assert_not_called() + recorded.record_llm_call.assert_not_called() + + +@pytest.mark.asyncio +async def test_call_is_counted_when_the_marker_cannot_be_written( + mocker: MockerFixture, recorded: Any +) -> None: + """The request fails after the call was made; the call is counted all the same.""" + charge = mocker.Mock() + + with pytest.raises(RuntimeError, match="the store is down"): + await _compact( + mocker, + _summarizing(mocker), + _long_conversation(), + charge=charge, + write_marker=mocker.AsyncMock( + side_effect=RuntimeError("the store is down") + ), + ) + + charge.assert_called_once_with(MODEL, SUMMARIZATION) + recorded.record_llm_token_usage.assert_called_once_with( + "openai", "gpt-4o-mini", 640, 72, ENDPOINT + ) + + +@pytest.mark.asyncio +async def test_call_that_returned_no_text_is_counted( + mocker: MockerFixture, recorded: Any +) -> None: + """The request fails for want of a summary; the call is counted all the same.""" + client = _answering(mocker, {SUMMARIZATION_PROMPT: _response(mocker, "", 640, 72)}) + charge = mocker.Mock() + + with pytest.raises(ValueError, match="no extractable text"): + await _compact(mocker, client, _long_conversation(), charge=charge) + + charge.assert_called_once_with(MODEL, SUMMARIZATION) + recorded.record_llm_call.assert_called_once_with("openai", "gpt-4o-mini", ENDPOINT) + + +@pytest.mark.asyncio +async def test_call_without_reported_usage_records_the_call_only( + mocker: MockerFixture, recorded: Any +) -> None: + """There are no tokens to record; the call is recorded and handed on.""" + client = _answering(mocker, {SUMMARIZATION_PROMPT: _response(mocker, "condensed")}) + charge = mocker.Mock() + + result = await _compact(mocker, client, _long_conversation(), charge=charge) + + assert result.summarization_usage == TokenCounter(llm_calls=1) + charge.assert_called_once_with(MODEL, TokenCounter(llm_calls=1)) + recorded.record_llm_token_usage.assert_not_called() + recorded.record_llm_call.assert_called_once_with("openai", "gpt-4o-mini", ENDPOINT) + + +@pytest.mark.asyncio +async def test_no_metrics_without_an_endpoint( + mocker: MockerFixture, recorded: Any +) -> None: + """A caller that does not say which endpoint it serves records no metrics.""" + charge = mocker.Mock() + + result = await _compact( + mocker, + _summarizing(mocker), + _long_conversation(), + charge=charge, + endpoint_path=None, + ) + + assert result.summarization_usage == SUMMARIZATION + charge.assert_called_once_with(MODEL, SUMMARIZATION) + recorded.record_llm_token_usage.assert_not_called() + recorded.record_llm_call.assert_not_called() + + +@pytest.mark.asyncio +async def test_nobody_is_charged_without_a_charge( + mocker: MockerFixture, recorded: Any +) -> None: + """An endpoint without a quota records the call and charges nobody.""" + result = await _compact(mocker, _summarizing(mocker), _long_conversation()) + + assert result.summarization_usage == SUMMARIZATION + recorded.record_llm_call.assert_called_once_with("openai", "gpt-4o-mini", ENDPOINT) + + +def _cache_with_two_summaries(mocker: MockerFixture) -> Any: + """Build a cache whose summaries cross the fold threshold with one more.""" + cache = mocker.Mock() + cache.get_summaries.return_value = [_summary("one", 20), _summary("two", 20)] + return cache + + +@pytest.mark.asyncio +async def test_summarization_and_fold_are_both_counted( + mocker: MockerFixture, recorded: Any +) -> None: + """A request that summarizes and then folds is charged for each of the calls.""" + client = _answering( + mocker, + { + SUMMARIZATION_PROMPT: _response(mocker, "third", 640, 72), + RECURSIVE_RESUMMARIZATION_PROMPT: _response(mocker, "folded", 310, 45), + }, + ) + charge = mocker.Mock() + + result = await _compact( + mocker, + client, + _long_conversation(), + charge=charge, + cache=_cache_with_two_summaries(mocker), + ) + + assert result.summarization_usage == TokenCounter( + input_tokens=950, output_tokens=117, llm_calls=2 + ) + assert charge.call_args_list == [ + mocker.call(MODEL, SUMMARIZATION), + mocker.call(MODEL, FOLD), + ] + assert recorded.record_llm_call.call_count == 2 + + +@pytest.mark.asyncio +async def test_summarization_is_counted_when_the_fold_fails( + mocker: MockerFixture, recorded: Any +) -> None: + """The summary is stored by then and never made again, so it is charged now.""" + client = _answering( + mocker, + { + SUMMARIZATION_PROMPT: _response(mocker, "third", 640, 72), + RECURSIVE_RESUMMARIZATION_PROMPT: RuntimeError("the model is down"), + }, + ) + charge = mocker.Mock() + + with pytest.raises(RuntimeError, match="the model is down"): + await _compact( + mocker, + client, + _long_conversation(), + charge=charge, + cache=_cache_with_two_summaries(mocker), + ) + + charge.assert_called_once_with(MODEL, SUMMARIZATION) + recorded.record_llm_call.assert_called_once_with("openai", "gpt-4o-mini", ENDPOINT) + + +@pytest.mark.asyncio +async def test_fold_that_could_not_be_stored_is_counted( + mocker: MockerFixture, recorded: Any +) -> None: + """The fold call was made and billed, even when its result cannot be kept.""" + _ = recorded + client = _answering( + mocker, + { + SUMMARIZATION_PROMPT: _response(mocker, "third", 640, 72), + RECURSIVE_RESUMMARIZATION_PROMPT: _response(mocker, "folded", 310, 45), + }, + ) + cache = _cache_with_two_summaries(mocker) + cache.replace_summaries.side_effect = CacheError("the cache is down") + charge = mocker.Mock() + + result = await _compact( + mocker, client, _long_conversation(), charge=charge, cache=cache + ) + + assert charge.call_args_list == [ + mocker.call(MODEL, SUMMARIZATION), + mocker.call(MODEL, FOLD), + ] + texts = [message.content for message in result.params.input] + assert not any("folded" in text for text in texts)