Skip to content

Commit 4a19bc4

Browse files
committed
Wired compaction turn persistence into agentic query flows
1 parent 0377864 commit 4a19bc4

10 files changed

Lines changed: 219 additions & 193 deletions

File tree

src/app/endpoints/query.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -233,7 +233,7 @@ async def query_endpoint_handler(
233233
responses_params,
234234
moderation_result,
235235
endpoint_path,
236-
compaction.original_input if compaction.compacted else None,
236+
compaction.original_input,
237237
)
238238

239239
if moderation_result.decision == "passed":

src/app/endpoints/streaming_query.py

Lines changed: 5 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -338,6 +338,7 @@ async def streaming_query_endpoint_handler( # pylint: disable=too-many-locals
338338
responses_params=responses_params,
339339
context=context,
340340
endpoint_path=endpoint_path,
341+
original_input=None,
341342
)
342343

343344
# Combine inline RAG results (BYOK + Solr) with tool-based results
@@ -353,6 +354,8 @@ async def streaming_query_endpoint_handler( # pylint: disable=too-many-locals
353354
responses_params=responses_params,
354355
turn_summary=turn_summary,
355356
background_topic_summary_tasks=_background_topic_summary_tasks,
357+
emit_start=True,
358+
original_input=None,
356359
),
357360
media_type=response_media_type,
358361
)
@@ -387,7 +390,6 @@ async def retrieve_response_generator(
387390
if context.moderation_result.decision == "blocked":
388391
turn_summary.llm_response = context.moderation_result.message
389392
turn_summary.id = context.moderation_result.moderation_id
390-
turn_summary.output_items = [context.moderation_result.refusal_response]
391393
# In compacted mode the conversation parameter was omitted, so the
392394
# refusal turn (with the original input) is persisted by
393395
# generate_response; storing it here too would duplicate it.
@@ -506,6 +508,7 @@ async def generate_response_with_compaction(
506508
responses_params=responses_params,
507509
context=context,
508510
endpoint_path=endpoint_path,
511+
original_input=compacted_original_input,
509512
)
510513
except HTTPException as e:
511514
yield http_exception_stream_event(e)
@@ -705,7 +708,7 @@ async def generate_response( # pylint: disable=too-many-arguments,too-many-posi
705708
if original_input is not None
706709
else context.query_request.query
707710
),
708-
turn_summary.output_items,
711+
[], # field was removed from TurnSummary
709712
)
710713
except Exception: # pylint: disable=broad-except
711714
logger.exception(
@@ -884,10 +887,6 @@ async def response_generator( # pylint: disable=too-many-branches,too-many-stat
884887
getattr(chunk, "response"), # noqa: B009
885888
)
886889
turn_summary.llm_response = turn_summary.llm_response or "".join(text_parts)
887-
# Capture structured output items for compacted-mode turn storage
888-
# (LCORE-1572), so the persisted turn keeps non-text output items
889-
# rather than being flattened to the response text.
890-
turn_summary.output_items = list(latest_response_object.output or [])
891890
event_id = chunk_id
892891
chunk_id += 1
893892
turn_summary.next_chunk_id = chunk_id
@@ -906,9 +905,6 @@ async def response_generator( # pylint: disable=too-many-branches,too-many-stat
906905
OpenAIResponseObject,
907906
getattr(chunk, "response"), # noqa: B009
908907
)
909-
# Capture any partial output items so a compacted-mode turn is not
910-
# persisted with empty output on these terminals (LCORE-1572).
911-
turn_summary.output_items = list(latest_response_object.output or [])
912908
error_message = (
913909
latest_response_object.error.message
914910
if latest_response_object.error

src/models/common/turn_summary.py

Lines changed: 0 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@
55

66
from typing import Any, Optional
77

8-
from llama_stack_api import OpenAIResponseOutput
98
from pydantic import AnyUrl, BaseModel, Field
109

1110
from utils.token_counter import TokenCounter
@@ -109,11 +108,6 @@ class TurnSummary(BaseModel):
109108
rag_chunks: list[RAGChunk] = Field(default_factory=list)
110109
referenced_documents: list[ReferencedDocument] = Field(default_factory=list)
111110
token_usage: TokenCounter = Field(default_factory=TokenCounter)
112-
output_items: list[OpenAIResponseOutput] = Field(
113-
default_factory=list,
114-
description="Structured response output items, captured for compacted-mode "
115-
"turn persistence (LCORE-1572). Empty on the non-compacted path.",
116-
)
117111
partial_tokens: list[str] = Field(
118112
default_factory=list,
119113
description="Accumulated text deltas during streaming, used to reconstruct "
Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,13 @@
11
"""Pydantic AI provider for Llama Stack."""
22

3-
from pydantic_ai_lightspeed.llamastack._model import LlamaStackResponsesModel
3+
from pydantic_ai_lightspeed.llamastack._model import (
4+
CompactionTurnContext,
5+
LlamaStackResponsesModel,
6+
)
47
from pydantic_ai_lightspeed.llamastack._provider import LlamaStackProvider
58

6-
__all__ = ["LlamaStackProvider", "LlamaStackResponsesModel"]
9+
__all__ = [
10+
"CompactionTurnContext",
11+
"LlamaStackProvider",
12+
"LlamaStackResponsesModel",
13+
]

src/pydantic_ai_lightspeed/llamastack/_model.py

Lines changed: 156 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -10,23 +10,30 @@
1010
deltas must be replayed with the matching suffix so pydantic_ai can append the
1111
streamed ``tool_args`` content to the correct part.
1212
13+
When compaction omits the ``conversation`` parameter from inference requests,
14+
``LlamaStackResponsesModel`` appends completed turns to the conversation via
15+
``CompactionTurnContext`` (in :meth:`_responses_create` for non-streaming rounds,
16+
and on ``response.completed`` for streaming rounds in :meth:`request_stream`).
17+
1318
This module provides ``LlamaStackResponsesModel`` which wraps the event stream to
1419
buffer those early delta events and replay them correctly once the item is announced.
1520
"""
1621

1722
from __future__ import annotations as _annotations
1823

1924
from collections import defaultdict
20-
from collections.abc import AsyncIterator
25+
from collections.abc import AsyncIterator, Sequence
2126
from contextlib import asynccontextmanager
22-
from typing import Any, cast
27+
from dataclasses import dataclass
28+
from typing import Any, Literal, Optional, cast, overload
2329

30+
from llama_stack_client import AsyncLlamaStackClient
2431
from openai import AsyncStream
2532
from openai.types import responses
2633
from pydantic_ai import UnexpectedModelBehavior
2734
from pydantic_ai._run_context import RunContext
2835
from pydantic_ai._utils import PeekableAsyncStream, Unset, number_to_datetime
29-
from pydantic_ai.messages import ModelMessage, ModelResponse
36+
from pydantic_ai.messages import ModelMessage, ModelRequest, ModelResponse
3037
from pydantic_ai.models import (
3138
ModelRequestParameters,
3239
StreamedResponse,
@@ -38,13 +45,38 @@
3845
OpenAIResponsesStreamedResponse,
3946
_map_api_errors,
4047
)
48+
from pydantic_ai.profiles import ModelProfileSpec
4149
from pydantic_ai.settings import ModelSettings
4250

4351
from log import get_logger
52+
from models.common.responses.types import ResponseInput
53+
from pydantic_ai_lightspeed.llamastack._provider import LlamaStackProvider
54+
from utils.conversations import append_turn_items_to_conversation
4455

4556
logger = get_logger(__name__)
4657

4758

59+
@dataclass
60+
class CompactionTurnContext:
61+
"""Mutable state for manually persisting compacted agent turns.
62+
63+
``latest_round_input`` is initialized to the real user query. The create patch
64+
leaves it unchanged on the first LLM round, then records pydantic-ai input
65+
for follow-up rounds after that turn is persisted.
66+
67+
Attributes:
68+
client: Llama Stack client used to append conversation items.
69+
conversation_id: Conversation to store turns against.
70+
latest_round_input: Input stored for the current or next inference round.
71+
original_input_persisted: Whether the first compacted round was appended.
72+
"""
73+
74+
client: AsyncLlamaStackClient
75+
conversation_id: str
76+
latest_round_input: ResponseInput
77+
original_input_persisted: bool = False
78+
79+
4880
class _FilteredResponseStream:
4981
"""Wraps an OpenAI AsyncStream to reorder spurious events from Llama Stack.
5082
@@ -58,13 +90,19 @@ class _FilteredResponseStream:
5890
a closing ``}`` to complete the outer JSON object that pydantic_ai opens.
5991
"""
6092

61-
def __init__(self, source: AsyncStream[responses.ResponseStreamEvent]) -> None:
93+
def __init__(
94+
self,
95+
source: AsyncStream[responses.ResponseStreamEvent],
96+
compaction: Optional[CompactionTurnContext] = None,
97+
) -> None:
6298
"""Wrap an existing stream with reordering logic.
6399
64100
Args:
65101
source: The raw OpenAI AsyncStream to reorder.
102+
compaction: Compaction state for turn persistence, if active.
66103
"""
67104
self._source = source
105+
self._compaction = compaction
68106
self._announced_item_ids: set[str] = set()
69107
self._buffered_deltas: dict[
70108
str, list[responses.ResponseFunctionCallArgumentsDeltaEvent]
@@ -112,6 +150,19 @@ async def _filtered_iter(
112150
self._buffered_deltas[event.item_id].append(event)
113151
continue
114152

153+
if (
154+
isinstance(event, responses.ResponseCompletedEvent)
155+
and self._compaction is not None
156+
):
157+
compaction = self._compaction
158+
await append_turn_items_to_conversation(
159+
compaction.client,
160+
compaction.conversation_id,
161+
compaction.latest_round_input,
162+
cast(Sequence[Any], event.response.output),
163+
)
164+
compaction.original_input_persisted = True
165+
115166
yield event
116167

117168
def _replay_buffered_deltas(
@@ -179,8 +230,108 @@ class LlamaStackResponsesModel(OpenAIResponsesModel):
179230
Overrides the streaming response processing to buffer and replay
180231
``ResponseFunctionCallArgumentsDeltaEvent`` events that Llama Stack emits
181232
before the corresponding ``McpCall`` or ``ResponseFunctionToolCall`` item.
233+
234+
When ``compaction`` is set, completed inference rounds are appended to the
235+
conversation because compacted mode omits the ``conversation`` parameter.
182236
"""
183237

238+
def __init__( # pylint: disable=too-many-arguments
239+
self,
240+
model_name: str,
241+
provider: LlamaStackProvider,
242+
profile: ModelProfileSpec | None = None,
243+
settings: ModelSettings | None = None,
244+
compaction: Optional[CompactionTurnContext] = None,
245+
) -> None:
246+
"""Initialize the model.
247+
248+
Args:
249+
model_name: Model identifier passed to pydantic-ai.
250+
provider: Pydantic AI provider or provider name.
251+
profile: Optional model profile override.
252+
settings: Optional pydantic-ai model settings.
253+
compaction: Compaction state when turns must be stored manually.
254+
"""
255+
super().__init__(
256+
model_name,
257+
provider=provider,
258+
profile=profile,
259+
settings=settings,
260+
)
261+
self.compaction = compaction
262+
263+
@overload
264+
async def _responses_create(
265+
self,
266+
messages: list[ModelRequest | ModelResponse],
267+
stream: Literal[False],
268+
model_settings: OpenAIResponsesModelSettings,
269+
model_request_parameters: ModelRequestParameters,
270+
) -> responses.Response: ...
271+
272+
@overload
273+
async def _responses_create(
274+
self,
275+
messages: list[ModelRequest | ModelResponse],
276+
stream: Literal[True],
277+
model_settings: OpenAIResponsesModelSettings,
278+
model_request_parameters: ModelRequestParameters,
279+
) -> AsyncStream[responses.ResponseStreamEvent]: ...
280+
281+
async def _responses_create(
282+
self,
283+
messages: list[ModelRequest | ModelResponse],
284+
stream: bool,
285+
model_settings: OpenAIResponsesModelSettings,
286+
model_request_parameters: ModelRequestParameters,
287+
) -> (
288+
responses.Response | AsyncStream[responses.ResponseStreamEvent] | ModelResponse
289+
):
290+
"""Create a Responses API request with compacted turn persistence.
291+
292+
After the first compacted round is persisted, records pydantic-ai input
293+
for follow-up tool-loop rounds. Non-streaming responses are appended
294+
immediately; streaming persistence is handled in :meth:`request_stream`.
295+
"""
296+
compaction = self.compaction
297+
if compaction is not None and compaction.original_input_persisted:
298+
request_params = await self._build_responses_request_params(
299+
messages,
300+
model_settings,
301+
model_request_parameters,
302+
self.profile,
303+
)
304+
compaction.latest_round_input = cast(ResponseInput, request_params.input)
305+
306+
result: (
307+
responses.Response
308+
| AsyncStream[responses.ResponseStreamEvent]
309+
| ModelResponse
310+
)
311+
if stream:
312+
result = await super()._responses_create(
313+
messages, True, model_settings, model_request_parameters
314+
)
315+
else:
316+
result = await super()._responses_create(
317+
messages, False, model_settings, model_request_parameters
318+
)
319+
320+
if (
321+
compaction is not None
322+
and not stream
323+
and isinstance(result, responses.Response)
324+
):
325+
await append_turn_items_to_conversation(
326+
compaction.client,
327+
compaction.conversation_id,
328+
compaction.latest_round_input,
329+
cast(Sequence[Any], result.output),
330+
)
331+
compaction.original_input_persisted = True
332+
333+
return result
334+
184335
async def request( # pylint: disable=unused-argument
185336
self,
186337
messages: list[ModelMessage],
@@ -274,7 +425,7 @@ async def request_stream( # pylint: disable=unused-argument
274425
messages, True, model_settings_cast, model_request_parameters
275426
)
276427

277-
filtered_stream = _FilteredResponseStream(response)
428+
filtered_stream = _FilteredResponseStream(response, self.compaction)
278429

279430
async with response:
280431
peekable: PeekableAsyncStream[

src/utils/agents/query.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -282,7 +282,7 @@ async def retrieve_agent_response(
282282
responses_params: ResponsesApiParams,
283283
moderation_result: ShieldModerationResult,
284284
endpoint_path: str,
285-
_original_input: Optional[ResponseInput] = None,
285+
original_input: Optional[ResponseInput] = None,
286286
) -> TurnSummary:
287287
"""Retrieve a turn summary from a blocking agent run.
288288
@@ -293,7 +293,7 @@ async def retrieve_agent_response(
293293
responses_params: Prepared Responses API parameters.
294294
moderation_result: Shield moderation outcome for the turn.
295295
endpoint_path: Endpoint path used for metric labeling.
296-
_original_input: Original user input before the explicit-input rewrite.
296+
original_input: Original user input before the explicit-input rewrite.
297297
298298
Returns:
299299
Turn summary for the completed agent run.
@@ -305,15 +305,17 @@ async def retrieve_agent_response(
305305
await append_turn_items_to_conversation(
306306
client,
307307
responses_params.conversation,
308-
responses_params.input,
308+
original_input or responses_params.input,
309309
[moderation_result.refusal_response],
310310
)
311311
return TurnSummary(
312312
id=moderation_result.moderation_id,
313313
llm_response=moderation_result.message,
314314
)
315315
try:
316-
agent = build_agent(client, responses_params, configuration.skills)
316+
agent = build_agent(
317+
client, responses_params, configuration.skills, original_input
318+
)
317319
logger.debug("Starting agent non-streaming response processing")
318320
run_result = await agent.run(cast(str, responses_params.input))
319321
except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as exc:

0 commit comments

Comments
 (0)