From d9bd3a701fde1320096b50e4321341f9d2d7670e Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Sat, 8 Aug 2026 14:43:29 +0300 Subject: [PATCH] Take the compaction turn boundary from the run, not the message shape --- CHANGELOG.md | 3 +- .../haiku/rag/capabilities/_base.py | 42 ++++---- tests/capabilities/test_capabilities.py | 96 +++++++++++++++++-- 3 files changed, 106 insertions(+), 35 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4a7b7930..f75350e7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,7 +8,8 @@ - `haiku.rag.store` no longer re-exports `Store`; import it from `haiku.rag.store.engine`. - `docling-local` reuses one docling `DocumentConverter` per set of conversion options instead of building one per document, so local layout, table and OCR models are no longer loaded per document. Conversions through a shared converter are serialized. - Prior-question tool output is trimmed from the model request only; `all_messages()` retains what the run gathered. -- A search result carrying page images no longer counts as a turn boundary; evidence retrieved earlier in the same turn survives. +- Evidence retrieved for the current question is no longer trimmed mid-question, whether or not it carries page images. +- `rag_cite` / `analysis_cite` returns are no longer replaced by the prior-question notice. ## [0.73.0] - 2026-08-06 diff --git a/haiku_rag_slim/haiku/rag/capabilities/_base.py b/haiku_rag_slim/haiku/rag/capabilities/_base.py index ae933464..41994ed8 100644 --- a/haiku_rag_slim/haiku/rag/capabilities/_base.py +++ b/haiku_rag_slim/haiku/rag/capabilities/_base.py @@ -16,7 +16,6 @@ from pydantic_ai.messages import ( ToolCallPart, ToolReturn, ToolReturnPart, - UserPromptPart, ) from pydantic_ai.models import ModelRequestContext from pydantic_ai.run import AgentRunResult @@ -82,22 +81,11 @@ PRIOR_TURN_NOTICE = ( ) -def _is_user_turn(message: ModelMessage) -> bool: - """Whether this message is a user turn rather than tool output. - - Content attached to a ``ToolReturn`` (page images) arrives as its own - ``UserPromptPart`` in the same ``ModelRequest`` as the ``ToolReturnPart``, - so a bare ``UserPromptPart`` check reads tool output as a new turn. - """ - if not isinstance(message, ModelRequest): - return False - return any(isinstance(part, UserPromptPart) for part in message.parts) and not any( - isinstance(part, ToolReturnPart) for part in message.parts - ) - - def _compact_old_tool_returns( - messages: list[ModelMessage], tool_names: frozenset[str] + messages: list[ModelMessage], + tool_names: frozenset[str], + *, + turn_start: int, ) -> list[ModelMessage]: """Remove bulky earlier-question evidence while retaining the current one. @@ -109,17 +97,19 @@ def _compact_old_tool_returns( they accumulate — but a follow-up about a figure already shown ("what colour is that box?") carries no terms that could retrieve it again, so removing the image turns an answerable question into a refusal. - """ - latest_user_message = -1 - for index, message in enumerate(messages): - if _is_user_turn(message): - latest_user_message = index - if latest_user_message < 0: + ``turn_start`` is how many messages existed when the current question + arrived, so everything below it belongs to an earlier one. The run reports + it rather than this function deriving it from message shape: a + ``UserPromptPart`` mid-question is as likely to be page images on a tool + return, or a notice a capability injected, and reading either as the next + question strips evidence the model is still answering from. + """ + if turn_start <= 0: return messages compacted = list(messages) - for index, message in enumerate(messages[:latest_user_message]): + for index, message in enumerate(messages[:turn_start]): if not isinstance(message, ModelRequest): continue parts = [ @@ -162,6 +152,7 @@ class RAGCapabilityBase[StateT: BaseModel](AbstractCapability[Any]): search_count: int = field(default=0, repr=False) request_count: int = field(default=0, repr=False) grace_requests_used: int = field(default=0, repr=False) + turn_start: int = field(default=0, repr=False) async def for_run(self, ctx: RunContext[Any]) -> "RAGCapabilityBase[StateT]": outer = getattr(ctx.deps, "state", None) @@ -179,6 +170,7 @@ class RAGCapabilityBase[StateT: BaseModel](AbstractCapability[Any]): search_count=0, request_count=0, grace_requests_used=0, + turn_start=len(ctx.messages), ) run_capability._sync_state() return run_capability @@ -202,7 +194,9 @@ class RAGCapabilityBase[StateT: BaseModel](AbstractCapability[Any]): destroy the host's record of what was retrieved. """ request_context.messages = _compact_old_tool_returns( - request_context.messages, self.tool_names + request_context.messages, + self.tool_names - {self._cite_tool_name}, + turn_start=self.turn_start, ) return await handler(request_context) diff --git a/tests/capabilities/test_capabilities.py b/tests/capabilities/test_capabilities.py index f107df51..aa360a73 100644 --- a/tests/capabilities/test_capabilities.py +++ b/tests/capabilities/test_capabilities.py @@ -803,7 +803,9 @@ def test_prior_turn_tool_results_are_compacted_but_current_evidence_is_kept(): ), ] - compacted = _compact_old_tool_returns(messages, frozenset({"rag_search"})) + compacted = _compact_old_tool_returns( + messages, frozenset({"rag_search"}), turn_start=3 + ) old_return = compacted[2].parts[0] current_return = compacted[5].parts[0] @@ -813,7 +815,7 @@ def test_prior_turn_tool_results_are_compacted_but_current_evidence_is_kept(): assert current_return.content == "current evidence" -def test_tool_results_are_unchanged_when_history_has_no_user_prompt(): +def test_nothing_is_compacted_on_the_first_question(): messages = [ ModelResponse(parts=[ToolCallPart("rag_search", {}, "current-call")]), ModelRequest( @@ -821,7 +823,9 @@ def test_tool_results_are_unchanged_when_history_has_no_user_prompt(): ), ] - compacted = _compact_old_tool_returns(messages, frozenset({"rag_search"})) + compacted = _compact_old_tool_returns( + messages, frozenset({"rag_search"}), turn_start=0 + ) assert compacted is messages current_return = compacted[1].parts[0] @@ -847,18 +851,33 @@ def _search_exchange(call_id: str, evidence: str, *, images: bool): ] -def test_image_bearing_tool_return_does_not_end_the_current_turn(): +@pytest.mark.parametrize( + ("label", "trailing"), + [ + ("page images on a tool return", [UserPromptPart(content=[PAGE_IMAGE])]), + ("a notice injected mid-run", [UserPromptPart("You answered without citing")]), + ], +) +def test_current_turn_evidence_survives_later_user_prompt_parts(label, trailing): + """Nothing that arrives mid-question may be read as the next question. + + Page images on a tool return and a notice this capability injects both + appear as a ``UserPromptPart`` after the question, so deriving the turn from + message shape stripped evidence the model was still answering from. + """ messages = [ ModelRequest(parts=[UserPromptPart("current question")]), *_search_exchange("first-call", "first evidence", images=False), - *_search_exchange("second-call", "second evidence", images=True), + ModelRequest(parts=trailing), ] - compacted = _compact_old_tool_returns(messages, frozenset({"rag_search"})) + compacted = _compact_old_tool_returns( + messages, frozenset({"rag_search"}), turn_start=0 + ) first_return = compacted[2].parts[0] assert isinstance(first_return, ToolReturnPart) - assert first_return.content == "first evidence" + assert first_return.content == "first evidence", label def test_prior_turn_images_outlive_their_tool_return(): @@ -874,7 +893,9 @@ def test_prior_turn_images_outlive_their_tool_return(): ModelRequest(parts=[UserPromptPart("current question")]), ] - compacted = _compact_old_tool_returns(messages, frozenset({"rag_search"})) + compacted = _compact_old_tool_returns( + messages, frozenset({"rag_search"}), turn_start=4 + ) old_return = compacted[2].parts[0] assert isinstance(old_return, ToolReturnPart) @@ -886,7 +907,7 @@ def test_prior_turn_images_outlive_their_tool_return(): ) -def test_user_attached_image_starts_a_turn_and_is_never_dropped(): +def test_user_attached_image_is_never_dropped(): messages = [ ModelRequest(parts=[UserPromptPart("old question")]), *_search_exchange("old-call", "old evidence", images=False), @@ -894,7 +915,9 @@ def test_user_attached_image_starts_a_turn_and_is_never_dropped(): ModelRequest(parts=[UserPromptPart(content=[PAGE_IMAGE, "what is this?"])]), ] - compacted = _compact_old_tool_returns(messages, frozenset({"rag_search"})) + compacted = _compact_old_tool_returns( + messages, frozenset({"rag_search"}), turn_start=4 + ) old_return = compacted[2].parts[0] assert isinstance(old_return, ToolReturnPart) @@ -940,3 +963,56 @@ async def test_compaction_never_reaches_the_stored_message_history(temp_db_path) ] assert "REAL EVIDENCE" in returns assert PRIOR_TURN_NOTICE not in returns + + +@pytest.mark.asyncio +async def test_prior_turn_search_is_compacted_but_its_cite_receipt_is_not(temp_db_path): + """Only evidence is compacted, and the boundary comes from the run. + + A cite acknowledgement is a receipt, not evidence: replacing it lengthened + the request and erased the record that citations had been registered. + """ + capability = create_rag( + db_path=temp_db_path, config=AppConfig(), defer_loading=False + ) + turns = iter( + [ + ModelResponse(parts=[ToolCallPart("rag_search", {"query": "q"}, "call-1")]), + ModelResponse( + parts=[ToolCallPart("rag_cite", {"chunk_ids": ["chunk-1"]}, "call-2")] + ), + ModelResponse(parts=[TextPart("first answer")]), + ModelResponse(parts=[TextPart("second answer")]), + ] + ) + wire: list[list[Any]] = [] + + async def model(messages, _info): + wire.append(list(messages)) + return next(turns) + + agent = Agent(FunctionModel(model), deps_type=Deps, capabilities=[capability]) + deps = Deps(state={"rag": RAGState().model_dump(mode="json")}) + + async def search(self, query: str, _limit: int | None) -> str: + """Record a result the way the real search does, so citing resolves.""" + cast(Any, self.state).searches[query] = [ + SearchResult(content="evidence", score=1.0, chunk_id="chunk-1") + ] + return "REAL EVIDENCE" + + with patch.object(RAGCapability, "_search", search): + first = await agent.run("old question", deps=deps) + await agent.run( + "current question", deps=deps, message_history=first.all_messages() + ) + + prior = { + part.tool_name: str(part.content) + for message in wire[-1] + if isinstance(message, ModelRequest) + for part in message.parts + if isinstance(part, ToolReturnPart) + } + assert prior["rag_search"] == PRIOR_TURN_NOTICE + assert prior["rag_cite"] != PRIOR_TURN_NOTICE