fix(markdown): a rewritten assistant reply reseeds the session instead of resuming it

This commit is contained in:
hh
2026-08-29 01:37:52 +02:00
parent 772db11c7f
commit d3a3e7b1f1
5 changed files with 60 additions and 15 deletions
+7 -1
View File
@@ -296,6 +296,7 @@ class ClaudeSdkBackend:
system: str | None = None, # noqa: ARG002 - the agent owns its prompt system: str | None = None, # noqa: ARG002 - the agent owns its prompt
conversation_id: str | None = None, conversation_id: str | None = None,
session_id: str | None = None, session_id: str | None = None,
reseed: bool = False,
capture: TurnCapture | None = None, capture: TurnCapture | None = None,
kind: str = "deep", kind: str = "deep",
pinned: bool = False, pinned: bool = False,
@@ -315,7 +316,9 @@ class ClaudeSdkBackend:
prior = history[:-1] prior = history[:-1]
key = conversation_id or fingerprint(prior) key = conversation_id or fingerprint(prior)
spec = _SessionSpec(kind=kind, pinned=pinned, tools=tools) spec = _SessionSpec(kind=kind, pinned=pinned, tools=tools)
live = await self._acquire(key, session_id=session_id, history=prior, spec=spec) live = await self._acquire(
key, session_id=session_id, history=prior, spec=spec, reseed=reseed
)
message_id = f"msg_{uuid.uuid4().hex}" message_id = f"msg_{uuid.uuid4().hex}"
yield build_message_start(message_id=message_id, model=self._agent.model) yield build_message_start(message_id=message_id, model=self._agent.model)
turn = _Turn() turn = _Turn()
@@ -484,10 +487,13 @@ class ClaudeSdkBackend:
session_id: str | None, session_id: str | None,
history: list[dict[str, Any]], history: list[dict[str, Any]],
spec: _SessionSpec, spec: _SessionSpec,
reseed: bool = False,
) -> Session: ) -> Session:
live = self._pool.get(key) live = self._pool.get(key)
if live is not None: if live is not None:
if not reseed:
return live return live
await self._pool.close(key)
resume = session_id resume = session_id
if resume is not None: if resume is not None:
await self.repair_session(resume) await self.repair_session(resume)
+22 -11
View File
@@ -87,6 +87,12 @@ class ForkOutcome:
messages: list[MessageParam] messages: list[MessageParam]
persist_messages: list[dict[str, Any]] persist_messages: list[dict[str, Any]]
divergence_index: int | None divergence_index: int | None
edited: bool = False
"""An earlier assistant turn's prose was rewritten in the file."""
@property
def reuse_session(self) -> bool:
return self.divergence_index is None and not self.edited
# ---- public store API --------------------------------------------------- # ---- public store API ---------------------------------------------------
@@ -407,7 +413,7 @@ def diff_and_fork(
new_user_turn = incoming[-1] new_user_turn = incoming[-1]
prior_incoming = incoming[:-1] prior_incoming = incoming[:-1]
spliced_groups, divergence = _walk_prefix(prior_incoming, stored_groups) spliced_groups, divergence, edited = _walk_prefix(prior_incoming, stored_groups)
if divergence is None and len(prior_incoming) < len(stored_groups): if divergence is None and len(prior_incoming) < len(stored_groups):
if _file_lags_store(stored_groups, len(prior_incoming), new_user_turn): if _file_lags_store(stored_groups, len(prior_incoming), new_user_turn):
@@ -433,6 +439,7 @@ def diff_and_fork(
messages=backend_msgs, messages=backend_msgs,
persist_messages=persist_msgs, persist_messages=persist_msgs,
divergence_index=divergence, divergence_index=divergence,
edited=edited,
) )
@@ -470,26 +477,29 @@ def _file_lags_store(
def _walk_prefix( def _walk_prefix(
prior_incoming: list[ParsedTurn], stored_groups: list[_StoredDisplayTurn] prior_incoming: list[ParsedTurn], stored_groups: list[_StoredDisplayTurn]
) -> tuple[list[list[dict[str, Any]]], int | None]: ) -> tuple[list[list[dict[str, Any]]], int | None, bool]:
"""Walk incoming vs stored side-by-side until first divergence. """Walk incoming vs stored side-by-side until first divergence.
Returns the spliced/matched group list (one entry per matched Returns the spliced/matched group list (one entry per matched
display turn, each carrying the raw messages we'll feed to the display turn, each carrying the raw messages we'll feed to the
backend for that turn) and the divergence index (``None`` if all backend for that turn), the divergence index (``None`` if all
of ``prior_incoming`` matched). of ``prior_incoming`` matched) and whether any assistant prose
was spliced in from the file - a rewritten reply keeps the
structure but must not resume the session that said otherwise.
""" """
from beaver_gateway.frontends.markdown.parser import TextSegment, ToolSegment from beaver_gateway.frontends.markdown.parser import TextSegment, ToolSegment
spliced_groups: list[list[dict[str, Any]]] = [] spliced_groups: list[list[dict[str, Any]]] = []
edited = False
for i, inc in enumerate(prior_incoming): for i, inc in enumerate(prior_incoming):
if i >= len(stored_groups): if i >= len(stored_groups):
return spliced_groups, i return spliced_groups, i, edited
st = stored_groups[i] st = stored_groups[i]
if inc.role != st.role: if inc.role != st.role:
return spliced_groups, i return spliced_groups, i, edited
if inc.role == "user": if inc.role == "user":
if inc.text != st.spoken_text: if inc.text != st.spoken_text:
return spliced_groups, i return spliced_groups, i, edited
spliced_groups.append(list(st.messages)) spliced_groups.append(list(st.messages))
continue continue
inc_skeleton = tuple( inc_skeleton = tuple(
@@ -499,17 +509,18 @@ def _walk_prefix(
# Files rendered without tool callouts (§3.10) carry no skeleton: # Files rendered without tool callouts (§3.10) carry no skeleton:
# prose alone decides whether the turn matched. # prose alone decides whether the turn matched.
if inc_skeleton and inc_skeleton != st.skeleton: if inc_skeleton and inc_skeleton != st.skeleton:
return spliced_groups, i return spliced_groups, i, edited
if inc.text == st.spoken_text: if inc.text == st.spoken_text:
spliced_groups.append(list(st.messages)) spliced_groups.append(list(st.messages))
continue continue
if inc_skeleton and inc_text_count != st.text_segment_count: if inc_skeleton and inc_text_count != st.text_segment_count:
return spliced_groups, i return spliced_groups, i, edited
spliced = _splice_assistant_group(stored_group=st, incoming=inc) spliced = _splice_assistant_group(stored_group=st, incoming=inc)
if spliced is None: if spliced is None:
return spliced_groups, i return spliced_groups, i, edited
spliced_groups.append(spliced) spliced_groups.append(spliced)
return spliced_groups, None edited = True
return spliced_groups, None, edited
def _assemble_tail( def _assemble_tail(
+1
View File
@@ -877,6 +877,7 @@ class Conversations:
messages=messages, messages=messages,
conversation_id=conv.external_id, conversation_id=conv.external_id,
session_id=resume if use_session else None, session_id=resume if use_session else None,
reseed=not use_session,
capture=capture, capture=capture,
kind=conv.kind, kind=conv.kind,
pinned=conv.kind == "master", pinned=conv.kind == "master",
@@ -400,7 +400,7 @@ class MarkdownFrontend(Frontend):
messages=outcome.messages, messages=outcome.messages,
origin="user", origin="user",
capture=capture, capture=capture,
use_session=outcome.divergence_index is None, use_session=outcome.reuse_session,
) )
try: try:
message = await self._stream_to_file( message = await self._stream_to_file(
@@ -623,7 +623,7 @@ class MarkdownFrontend(Frontend):
messages=outcome.messages, messages=outcome.messages,
origin="user", origin="user",
capture=capture, capture=capture,
use_session=outcome.divergence_index is None, use_session=outcome.reuse_session,
) )
acc = StreamAccumulator() acc = StreamAccumulator()
+27
View File
@@ -228,3 +228,30 @@ async def test_anthropic_turns_become_one_deep_conversation(stack: Stack) -> Non
headers=AUTH, headers=AUTH,
) )
assert r.status_code == 400 and "runs on 'd'" in r.json()["detail"] assert r.status_code == 400 and "runs on 'd'" in r.json()["detail"]
async def test_markdown_edited_reply_reseeds_the_session(stack: Stack) -> None:
async with stack.client(stack.markdown) as c:
r = await c.post(
"/chat",
json={"filename": "e.md", "agent": "d", "content": "### User:\n\nanimal\n"},
headers=AUTH,
)
assert r.status_code == 200, r.text
content = r.json()["new_content"]
assert "ok:animal" in content and len(ScriptedClient.instances) == 1
edited = content.replace("ok:animal", "penguin") + "\nwhich one?\n"
r = await c.post(
"/chat",
json={"filename": "e.md", "agent": "d", "content": edited},
headers=AUTH,
)
assert r.status_code == 200, r.text
assert len(ScriptedClient.instances) == 2
first, second = ScriptedClient.instances
assert not first.connected and second.prompts == ["which one?"]
assert second.options.resume not in (None, first.session_id)
assert "penguin" in repr(vars(stack.world.store))
async with stack.world.db.session() as session:
stored = await load_messages(session, conversation_id=1)
assert stored[1]["content"] == [{"type": "text", "text": "penguin"}]