Give the capsule to one return, and the run its own copy of state
A request can carry several of this capability's returns — a model can call search twice in one response — and the carrier was identified by message alone, so each of them received the whole capsule. It is selected by message and part now, so exactly one carries it however many share the request. The chat TUI passed its persisted state into the run, so tool synchronisation mutated it in place while the message history was promoted only on success. A cancelled or failed run therefore kept the evidence the tools had recorded and discarded the messages that justified it, and the next question derived its identity from the shorter history: behind the recorded epoch, refused as non-append-only, the conversation unusable until cleared. The run gets a copy, promoted with the messages or not at all. Five decorators had been left attached to a helper by an earlier extraction, which pytest does not collect, so the resume case they carried was silently untested. The wire test covers both resume shapes again, no prompt and deferred results.
This commit is contained in:
parent
a69f3a8a98
commit
11b1bfbc94
5 changed files with 154 additions and 28 deletions
|
|
@ -253,16 +253,17 @@ def compact_history(
|
|||
if not isinstance(message, ModelRequest):
|
||||
continue
|
||||
parts: list[Any] = []
|
||||
for part in message.parts:
|
||||
for position, part in enumerate(message.parts):
|
||||
if isinstance(part, ToolReturnPart) and part.tool_name in owned_tools:
|
||||
body = capsule_text or RECEIPT if index == carrier else RECEIPT
|
||||
carries = (index, position) == carrier
|
||||
body = capsule_text or RECEIPT if carries else RECEIPT
|
||||
parts.append(replace(part, content=body))
|
||||
elif isinstance(part, UserPromptPart):
|
||||
if (kept := _strip_our_pictures(part)) is not None:
|
||||
parts.append(kept)
|
||||
else:
|
||||
parts.append(part)
|
||||
if index == carrier and capsule_images:
|
||||
if carrier is not None and index == carrier[0] and capsule_images:
|
||||
parts.append(UserPromptPart(content=list(capsule_images)))
|
||||
if not parts:
|
||||
# A request with no parts is not a message; whatever emptied it was not
|
||||
|
|
@ -275,15 +276,21 @@ def compact_history(
|
|||
|
||||
def _newest_owned_return(
|
||||
messages: list[ModelMessage], boundary: int, owned_tools: frozenset[str]
|
||||
) -> int | None:
|
||||
"""Index of the last earlier request holding one of our evidence returns."""
|
||||
) -> tuple[int, int] | None:
|
||||
"""Where the last of our evidence returns is, as message and part.
|
||||
|
||||
The part matters: a model can call search twice in one response, so one request
|
||||
can hold several of our returns, and giving the capsule to each duplicates the
|
||||
whole of it.
|
||||
"""
|
||||
for index in range(min(boundary, len(messages)) - 1, -1, -1):
|
||||
message = messages[index]
|
||||
if isinstance(message, ModelRequest) and any(
|
||||
isinstance(part, ToolReturnPart) and part.tool_name in owned_tools
|
||||
for part in message.parts
|
||||
):
|
||||
return index
|
||||
if not isinstance(message, ModelRequest):
|
||||
continue
|
||||
for position in range(len(message.parts) - 1, -1, -1):
|
||||
part = message.parts[position]
|
||||
if isinstance(part, ToolReturnPart) and part.tool_name in owned_tools:
|
||||
return index, position
|
||||
return None
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import uuid
|
||||
from collections.abc import Iterable, Sequence
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
|
@ -208,7 +209,12 @@ class ChatApp(App):
|
|||
await chat_history.show_thinking()
|
||||
|
||||
message = None
|
||||
deps = ChatDeps(state=self._state)
|
||||
# The run gets a copy: state and message history have to advance together.
|
||||
# A cancelled or failed run discards its messages, and state that advanced
|
||||
# anyway would leave the next question deriving its identity from a shorter
|
||||
# history than the evidence already recorded — refused as non-append-only,
|
||||
# with the conversation stuck until it is cleared.
|
||||
deps = ChatDeps(state=deepcopy(self._state))
|
||||
|
||||
try:
|
||||
async with self._agent.run_stream_events(
|
||||
|
|
|
|||
|
|
@ -872,20 +872,6 @@ def _in_flight_history() -> list[Any]:
|
|||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"resume_kwargs",
|
||||
[
|
||||
pytest.param({}, id="no prompt"),
|
||||
pytest.param(
|
||||
{"deferred_tool_results": DeferredToolResults()}, id="deferred results"
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
def _record(deps: Deps, namespace: str) -> CapabilityEvidenceRecord:
|
||||
return CapabilityEvidenceRecord.model_validate(deps.state[namespace]["evidence"])
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from typing import Any, cast
|
|||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic_ai import Agent, RunContext
|
||||
from pydantic_ai import Agent, DeferredToolResults, RunContext
|
||||
from pydantic_ai.messages import (
|
||||
BinaryContent,
|
||||
ModelRequest,
|
||||
|
|
@ -346,9 +346,18 @@ async def test_with_the_compactor_a_new_question_compacts_the_previous_one(
|
|||
assert returns_of(wire[-1]) == [RECEIPT]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"resume_kwargs",
|
||||
[
|
||||
pytest.param({}, id="no prompt"),
|
||||
pytest.param(
|
||||
{"deferred_tool_results": DeferredToolResults()}, id="deferred results"
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_resumed_question_keeps_the_evidence_it_is_answering_from(
|
||||
temp_db_path,
|
||||
temp_db_path, resume_kwargs
|
||||
):
|
||||
"""The boundary is the stored identity of the question in progress.
|
||||
|
||||
|
|
@ -370,7 +379,9 @@ async def test_a_resumed_question_keeps_the_evidence_it_is_answering_from(
|
|||
*in_flight_history(),
|
||||
]
|
||||
|
||||
await agent.run(deps=resuming_deps(question=4), message_history=history)
|
||||
await agent.run(
|
||||
deps=resuming_deps(question=4), message_history=history, **resume_kwargs
|
||||
)
|
||||
|
||||
assert returns_of(wire[-1]) == [RECEIPT, "EVIDENCE FOR THE LIVE TURN"]
|
||||
|
||||
|
|
@ -653,3 +664,66 @@ def test_a_request_is_never_left_with_no_parts():
|
|||
|
||||
assert all(message.parts for message in compacted)
|
||||
assert compacted[0] is ours_alone
|
||||
|
||||
|
||||
def test_two_owned_returns_in_one_request_yield_one_capsule():
|
||||
"""A model can search twice in one response, so a request can hold two returns.
|
||||
|
||||
Identifying the carrier by message alone gave every return in it the capsule,
|
||||
which duplicates the whole thing — and it is unbounded.
|
||||
"""
|
||||
messages = [
|
||||
ModelRequest(parts=[UserPromptPart("first")]),
|
||||
ModelResponse(
|
||||
parts=[
|
||||
ToolCallPart("rag_search", {"query": "a"}, "call-1"),
|
||||
ToolCallPart("rag_search", {"query": "b"}, "call-2"),
|
||||
]
|
||||
),
|
||||
ModelRequest(
|
||||
parts=[
|
||||
ToolReturnPart("rag_search", "FIRST EVIDENCE", "call-1"),
|
||||
ToolReturnPart("rag_search", "SECOND EVIDENCE", "call-2"),
|
||||
]
|
||||
),
|
||||
ModelResponse(parts=[TextPart("an answer")]),
|
||||
ModelRequest(parts=[UserPromptPart("second")]),
|
||||
]
|
||||
|
||||
compacted = compact_history(
|
||||
messages, boundary=4, owned_tools=OWNED, capsule_text="CAPSULE"
|
||||
)
|
||||
|
||||
assert returns_of(compacted) == [RECEIPT, "CAPSULE"]
|
||||
|
||||
|
||||
def test_the_capsule_is_attached_beside_the_newest_return_of_that_request():
|
||||
messages = [
|
||||
ModelRequest(parts=[UserPromptPart("first")]),
|
||||
ModelResponse(
|
||||
parts=[
|
||||
ToolCallPart("rag_search", {"query": "a"}, "call-1"),
|
||||
ToolCallPart("rag_search", {"query": "b"}, "call-2"),
|
||||
]
|
||||
),
|
||||
ModelRequest(
|
||||
parts=[
|
||||
ToolReturnPart("rag_search", "FIRST", "call-1"),
|
||||
ToolReturnPart("rag_search", "SECOND", "call-2"),
|
||||
]
|
||||
),
|
||||
ModelResponse(parts=[TextPart("an answer")]),
|
||||
ModelRequest(parts=[UserPromptPart("second")]),
|
||||
]
|
||||
fresh = BinaryContent(data=b"cited-bytes", media_type="image/png")
|
||||
|
||||
compacted = compact_history(
|
||||
messages,
|
||||
boundary=4,
|
||||
owned_tools=OWNED,
|
||||
capsule_text="CAPSULE",
|
||||
capsule_images=[picture_label("cited", "#/pictures/1"), fresh],
|
||||
)
|
||||
|
||||
assert images_of(compacted) == [fresh]
|
||||
assert returns_of(compacted) == [RECEIPT, "CAPSULE"]
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
|
|
@ -412,3 +413,55 @@ async def test_chat_app_open_failure_surfaces_real_error(tmp_path: Path):
|
|||
with pytest.raises(FileNotFoundError):
|
||||
async with app.run_test():
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_cancelled_run_does_not_advance_persisted_state(temp_db_path: Path):
|
||||
"""State and message history have to move together, or the thread bricks.
|
||||
|
||||
A cancelled run keeps whatever the tools wrote but discards the run's messages.
|
||||
If the state advanced, the next question derives its identity from the shorter
|
||||
history, lands behind the recorded evidence epoch, and is refused as
|
||||
non-append-only — leaving the conversation unusable until cleared.
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
app, mock_client = _make_app(temp_db_path)
|
||||
|
||||
class CancellingRun:
|
||||
"""A run that writes evidence through the tools, then is cancelled."""
|
||||
|
||||
def __init__(self, deps):
|
||||
self._deps = deps
|
||||
|
||||
async def __aenter__(self):
|
||||
self._deps.state["rag"] = {
|
||||
"evidence": {"question": 0, "latest_evidence_epoch": 7}
|
||||
}
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_):
|
||||
return False
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise asyncio.CancelledError
|
||||
|
||||
with patch("haiku.rag.chat.app.HaikuRAG", return_value=mock_client):
|
||||
async with app.run_test():
|
||||
app._state = {
|
||||
"rag": {"evidence": {"question": 0, "latest_evidence_epoch": 0}}
|
||||
}
|
||||
before = deepcopy(app._state)
|
||||
|
||||
class Agent:
|
||||
def run_stream_events(self, *_, deps, **__):
|
||||
return CancellingRun(deps)
|
||||
|
||||
app._agent = Agent() # type: ignore[assignment]
|
||||
await app._run_agent("a question")
|
||||
|
||||
assert app._state == before
|
||||
assert app._messages == []
|
||||
|
|
|
|||
Loading…
Reference in a new issue