Merge pull request #93784 from NousResearch/salv/81234-retry-carrier
fix: /retry and /undo no longer replay an older message after compaction (#81233, salvage #81234)
This commit is contained in:
@@ -771,6 +771,19 @@ def repair_message_sequence(agent, messages: List[Dict]) -> int:
|
||||
and merged[-1].get("role") == "user"
|
||||
):
|
||||
prev = merged[-1]
|
||||
# A summary carrier followed by a new user row is a deliberate
|
||||
# durable shape after retry/rewind. Do not absorb the fresh ask
|
||||
# into the already-persisted carrier: mutating that dict can make
|
||||
# the only in-memory copy diverge from its durable row. Provider
|
||||
# sanitizers merge copies later when strict alternation requires
|
||||
# it, without rewriting either durable message.
|
||||
from agent.context_compressor import split_user_originated_turn
|
||||
|
||||
handoff, _ = split_user_originated_turn(prev)
|
||||
if handoff is not None:
|
||||
merged.append(msg)
|
||||
continue
|
||||
|
||||
prev_content = prev.get("content", "")
|
||||
new_content = msg.get("content", "")
|
||||
# Only merge plain-text content; leave multimodal (list)
|
||||
@@ -1428,8 +1441,6 @@ def drop_thinking_only_and_merge_users(
|
||||
)
|
||||
]
|
||||
dropped = len(messages) - len(kept)
|
||||
if dropped == 0:
|
||||
return messages
|
||||
|
||||
# Pass 2: merge any newly-adjacent user messages.
|
||||
merged: List[Dict[str, Any]] = []
|
||||
@@ -1479,6 +1490,9 @@ def drop_thinking_only_and_merge_users(
|
||||
else:
|
||||
merged.append(m)
|
||||
|
||||
if dropped == 0 and merges == 0:
|
||||
return messages
|
||||
|
||||
_ra().logger.debug(
|
||||
"Pre-call sanitizer: dropped %d thinking-only assistant turn(s), "
|
||||
"merged %d adjacent user message(s)",
|
||||
|
||||
+222
-18
@@ -16,6 +16,7 @@ Improvements over v2:
|
||||
- Richer tool call/result detail in summarizer input
|
||||
"""
|
||||
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
@@ -8103,6 +8104,218 @@ def is_compaction_summary_message(message: Any) -> bool:
|
||||
return ContextCompressor._is_context_summary_content(content)
|
||||
|
||||
|
||||
# Display metadata that describes the durable message independently of the
|
||||
# compaction wrapper. Other metadata may describe a synthetic timeline event
|
||||
# and must not make that event look human after projection.
|
||||
SUMMARY_CARRIER_DURABLE_DISPLAY_METADATA_KEYS = ("reactions",)
|
||||
|
||||
|
||||
def _handoff_only_content(content: Any) -> Any:
|
||||
"""Project summary-bearing content to the synthetic handoff alone.
|
||||
|
||||
The compressor has two composite layouts. Ordinary merge-into-tail keeps
|
||||
the live content before ``_MERGED_SUMMARY_DELIMITER``; the force-user-
|
||||
leading layout keeps it after ``_SUMMARY_END_MARKER``. This is the inverse
|
||||
of ``_strip_context_summary_handoff_message`` and deliberately never keeps
|
||||
live media blocks.
|
||||
"""
|
||||
if isinstance(content, str):
|
||||
if _MERGED_SUMMARY_DELIMITER in content:
|
||||
suffix = content.split(_MERGED_SUMMARY_DELIMITER, 1)[1].lstrip()
|
||||
marker_idx = suffix.find(_SUMMARY_END_MARKER)
|
||||
if marker_idx >= 0:
|
||||
return suffix[: marker_idx + len(_SUMMARY_END_MARKER)]
|
||||
return suffix
|
||||
marker_idx = content.find(_SUMMARY_END_MARKER)
|
||||
if marker_idx >= 0:
|
||||
return content[: marker_idx + len(_SUMMARY_END_MARKER)]
|
||||
return content
|
||||
|
||||
if not isinstance(content, list):
|
||||
return content
|
||||
|
||||
# Ordinary merge: the summary suffix begins in the delimiter-bearing text
|
||||
# part. Do not retain later parts: malformed/legacy rows may carry live
|
||||
# media there rather than synthetic scaffold content.
|
||||
for item in content:
|
||||
text = (
|
||||
item
|
||||
if isinstance(item, str)
|
||||
else item.get("text")
|
||||
if isinstance(item, dict)
|
||||
else None
|
||||
)
|
||||
if not isinstance(text, str) or _MERGED_SUMMARY_DELIMITER not in text:
|
||||
continue
|
||||
suffix = text.split(_MERGED_SUMMARY_DELIMITER, 1)[1].lstrip()
|
||||
marker_idx = suffix.find(_SUMMARY_END_MARKER)
|
||||
if marker_idx >= 0:
|
||||
suffix = suffix[: marker_idx + len(_SUMMARY_END_MARKER)]
|
||||
if not suffix:
|
||||
return []
|
||||
if isinstance(item, dict):
|
||||
copied = item.copy()
|
||||
copied["text"] = suffix
|
||||
return [copied]
|
||||
return [suffix]
|
||||
|
||||
# Force-user-leading merge: keep textual parts through the end marker and
|
||||
# truncate the marker-bearing part before the live ask.
|
||||
projected: list[Any] = []
|
||||
for item in content:
|
||||
text = (
|
||||
item
|
||||
if isinstance(item, str)
|
||||
else item.get("text")
|
||||
if isinstance(item, dict)
|
||||
else None
|
||||
)
|
||||
if isinstance(text, str) and _SUMMARY_END_MARKER in text:
|
||||
prefix = text.split(_SUMMARY_END_MARKER, 1)[0] + _SUMMARY_END_MARKER
|
||||
if isinstance(item, dict):
|
||||
copied = item.copy()
|
||||
copied["text"] = prefix
|
||||
projected.append(copied)
|
||||
else:
|
||||
projected.append(prefix)
|
||||
return projected
|
||||
if isinstance(text, str):
|
||||
projected.append(item.copy() if isinstance(item, dict) else item)
|
||||
return projected
|
||||
|
||||
|
||||
def split_user_originated_turn(
|
||||
message: Any,
|
||||
) -> tuple[Optional[Dict[str, Any]], Optional[Dict[str, Any]]]:
|
||||
"""Split a user row into hidden handoff scaffold and canonical live view.
|
||||
|
||||
Returns ``(handoff_only, live_view)``. A normal human row has no
|
||||
handoff; a pure compaction handoff has no live view; a composite carrier
|
||||
has both. Rewritten projections are fresh dictionaries and never retain
|
||||
stale API-content or physical persistence identity.
|
||||
"""
|
||||
if not isinstance(message, dict) or message.get("role") != "user":
|
||||
return None, None
|
||||
|
||||
is_summary = is_compaction_summary_message(message)
|
||||
handoff: Optional[Dict[str, Any]] = None
|
||||
if is_summary:
|
||||
handoff = {
|
||||
"role": "user",
|
||||
"content": _handoff_only_content(message.get("content")),
|
||||
COMPRESSED_SUMMARY_METADATA_KEY: True,
|
||||
"display_kind": "hidden",
|
||||
}
|
||||
if COMPRESSED_SUMMARY_HAS_USER_TURN_KEY in message:
|
||||
handoff[COMPRESSED_SUMMARY_HAS_USER_TURN_KEY] = bool(
|
||||
message.get(COMPRESSED_SUMMARY_HAS_USER_TURN_KEY)
|
||||
)
|
||||
if message.get(MICRO_COMPACT_MARKER_KEY):
|
||||
handoff[MICRO_COMPACT_MARKER_KEY] = True
|
||||
if message.get("timestamp") is not None:
|
||||
handoff["timestamp"] = message["timestamp"]
|
||||
drop_stale_api_content(handoff)
|
||||
|
||||
# Hidden is the legacy physical wrapper used for compaction rows and
|
||||
# does not hide a successfully unwrapped human payload. Other typed
|
||||
# display rows are synthetic timeline events, never user input.
|
||||
display_kind = message.get("display_kind")
|
||||
if display_kind and display_kind != "hidden":
|
||||
return handoff, None
|
||||
candidate = ContextCompressor._strip_context_summary_handoff_message(message)
|
||||
if candidate is None:
|
||||
return handoff, None
|
||||
else:
|
||||
if message.get("display_kind"):
|
||||
return None, None
|
||||
candidate = message.copy()
|
||||
|
||||
candidate.pop(COMPRESSED_SUMMARY_METADATA_KEY, None)
|
||||
candidate.pop(COMPRESSED_SUMMARY_HAS_USER_TURN_KEY, None)
|
||||
candidate.pop(MICRO_COMPACT_MARKER_KEY, None)
|
||||
candidate.pop(_DB_PERSISTED_MARKER, None)
|
||||
if is_summary:
|
||||
candidate.pop("_row_id", None)
|
||||
candidate.pop("display_kind", None)
|
||||
candidate.pop("display_metadata", None)
|
||||
carrier_metadata = message.get("display_metadata")
|
||||
if isinstance(carrier_metadata, dict):
|
||||
durable_metadata = {
|
||||
key: copy.deepcopy(carrier_metadata[key])
|
||||
for key in SUMMARY_CARRIER_DURABLE_DISPLAY_METADATA_KEYS
|
||||
if key in carrier_metadata
|
||||
}
|
||||
if durable_metadata:
|
||||
candidate["display_metadata"] = durable_metadata
|
||||
drop_stale_api_content(candidate)
|
||||
if ContextCompressor._is_synthetic_compression_user_turn(candidate):
|
||||
return handoff, None
|
||||
if not ContextCompressor._is_actionable_user_turn(candidate):
|
||||
return handoff, None
|
||||
return handoff, candidate
|
||||
|
||||
|
||||
def user_originated_turn_view(message: Any) -> Optional[Dict[str, Any]]:
|
||||
"""Return the live human-authored projection of a user row, if any."""
|
||||
return split_user_originated_turn(message)[1]
|
||||
|
||||
|
||||
def history_before_user_originated_turn(
|
||||
messages: List[Dict[str, Any]],
|
||||
index: int,
|
||||
) -> tuple[List[Dict[str, Any]], Dict[str, Any]]:
|
||||
"""Return a rewind prefix and canonical live view for ``index``.
|
||||
|
||||
When the selected row is a composite carrier, the hidden handoff scaffold
|
||||
remains at the new history head while the live ask and later rows are
|
||||
removed. This retains the only representation of already-compacted turns.
|
||||
"""
|
||||
if index < 0 or index >= len(messages):
|
||||
raise IndexError("user turn index is outside the transcript")
|
||||
handoff, live_view = split_user_originated_turn(messages[index])
|
||||
if live_view is None:
|
||||
raise ValueError("selected row is not a user-originated turn")
|
||||
prefix = [message.copy() for message in messages[:index]]
|
||||
if handoff is not None:
|
||||
prefix.append(handoff)
|
||||
return prefix, live_view
|
||||
|
||||
|
||||
def retryable_user_text(content: Any) -> str:
|
||||
"""Return lossless retry text or raise before destructive mutation.
|
||||
|
||||
Retry has no attachment replay protocol. Media and unknown structured
|
||||
parts therefore fail closed; already-persisted strings are replayed as
|
||||
text, including any textual degradation labels. Structured content is
|
||||
flattened only when every part is plain text.
|
||||
"""
|
||||
if isinstance(content, str):
|
||||
text = content
|
||||
elif isinstance(content, list):
|
||||
chunks: list[str] = []
|
||||
for part in content:
|
||||
if isinstance(part, str):
|
||||
chunks.append(part)
|
||||
continue
|
||||
if not isinstance(part, dict):
|
||||
raise ValueError("retry does not support non-text content")
|
||||
if part.get("type") not in {"text", "input_text", "output_text"}:
|
||||
raise ValueError("retry does not support media or unknown content parts")
|
||||
if set(part) - {"type", "text"}:
|
||||
raise ValueError("retry cannot losslessly flatten annotated text parts")
|
||||
part_text = part.get("text")
|
||||
if not isinstance(part_text, str):
|
||||
raise ValueError("retry text parts must contain text")
|
||||
chunks.append(part_text)
|
||||
text = "".join(chunks)
|
||||
else:
|
||||
raise ValueError("retry does not support non-text content")
|
||||
|
||||
if not text.strip():
|
||||
raise ValueError("retry found no text to send")
|
||||
return text
|
||||
|
||||
|
||||
def _handoff_carries_live_user_content(message: Any) -> bool:
|
||||
"""Return True when a summary-bearing row still carries a live user ask.
|
||||
|
||||
@@ -8112,13 +8325,11 @@ def _handoff_carries_live_user_content(message: Any) -> bool:
|
||||
shape must remain actionable (#80622 must not treat them as sole-handoff).
|
||||
|
||||
Delegates to ``_strip_context_summary_handoff_message`` — the canonical
|
||||
"does anything survive once the handoff is removed" logic (it also
|
||||
handles multimodal list content and returns ``None`` for a merged-shaped
|
||||
row whose preserved prior tail is EMPTY, which a bare
|
||||
``classify_summary_content(...) == "merged"`` check would wrongly treat
|
||||
as live). Callers must pre-filter with ``is_compaction_summary_message``:
|
||||
for non-summary rows the strip helper returns the message unchanged,
|
||||
which would read as "carries live content" here.
|
||||
"does anything survive once the handoff is removed" logic. This helper
|
||||
also applies to merged assistant carriers whose pending tool calls keep an
|
||||
exchange in flight, so it must not use the user-row-only display projection.
|
||||
Callers must pre-filter with ``is_compaction_summary_message`` because a
|
||||
non-summary row is returned unchanged by the strip helper.
|
||||
"""
|
||||
if not isinstance(message, dict):
|
||||
return False
|
||||
@@ -8197,15 +8408,8 @@ def is_user_originated_turn(message: Any) -> bool:
|
||||
this instead of ``role == "user" and not display_kind`` — standalone
|
||||
handoffs with ``_compressed_summary_has_user_turn`` were previously left
|
||||
without ``display_kind=hidden`` and could be mistaken for real asks (#80622).
|
||||
Summary-bearing rows are never user-originated, even when they embed a
|
||||
live ask after the end marker (callers that need that text should unwrap).
|
||||
Summary-bearing rows count only when their canonical live-user projection
|
||||
recovers an actionable ask. Pure handoffs and typed synthetic rows never
|
||||
count.
|
||||
"""
|
||||
if not isinstance(message, dict) or message.get("role") != "user":
|
||||
return False
|
||||
if message.get("display_kind"):
|
||||
return False
|
||||
if is_compaction_summary_message(message):
|
||||
return False
|
||||
if ContextCompressor._is_synthetic_compression_user_turn(message):
|
||||
return False
|
||||
return ContextCompressor._is_actionable_user_turn(message)
|
||||
return user_originated_turn_view(message) is not None
|
||||
|
||||
+10
-2
@@ -274,18 +274,26 @@ def reanchor_current_turn_user_idx(messages: List[Any], user_message: Any) -> in
|
||||
scaffolding, not the active ask. Returns -1 when the list has no
|
||||
user-originated message at all.
|
||||
"""
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
|
||||
fallback = -1
|
||||
for i in range(len(messages) - 1, -1, -1):
|
||||
msg = messages[i]
|
||||
if not (isinstance(msg, dict) and msg.get("role") == "user"):
|
||||
continue
|
||||
# Typed synthetic current events still need their physical persistence
|
||||
# anchor when their raw content is unchanged. They are not eligible
|
||||
# for the human-only fallback below.
|
||||
if msg.get("content") == user_message:
|
||||
return i
|
||||
live_view = user_originated_turn_view(msg)
|
||||
if live_view is None:
|
||||
continue
|
||||
if live_view.get("content") == user_message:
|
||||
return i
|
||||
# Prefer a real human turn over a synthetic handoff / continuation
|
||||
# marker when the exact content was rewritten by merge-into-tail.
|
||||
if fallback < 0 and is_user_originated_turn(msg):
|
||||
if fallback < 0:
|
||||
fallback = i
|
||||
return fallback
|
||||
|
||||
|
||||
@@ -37,6 +37,7 @@ import {
|
||||
applyBranchVisibility,
|
||||
applyReloadOptimistic,
|
||||
applyRewindOptimistic,
|
||||
durableRowIdsForRebind,
|
||||
finalizeInterruptedMessages,
|
||||
planEdit,
|
||||
planReload,
|
||||
@@ -415,7 +416,8 @@ export function useSessionTileActions({ requestGateway, runtimeId, scope, stored
|
||||
interruptFirst: boolean,
|
||||
truncateMessageId?: string,
|
||||
truncateRowId?: number,
|
||||
sourceText?: string
|
||||
sourceText?: string,
|
||||
rebindRowIds?: readonly number[]
|
||||
) =>
|
||||
runRewindSubmit(
|
||||
requestGateway,
|
||||
@@ -429,7 +431,8 @@ export function useSessionTileActions({ requestGateway, runtimeId, scope, stored
|
||||
onSessionRecovered: bindRecoveredRuntime
|
||||
},
|
||||
truncateRowId,
|
||||
sourceText
|
||||
sourceText,
|
||||
rebindRowIds
|
||||
),
|
||||
[bindRecoveredRuntime, requestGateway]
|
||||
)
|
||||
@@ -477,7 +480,8 @@ export function useSessionTileActions({ requestGateway, runtimeId, scope, stored
|
||||
false,
|
||||
plan.truncateMessageId,
|
||||
plan.truncateRowId,
|
||||
plan.sourceText
|
||||
plan.sourceText,
|
||||
durableRowIdsForRebind(state.messages)
|
||||
)
|
||||
)
|
||||
} catch (err) {
|
||||
@@ -513,7 +517,8 @@ export function useSessionTileActions({ requestGateway, runtimeId, scope, stored
|
||||
interruptFirst,
|
||||
plan.truncateMessageId,
|
||||
plan.truncateRowId,
|
||||
plan.sourceText
|
||||
plan.sourceText,
|
||||
durableRowIdsForRebind(messages)
|
||||
)
|
||||
)
|
||||
} catch (err) {
|
||||
@@ -561,7 +566,8 @@ export function useSessionTileActions({ requestGateway, runtimeId, scope, stored
|
||||
interruptFirst,
|
||||
plan.truncateMessageId,
|
||||
plan.truncateRowId,
|
||||
plan.sourceText
|
||||
plan.sourceText,
|
||||
durableRowIdsForRebind(messages)
|
||||
)
|
||||
)
|
||||
} catch (err) {
|
||||
|
||||
@@ -57,6 +57,7 @@ import {
|
||||
applyBranchVisibility,
|
||||
applyReloadOptimistic,
|
||||
applyRewindOptimistic,
|
||||
durableRowIdsForRebind,
|
||||
finalizeInterruptedMessages,
|
||||
planEdit,
|
||||
planReload,
|
||||
@@ -847,7 +848,8 @@ export function usePromptActions({
|
||||
truncateMessageId: string | undefined,
|
||||
interruptFirst: boolean,
|
||||
truncateRowId?: number,
|
||||
sourceText?: string
|
||||
sourceText?: string,
|
||||
rebindRowIds?: readonly number[]
|
||||
) =>
|
||||
runRewindSubmit(
|
||||
requestGateway,
|
||||
@@ -864,7 +866,8 @@ export function usePromptActions({
|
||||
}
|
||||
},
|
||||
truncateRowId,
|
||||
sourceText
|
||||
sourceText,
|
||||
rebindRowIds
|
||||
),
|
||||
[activeSessionIdRef, requestGateway, selectedStoredSessionIdRef]
|
||||
)
|
||||
@@ -879,7 +882,8 @@ export function usePromptActions({
|
||||
return
|
||||
}
|
||||
|
||||
const plan = planReload($messages.get(), parentId)
|
||||
const messages = $messages.get()
|
||||
const plan = planReload(messages, parentId)
|
||||
|
||||
if (!plan) {
|
||||
return
|
||||
@@ -896,7 +900,8 @@ export function usePromptActions({
|
||||
plan.truncateMessageId,
|
||||
false,
|
||||
plan.truncateRowId,
|
||||
plan.sourceText
|
||||
plan.sourceText,
|
||||
durableRowIdsForRebind(messages)
|
||||
)
|
||||
|
||||
applySurvivorRowIds(sessionId, survivorRowIds)
|
||||
@@ -964,7 +969,8 @@ export function usePromptActions({
|
||||
plan.truncateMessageId,
|
||||
interruptFirst,
|
||||
plan.truncateRowId,
|
||||
plan.sourceText
|
||||
plan.sourceText,
|
||||
durableRowIdsForRebind(messages)
|
||||
)
|
||||
|
||||
applySurvivorRowIds(sessionId, survivorRowIds)
|
||||
@@ -1049,7 +1055,8 @@ export function usePromptActions({
|
||||
plan.truncateMessageId,
|
||||
interruptFirst,
|
||||
plan.truncateRowId,
|
||||
plan.sourceText
|
||||
plan.sourceText,
|
||||
durableRowIdsForRebind(messages)
|
||||
)
|
||||
|
||||
applySurvivorRowIds(sessionId, survivorRowIds)
|
||||
@@ -1082,7 +1089,8 @@ export function usePromptActions({
|
||||
retryPlan.truncateMessageId,
|
||||
false,
|
||||
retryPlan.truncateRowId,
|
||||
retryPlan.sourceText
|
||||
retryPlan.sourceText,
|
||||
durableRowIdsForRebind(refreshed)
|
||||
)
|
||||
|
||||
applySurvivorRowIds(sessionId, survivorRowIds)
|
||||
|
||||
@@ -163,6 +163,21 @@ describe('survivorRowIdsFrom', () => {
|
||||
it('keeps integer ids and nulls anything else', () => {
|
||||
expect(survivorRowIdsFrom({ survivor_user_row_ids: [7, null, 9.5, '11', 12] })).toEqual([7, null, null, null, 12])
|
||||
})
|
||||
|
||||
it('prefers the explicit old-to-new row id map', () => {
|
||||
const parsed = survivorRowIdsFrom({
|
||||
survivor_user_row_ids: [100, 200],
|
||||
survivor_row_id_map: { '11': 101, '21': 201, '31': null }
|
||||
})
|
||||
|
||||
expect(parsed).toEqual(
|
||||
new Map([
|
||||
[11, 101],
|
||||
[21, 201],
|
||||
[31, null]
|
||||
])
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
describe('rebindSurvivorRowIds', () => {
|
||||
@@ -217,6 +232,30 @@ describe('rebindSurvivorRowIds', () => {
|
||||
|
||||
expect(rebindSurvivorRowIds(messages, [7])[0]).toBe(messages[0])
|
||||
})
|
||||
|
||||
it('rebinds a paged tail by previous row id instead of global position', () => {
|
||||
const messages = [
|
||||
user('archived-u250', 800),
|
||||
user('tail-u298', 901),
|
||||
assistant('tail-a298', 902),
|
||||
user('edited', 903)
|
||||
]
|
||||
|
||||
const rebound = rebindSurvivorRowIds(
|
||||
messages,
|
||||
new Map([
|
||||
[901, 1001],
|
||||
[902, 1002],
|
||||
[903, null]
|
||||
])
|
||||
)
|
||||
|
||||
expect(rebound[0]).toBe(messages[0])
|
||||
expect(rebound[0].rowId).toBe(800)
|
||||
expect(rebound[1].rowId).toBe(1001)
|
||||
expect(rebound[2].rowId).toBe(1002)
|
||||
expect(rebound[3].rowId).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe('finalizeInterruptedMessages', () => {
|
||||
@@ -478,7 +517,18 @@ describe('runRewindSubmit durable-address discipline (#87059)', () => {
|
||||
it('leaves a bound durable rowId untouched (no extra history call) and drops the client ordinal', async () => {
|
||||
const calls: Call[] = []
|
||||
|
||||
await runRewindSubmit(makeGateway(calls), 'sid', 'fixed prompt', 1, undefined, false, undefined, 13, 'typo prompt')
|
||||
await runRewindSubmit(
|
||||
makeGateway(calls),
|
||||
'sid',
|
||||
'fixed prompt',
|
||||
1,
|
||||
undefined,
|
||||
false,
|
||||
undefined,
|
||||
13,
|
||||
'typo prompt',
|
||||
[11, 12, 13]
|
||||
)
|
||||
|
||||
expect(calls.some(call => call.method === 'session.history')).toBe(false)
|
||||
|
||||
@@ -486,6 +536,8 @@ describe('runRewindSubmit durable-address discipline (#87059)', () => {
|
||||
|
||||
expect(submit?.params?.truncate_before_row_id).toBe(13)
|
||||
expect(submit?.params?.truncate_before_user_ordinal).toBeUndefined()
|
||||
expect(submit?.params?.confirm_empty_truncate).toBe(true)
|
||||
expect(submit?.params?.rebind_survivor_row_ids).toEqual([11, 12, 13])
|
||||
})
|
||||
|
||||
it('drops the client ordinal whenever a durable row id is present, including a complete live transcript (#88082, #89244)', async () => {
|
||||
|
||||
@@ -35,24 +35,41 @@ import {
|
||||
type RequestGateway = <T = unknown>(method: string, params?: Record<string, unknown>, timeoutMs?: number) => Promise<T>
|
||||
|
||||
/**
|
||||
* Post-rewrite durable ids of the surviving visible user turns, in visible-user
|
||||
* ordinal order — the gateway's `survivor_user_row_ids` on a truncating
|
||||
* `prompt.submit`. A rewind's `replace_messages` re-inserts the kept prefix as
|
||||
* NEW SQLite rows, so every pre-rewind `ChatMessage.rowId` on a surviving
|
||||
* bubble is stale the moment the rewind lands; targeting one on the next
|
||||
* rewind/edit/regenerate gets a fail-closed 4018 from the gateway. `null`
|
||||
* means that turn has no durable id (drop the cached one, don't keep a stale
|
||||
* one). Absent entirely = the submit didn't truncate a durable session (or an
|
||||
* older gateway) — leave state untouched.
|
||||
* Post-rewrite durable identity information from a truncating `prompt.submit`.
|
||||
* New gateways return an old-to-new map for every surviving physical row;
|
||||
* older gateways return visible-user ids in ordinal order. A rewind's
|
||||
* `replace_messages` re-inserts the kept prefix as NEW SQLite rows, so every
|
||||
* cached rowId is stale the moment the rewind lands.
|
||||
*/
|
||||
export type SurvivorUserRowIds = readonly (null | number)[]
|
||||
export type SurvivorUserRowIds = readonly (null | number)[] | Map<number, null | number>
|
||||
|
||||
interface PromptSubmitResult {
|
||||
status?: string
|
||||
survivor_user_row_ids?: unknown
|
||||
survivor_row_id_map?: unknown
|
||||
}
|
||||
|
||||
export function survivorRowIdsFrom(result: PromptSubmitResult | undefined): SurvivorUserRowIds | undefined {
|
||||
const rawMap = result?.survivor_row_id_map
|
||||
|
||||
if (rawMap && typeof rawMap === 'object' && !Array.isArray(rawMap)) {
|
||||
const parsed = new Map<number, null | number>()
|
||||
|
||||
for (const [oldId, newId] of Object.entries(rawMap)) {
|
||||
const previous = Number(oldId)
|
||||
|
||||
if (Number.isInteger(previous)) {
|
||||
if (newId === null) {
|
||||
parsed.set(previous, null)
|
||||
} else if (typeof newId === 'number' && Number.isInteger(newId)) {
|
||||
parsed.set(previous, newId)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return parsed
|
||||
}
|
||||
|
||||
const raw = result?.survivor_user_row_ids
|
||||
|
||||
if (!Array.isArray(raw)) {
|
||||
@@ -71,6 +88,28 @@ export function survivorRowIdsFrom(result: PromptSubmitResult | undefined): Surv
|
||||
* stale id now addresses an archived row and would be refused with 4018.
|
||||
*/
|
||||
export function rebindSurvivorRowIds(messages: ChatMessage[], survivorRowIds: SurvivorUserRowIds): ChatMessage[] {
|
||||
if (survivorRowIds instanceof Map) {
|
||||
return messages.map(message => {
|
||||
if (message.rowId === undefined) {
|
||||
return message
|
||||
}
|
||||
|
||||
if (!survivorRowIds.has(message.rowId)) {
|
||||
// Requested rows outside the active tip (compacted/ancestor history)
|
||||
// were not rewritten and keep their durable identity.
|
||||
return message
|
||||
}
|
||||
|
||||
const next = survivorRowIds.get(message.rowId)
|
||||
|
||||
return next === null || next === undefined
|
||||
? { ...message, rowId: undefined }
|
||||
: message.rowId === next
|
||||
? message
|
||||
: { ...message, rowId: next }
|
||||
})
|
||||
}
|
||||
|
||||
// Same ordinal space as the truncate math: visible AND persisted (failed
|
||||
// turns never reached the gateway, so they hold no survivor slot).
|
||||
const indices = new Set(visibleUserMessageIndices(messages))
|
||||
@@ -92,6 +131,16 @@ export function rebindSurvivorRowIds(messages: ChatMessage[], survivorRowIds: Su
|
||||
})
|
||||
}
|
||||
|
||||
export function durableRowIdsForRebind(messages: readonly ChatMessage[]): number[] {
|
||||
return [
|
||||
...new Set(
|
||||
messages.flatMap(message =>
|
||||
typeof message.rowId === 'number' && Number.isInteger(message.rowId) ? [message.rowId] : []
|
||||
)
|
||||
)
|
||||
]
|
||||
}
|
||||
|
||||
/**
|
||||
* Renderer-synthetic message ids (`${timestamp}-${index}-${role}` from
|
||||
* chat-messages/hydration.ts, plus older `user-…` / `assistant-…` shapes). Gateway
|
||||
@@ -223,7 +272,8 @@ export async function runRewindSubmit(
|
||||
interruptFirst: boolean,
|
||||
recovery?: { storedSessionId?: null | string; onSessionRecovered?: (sessionId: string) => void },
|
||||
truncateRowId?: number,
|
||||
sourceText?: string
|
||||
sourceText?: string,
|
||||
rebindRowIds?: readonly number[]
|
||||
): Promise<SurvivorUserRowIds | undefined> {
|
||||
// Recovery may rebind the live id mid-flight; interrupt/submit must both
|
||||
// follow it rather than pinning the dead one.
|
||||
@@ -248,6 +298,18 @@ export async function runRewindSubmit(
|
||||
typeof truncateRowId === 'number' ||
|
||||
(typeof truncateMessageId === 'string' && truncateMessageId.length > 0 && !isSyntheticRendererId(truncateMessageId))
|
||||
|
||||
if (wantsTruncation && hasDurableAddress) {
|
||||
// A durable row or platform-message id is globally stable, while the
|
||||
// rendered ordinal can be relative to a paged transcript tail. Do not
|
||||
// cross-check those different spaces; the gateway still validates the
|
||||
// durable address against its full transcript.
|
||||
resolvedOrdinal = undefined
|
||||
|
||||
if (typeof resolvedRowId === 'number' && Number.isInteger(resolvedRowId)) {
|
||||
resolvedMessageId = undefined
|
||||
}
|
||||
}
|
||||
|
||||
if (wantsTruncation && !hasDurableAddress) {
|
||||
resolvedRowId =
|
||||
sourceText === undefined
|
||||
@@ -262,18 +324,6 @@ export async function runRewindSubmit(
|
||||
resolvedMessageId = undefined
|
||||
}
|
||||
|
||||
// Durable id present (#88082, #89244): drop the client ordinal. Renderer
|
||||
// ordinals and gateway tip ordinals are different spaces — tail-only
|
||||
// prefetch (window-relative), in-place compact (display-lineage scrollback
|
||||
// vs tip, prefix_user_count structurally 0), and any later visibility
|
||||
// skew. The cut is aimed by the resolved durable id; sending a divergent
|
||||
// ordinal trips the gateway's 4030 cross-check. Unknown ids still fail
|
||||
// closed at 4018. The ordinal is only a tripwire, and it is calibrated in
|
||||
// a space the gateway cannot reliably compute.
|
||||
if (wantsTruncation && hasDurableAddress) {
|
||||
resolvedOrdinal = undefined
|
||||
}
|
||||
|
||||
const interrupt = async () => {
|
||||
try {
|
||||
await requestGateway('session.interrupt', { session_id: liveSessionId })
|
||||
@@ -290,11 +340,17 @@ export async function runRewindSubmit(
|
||||
text,
|
||||
...truncateSubmitParams(resolvedOrdinal, resolvedMessageId, resolvedRowId),
|
||||
// A first-turn rewind resolves to an empty transcript, which the
|
||||
// gateway additionally gates behind confirm_empty_truncate. When the
|
||||
// client ordinal is dropped (durable-id path), carry the flag from
|
||||
// the caller's ordinal-0 belief: required when right, ignored by the
|
||||
// gateway when the cut isn't actually empty.
|
||||
...(resolvedOrdinal === undefined && truncateOrdinal === 0 ? { confirm_empty_truncate: true } : {})
|
||||
// gateway additionally gates behind confirm_empty_truncate. In
|
||||
// resolved-row-id mode the tail-local ordinal was dropped (see
|
||||
// above). The durable target is authoritative, so explicitly allow a
|
||||
// first-active-tip cut; the gateway ignores this when the prefix is
|
||||
// non-empty.
|
||||
...(resolvedOrdinal === undefined && (resolvedRowId !== undefined || truncateOrdinal === 0)
|
||||
? { confirm_empty_truncate: true }
|
||||
: {}),
|
||||
...(rebindRowIds?.length
|
||||
? { rebind_survivor_row_ids: [...new Set(rebindRowIds.filter(Number.isInteger))] }
|
||||
: {})
|
||||
},
|
||||
PROMPT_SUBMIT_REQUEST_TIMEOUT_MS
|
||||
)
|
||||
|
||||
@@ -365,6 +365,28 @@ describe('toChatMessages', () => {
|
||||
}
|
||||
})
|
||||
|
||||
it('projects persisted composite compaction carriers to their live user turn', () => {
|
||||
const messages = toChatMessages([
|
||||
{
|
||||
id: 71,
|
||||
role: 'user',
|
||||
content: 'internal summary scaffold\n\nREAL ASK',
|
||||
display_content: 'REAL ASK',
|
||||
timestamp: 1
|
||||
},
|
||||
{
|
||||
id: 72,
|
||||
role: 'user',
|
||||
content: 'prior live ask\n\ninternal summary scaffold',
|
||||
display_content: 'prior live ask',
|
||||
timestamp: 2
|
||||
}
|
||||
])
|
||||
|
||||
expect(messages.map(chatMessageText)).toEqual(['REAL ASK', 'prior live ask'])
|
||||
expect(messages.map(message => message.rowId)).toEqual([71, 72])
|
||||
})
|
||||
|
||||
it('projects durable timeline kinds without inspecting their text', () => {
|
||||
const messages = toChatMessages([
|
||||
{ role: 'user', content: 'real user turn', timestamp: 1 },
|
||||
|
||||
@@ -189,7 +189,10 @@ export function toChatMessages(messages: SessionMessage[]): ChatMessage[] {
|
||||
return
|
||||
}
|
||||
|
||||
const content = message.content || message.text || message.context || message.name
|
||||
const content =
|
||||
message.display_content !== undefined
|
||||
? message.display_content
|
||||
: message.content || message.text || message.context || message.name
|
||||
|
||||
const rawDisplayContent = transcriptContent(
|
||||
message.display_kind,
|
||||
|
||||
@@ -573,6 +573,8 @@ export interface SessionMessage {
|
||||
args?: unknown
|
||||
codex_reasoning_items?: unknown
|
||||
content: unknown
|
||||
/** Backend-projected user-visible content when a physical row also carries internal model scaffolding. */
|
||||
display_content?: unknown
|
||||
context?: unknown
|
||||
name?: string
|
||||
reasoning?: null | string
|
||||
|
||||
@@ -10174,6 +10174,118 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
print(f" Resume the live session with: hermes --resume {self.session_id}")
|
||||
except Exception as e:
|
||||
print(f"(x_x) Failed to save: {e}")
|
||||
|
||||
def _rewind_persisted_user_turn(
|
||||
self,
|
||||
*,
|
||||
warm_history: List[Dict[str, Any]],
|
||||
user_ordinal: int,
|
||||
warm_live_view: Dict[str, Any],
|
||||
) -> tuple[List[Dict[str, Any]], Dict[str, Any], Dict[str, Any]]:
|
||||
"""Bind one warm user ordinal to a durable row and rewind it atomically."""
|
||||
if self._session_db is None or not self.session_id:
|
||||
raise RuntimeError("session database is unavailable")
|
||||
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
split_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from agent.memory_manager import sanitize_context
|
||||
from agent.tool_dispatch_helpers import (
|
||||
_is_multimodal_tool_result,
|
||||
_multimodal_text_summary,
|
||||
)
|
||||
from run_agent import _is_ephemeral_scaffolding
|
||||
|
||||
def _persistence_content(content: Any) -> Any:
|
||||
"""Project warm content exactly as the session DB flush does."""
|
||||
if _is_multimodal_tool_result(content):
|
||||
return _multimodal_text_summary(content)
|
||||
if isinstance(content, list):
|
||||
text_parts = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
text_parts.append(str(part.get("text", "")))
|
||||
elif isinstance(part, dict) and part.get("type") in {
|
||||
"image",
|
||||
"image_url",
|
||||
"input_image",
|
||||
}:
|
||||
text_parts.append("[screenshot]")
|
||||
return "\n".join(text_parts) if text_parts else None
|
||||
return content
|
||||
|
||||
def _comparison_content(message: Dict[str, Any]) -> Any:
|
||||
content = _persistence_content(message.get("content"))
|
||||
if message.get("role") in {"user", "assistant"} and isinstance(
|
||||
content, str
|
||||
):
|
||||
return sanitize_context(content).strip()
|
||||
return content
|
||||
|
||||
expected_active_ids = self._session_db.get_active_message_ids(
|
||||
self.session_id
|
||||
)
|
||||
durable = self._session_db.get_messages_as_conversation(
|
||||
self.session_id,
|
||||
include_row_ids=True,
|
||||
)
|
||||
warm_persistence_history = [
|
||||
message
|
||||
for message in warm_history
|
||||
if not _is_ephemeral_scaffolding(message)
|
||||
]
|
||||
warm_user_indices = [
|
||||
index
|
||||
for index, message in enumerate(warm_persistence_history)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
durable_user_indices = [
|
||||
index
|
||||
for index, message in enumerate(durable)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if len(durable_user_indices) != len(warm_user_indices):
|
||||
raise RuntimeError(
|
||||
"session history changed before the rewind could be persisted"
|
||||
)
|
||||
if user_ordinal < 0 or user_ordinal >= len(durable_user_indices):
|
||||
raise RuntimeError("persisted rewind target is no longer available")
|
||||
|
||||
warm_prefix, _ = history_before_user_originated_turn(
|
||||
warm_persistence_history, warm_user_indices[user_ordinal]
|
||||
)
|
||||
durable_target_index = durable_user_indices[user_ordinal]
|
||||
durable_target = durable[durable_target_index]
|
||||
durable_prefix, durable_live_view = history_before_user_originated_turn(
|
||||
durable, durable_target_index
|
||||
)
|
||||
if _comparison_content(durable_live_view) != _comparison_content(
|
||||
warm_live_view
|
||||
):
|
||||
raise RuntimeError(
|
||||
"session history changed before the rewind could be persisted"
|
||||
)
|
||||
target_row_id = durable_target.get("_row_id")
|
||||
if not isinstance(target_row_id, int):
|
||||
raise RuntimeError("persisted rewind target has no row identity")
|
||||
scaffold, _ = split_user_originated_turn(durable_target)
|
||||
result = self._session_db.rewind_to_message(
|
||||
self.session_id,
|
||||
target_row_id,
|
||||
preserve_compaction_handoff=scaffold is not None,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content=durable_live_view.get("content"),
|
||||
)
|
||||
if scaffold is not None:
|
||||
replacement_id = result.get("replacement_message_id")
|
||||
if not isinstance(replacement_id, int) or not durable_prefix:
|
||||
raise RuntimeError("rewind did not retain its compaction handoff")
|
||||
durable_prefix[-1]["_row_id"] = replacement_id
|
||||
durable_prefix[-1]["_db_persisted"] = True
|
||||
warm_prefix[-1] = durable_prefix[-1]
|
||||
return warm_prefix, durable_live_view, result
|
||||
|
||||
def retry_last(self):
|
||||
"""Retry the last user message by removing the last exchange and re-sending.
|
||||
@@ -10188,25 +10300,71 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
|
||||
# Walk backwards to the last *real* user message. Timeline bookkeeping
|
||||
# rows (display_kind set) are role=user but are not user turns — match
|
||||
# CLI resume counting and list_recent_user_messages. Compaction
|
||||
# CLI resume counting and user_originated_turn_view. Compaction
|
||||
# handoffs are excluded too (durable role=user, sometimes without
|
||||
# display_kind on legacy sessions; #80622).
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
retryable_user_text,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from agent.memory_manager import sanitize_context
|
||||
from run_agent import _is_ephemeral_scaffolding
|
||||
|
||||
last_user_idx = None
|
||||
for i in range(len(self.conversation_history) - 1, -1, -1):
|
||||
msg = self.conversation_history[i]
|
||||
if is_user_originated_turn(msg):
|
||||
last_user_idx = i
|
||||
break
|
||||
warm_history = list(self.conversation_history)
|
||||
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(warm_history)
|
||||
if not _is_ephemeral_scaffolding(message)
|
||||
and user_originated_turn_view(message) is not None
|
||||
]
|
||||
|
||||
if last_user_idx is None:
|
||||
if not user_indices:
|
||||
print("(._.) No user message found to retry.")
|
||||
return None
|
||||
last_user_idx = user_indices[-1]
|
||||
|
||||
# Extract the message text and remove everything from that point forward
|
||||
last_message = self.conversation_history[last_user_idx].get("content", "")
|
||||
self.conversation_history = self.conversation_history[:last_user_idx]
|
||||
# Resolve a lossless live payload before touching either persistence or
|
||||
# memory. A force-user-leading compaction row is one physical carrier:
|
||||
# its historical handoff remains in the prefix while only the embedded
|
||||
# human ask is retried. Media cannot be replayed by /retry, so fail
|
||||
# closed before archiving anything.
|
||||
try:
|
||||
truncated, live_view = history_before_user_originated_turn(
|
||||
warm_history, last_user_idx
|
||||
)
|
||||
live_content = live_view.get("content")
|
||||
if isinstance(live_content, str):
|
||||
live_content = sanitize_context(live_content).strip()
|
||||
last_message = retryable_user_text(live_content)
|
||||
except ValueError as exc:
|
||||
print(f"(._.) Cannot retry that message safely: {exc}")
|
||||
return None
|
||||
|
||||
# Persist the rewind before publishing the shorter in-memory view.
|
||||
# The DB owns the physical carrier split so the archived original and
|
||||
# retained scaffold are committed atomically. A plain user row keeps
|
||||
# the legacy rewind shape (no replacement scaffold).
|
||||
if self._session_db is not None and self.session_id:
|
||||
try:
|
||||
truncated, _, _ = self._rewind_persisted_user_turn(
|
||||
warm_history=warm_history,
|
||||
user_ordinal=len(user_indices) - 1,
|
||||
warm_live_view=live_view,
|
||||
)
|
||||
except Exception as exc:
|
||||
print(f"(x_x) Retry rewind failed; history was not changed: {exc}")
|
||||
return None
|
||||
|
||||
self.conversation_history = truncated
|
||||
if self.agent is not None:
|
||||
if hasattr(self.agent, "_session_messages"):
|
||||
self.agent._session_messages = self.conversation_history
|
||||
if hasattr(self.agent, "_last_flushed_db_idx"):
|
||||
self.agent._last_flushed_db_idx = len(self.conversation_history)
|
||||
if hasattr(self.agent, "_db_flush_scan_prefix"):
|
||||
self.agent._db_flush_scan_prefix = self.conversation_history[:]
|
||||
|
||||
print(f"(^_^)b Retrying: \"{last_message[:60]}{'...' if len(last_message) > 60 else ''}\"")
|
||||
return last_message
|
||||
@@ -10243,61 +10401,64 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
|
||||
# Walk backwards collecting the indices of the last N *real* user
|
||||
# messages (exclude display_kind timeline rows and compaction
|
||||
# handoffs — same predicate as list_recent_user_messages, resume
|
||||
# handoffs — same predicate as user_originated_turn_view, resume
|
||||
# turn counting, and /retry; #80622).
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from run_agent import _is_ephemeral_scaffolding
|
||||
|
||||
user_indices = []
|
||||
for i in range(len(self.conversation_history) - 1, -1, -1):
|
||||
msg = self.conversation_history[i]
|
||||
if is_user_originated_turn(msg):
|
||||
user_indices.append(i)
|
||||
if len(user_indices) >= n:
|
||||
break
|
||||
warm_history = list(self.conversation_history)
|
||||
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(warm_history)
|
||||
if not _is_ephemeral_scaffolding(message)
|
||||
and user_originated_turn_view(message) is not None
|
||||
]
|
||||
|
||||
if not user_indices:
|
||||
print("(._.) No user message found to undo.")
|
||||
return
|
||||
|
||||
# The oldest of the collected user messages is our truncation point.
|
||||
cut_idx = user_indices[-1]
|
||||
turns_undone = len(user_indices)
|
||||
turns_undone = min(n, len(user_indices))
|
||||
target_ordinal = len(user_indices) - turns_undone
|
||||
cut_idx = user_indices[target_ordinal]
|
||||
|
||||
removed_count = len(self.conversation_history) - cut_idx
|
||||
removed_msg = self.conversation_history[cut_idx].get("content", "")
|
||||
removed_text = self._undo_content_to_text(removed_msg)
|
||||
|
||||
# Truncate the in-memory history to before that user message.
|
||||
self.conversation_history = self.conversation_history[:cut_idx]
|
||||
removed_count = len(warm_history) - cut_idx
|
||||
truncated, live_view = history_before_user_originated_turn(
|
||||
warm_history, cut_idx
|
||||
)
|
||||
removed_text = self._undo_content_to_text(live_view.get("content"))
|
||||
|
||||
# Soft-delete the truncated rows on disk so re-prompts and search
|
||||
# see the clean transcript while the rows survive for audit.
|
||||
rewound_rows = 0
|
||||
if self._session_db is not None and self.session_id:
|
||||
try:
|
||||
recents = self._session_db.list_recent_user_messages(
|
||||
self.session_id, limit=max(turns_undone, 10)
|
||||
truncated, durable_live_view, result = (
|
||||
self._rewind_persisted_user_turn(
|
||||
warm_history=warm_history,
|
||||
user_ordinal=target_ordinal,
|
||||
warm_live_view=live_view,
|
||||
)
|
||||
)
|
||||
if recents:
|
||||
target_idx = min(turns_undone - 1, len(recents) - 1)
|
||||
target_id = recents[target_idx]["id"]
|
||||
result = self._session_db.rewind_to_message(
|
||||
self.session_id, target_id
|
||||
)
|
||||
rewound_rows = result.get("rewound_count", 0)
|
||||
# Prefer the DB's decoded target text for the prefill —
|
||||
# it's the canonical persisted copy.
|
||||
db_text = self._undo_content_to_text(
|
||||
(result.get("target_message") or {}).get("content")
|
||||
)
|
||||
if db_text:
|
||||
removed_text = db_text
|
||||
except ValueError as e:
|
||||
# Non-user target / cross-session — keep the in-memory undo
|
||||
# but skip the soft-delete; surface a debug-level note.
|
||||
logger.debug("undo: soft-delete skipped: %s", e)
|
||||
# Canonicalize the editable prefill before mutation. The raw
|
||||
# physical carrier contains the reference summary wrapper.
|
||||
durable_text = self._undo_content_to_text(
|
||||
durable_live_view.get("content")
|
||||
)
|
||||
if durable_text:
|
||||
removed_text = durable_text
|
||||
rewound_rows = result.get("rewound_count", 0)
|
||||
except Exception as e:
|
||||
logger.debug("undo: soft-delete failed: %s", e)
|
||||
logger.debug("undo: durable rewind failed: %s", e)
|
||||
print(f"(x_x) Undo failed; history was not changed: {e}")
|
||||
return
|
||||
|
||||
# Publish only after the durable rewind succeeds (or no store exists).
|
||||
self.conversation_history = truncated
|
||||
|
||||
# Agent surgery: invalidate the system-prompt cache and reset the
|
||||
# flush index so the next turn re-flushes from the truncated head.
|
||||
@@ -10312,6 +10473,10 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
|
||||
self.agent._last_flushed_db_idx = len(self.conversation_history)
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(self.agent, "_session_messages"):
|
||||
self.agent._session_messages = self.conversation_history
|
||||
if hasattr(self.agent, "_db_flush_scan_prefix"):
|
||||
self.agent._db_flush_scan_prefix = self.conversation_history[:]
|
||||
# Notify memory providers — same hook /branch fires, with the
|
||||
# rewound flag so per-turn document caches invalidate (#6672, #21910).
|
||||
try:
|
||||
|
||||
+111
-50
@@ -3650,17 +3650,21 @@ class SessionStore:
|
||||
entry = self._entries.get(session_key)
|
||||
return getattr(entry, "session_id", None) if entry else None
|
||||
|
||||
def append_to_transcript(self, session_id: str, message: Dict[str, Any], skip_db: bool = False) -> None:
|
||||
"""Serialize transcript draining across queue migration boundaries."""
|
||||
if not self._db or skip_db:
|
||||
return
|
||||
def _get_transcript_drain_lock(self):
|
||||
"""Return the lock that serializes pending-queue drain boundaries."""
|
||||
drain_lock = getattr(self, "_transcript_drain_lock", None)
|
||||
if drain_lock is None:
|
||||
# Compatibility for old in-memory/test instances created via
|
||||
# object.__new__ before this field existed.
|
||||
drain_lock = threading.RLock()
|
||||
self._transcript_drain_lock = drain_lock
|
||||
with drain_lock:
|
||||
return drain_lock
|
||||
|
||||
def append_to_transcript(self, session_id: str, message: Dict[str, Any], skip_db: bool = False) -> None:
|
||||
"""Serialize transcript draining across queue migration boundaries."""
|
||||
if not self._db or skip_db:
|
||||
return
|
||||
with self._get_transcript_drain_lock():
|
||||
reroutes = getattr(self, "_transcript_reroutes", None)
|
||||
if reroutes is None:
|
||||
reroutes = {}
|
||||
@@ -3999,6 +4003,7 @@ class SessionStore:
|
||||
session_id: str,
|
||||
messages: List[Dict[str, Any]],
|
||||
active_only: bool = False,
|
||||
reject_active_turn_lease: bool = False,
|
||||
) -> bool:
|
||||
"""Replace the entire transcript for a session with new messages.
|
||||
|
||||
@@ -4019,16 +4024,26 @@ class SessionStore:
|
||||
change on top of a failed write — e.g. /compress repointing the live
|
||||
session onto a fresh session_id — must check it so they can surface an
|
||||
error instead of silently dropping the conversation.
|
||||
|
||||
``reject_active_turn_lease`` is for user-initiated rewrites that do not
|
||||
own the cross-process turn lease. It leaves internal rewrite policy
|
||||
unchanged for existing callers unless they opt in explicitly.
|
||||
"""
|
||||
if not self._db:
|
||||
return True
|
||||
self._clear_dirty_transcript(session_id)
|
||||
try:
|
||||
self._db.replace_messages(session_id, messages, active_only=active_only)
|
||||
with self._get_transcript_drain_lock():
|
||||
try:
|
||||
self._db.replace_messages(
|
||||
session_id,
|
||||
messages,
|
||||
active_only=active_only,
|
||||
reject_active_turn_lease=reject_active_turn_lease,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug("Failed to rewrite transcript in DB: %s", e)
|
||||
return False
|
||||
self._clear_dirty_transcript(session_id)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.debug("Failed to rewrite transcript in DB: %s", e)
|
||||
return False
|
||||
|
||||
def load_transcript(self, session_id: str) -> List[Dict[str, Any]]:
|
||||
"""Load all messages from a session's transcript.
|
||||
@@ -4080,7 +4095,13 @@ class SessionStore:
|
||||
)
|
||||
return []
|
||||
|
||||
def rewind_session(self, session_id: str, n: int = 1) -> Optional[Dict[str, Any]]:
|
||||
def rewind_session(
|
||||
self,
|
||||
session_id: str,
|
||||
n: int = 1,
|
||||
*,
|
||||
require_retryable_composite: bool = False,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Back up ``n`` user turns via soft-delete, keeping rows for audit.
|
||||
|
||||
Unlike :meth:`rewrite_transcript` (a hard replace used by /retry),
|
||||
@@ -4091,47 +4112,87 @@ class SessionStore:
|
||||
Returns a dict ``{"rewound_count", "turns_undone", "target_text"}`` on
|
||||
success, or ``None`` if there's no DB or no user message to back up to.
|
||||
``n`` clamps to the oldest user turn when it exceeds the turn count.
|
||||
``require_retryable_composite`` is the gateway ``/retry`` guard: the
|
||||
selected current turn must still be a composite carrier, and its live
|
||||
payload must be losslessly replayable as text before anything changes.
|
||||
"""
|
||||
if not self._db:
|
||||
return None
|
||||
self._clear_dirty_transcript(session_id)
|
||||
if n < 1:
|
||||
n = 1
|
||||
try:
|
||||
recents = self._db.list_recent_user_messages(session_id, limit=max(n, 10))
|
||||
except Exception as e:
|
||||
logger.debug("rewind_session: failed to list user messages: %s", e)
|
||||
return None
|
||||
if not recents:
|
||||
return None
|
||||
target_idx = min(n - 1, len(recents) - 1)
|
||||
target_id = recents[target_idx]["id"]
|
||||
try:
|
||||
result = self._db.rewind_to_message(session_id, target_id)
|
||||
except ValueError as e:
|
||||
logger.debug("rewind_session: %s", e)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug("rewind_session: rewind_to_message failed: %s", e)
|
||||
return None
|
||||
target_msg = result.get("target_message") or {}
|
||||
content = target_msg.get("content") or ""
|
||||
if isinstance(content, list):
|
||||
parts = [
|
||||
p.get("text", "")
|
||||
for p in content
|
||||
if isinstance(p, dict) and p.get("type") == "text"
|
||||
]
|
||||
target_text = "\n".join(t for t in parts if t)
|
||||
elif isinstance(content, str):
|
||||
target_text = content
|
||||
else:
|
||||
target_text = ""
|
||||
return {
|
||||
"rewound_count": result.get("rewound_count", 0),
|
||||
"turns_undone": target_idx + 1,
|
||||
"target_text": target_text,
|
||||
}
|
||||
with self._get_transcript_drain_lock():
|
||||
if n < 1:
|
||||
n = 1
|
||||
from agent.context_compressor import (
|
||||
retryable_user_text,
|
||||
split_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
|
||||
try:
|
||||
expected_active_ids = self._db.get_active_message_ids(session_id)
|
||||
durable = self._db.get_messages_as_conversation(
|
||||
session_id,
|
||||
include_row_ids=True,
|
||||
)
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(durable)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if not user_indices:
|
||||
return None
|
||||
turns_undone = min(n, len(user_indices))
|
||||
target = durable[user_indices[-turns_undone]]
|
||||
target_id = target.get("_row_id")
|
||||
if not isinstance(target_id, int):
|
||||
return None
|
||||
handoff, target_view = split_user_originated_turn(target)
|
||||
if target_view is None:
|
||||
return None
|
||||
if require_retryable_composite and handoff is None:
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug("rewind_session: failed to resolve canonical target: %s", e)
|
||||
return None
|
||||
if require_retryable_composite:
|
||||
# Keep replay-policy failures distinct from persistence errors
|
||||
# so /retry can explain why the selected carrier is unsafe.
|
||||
target_text = retryable_user_text(target_view.get("content"))
|
||||
try:
|
||||
result = self._db.rewind_to_message(
|
||||
session_id,
|
||||
target_id,
|
||||
preserve_compaction_handoff=handoff is not None,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content=target_view.get("content"),
|
||||
)
|
||||
except ValueError as e:
|
||||
logger.debug("rewind_session: %s", e)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug("rewind_session: rewind_to_message failed: %s", e)
|
||||
return None
|
||||
self._clear_dirty_transcript(session_id)
|
||||
# ``target_view`` is the canonical live projection of the physical DB
|
||||
# row. For a composite carrier, the raw target contains the historical
|
||||
# summary wrapper and must never be echoed back as the editable prompt.
|
||||
if not require_retryable_composite:
|
||||
content = target_view.get("content") or ""
|
||||
if isinstance(content, list):
|
||||
parts = [
|
||||
p.get("text", "")
|
||||
for p in content
|
||||
if isinstance(p, dict) and p.get("type") == "text"
|
||||
]
|
||||
target_text = "\n".join(t for t in parts if t)
|
||||
elif isinstance(content, str):
|
||||
target_text = content
|
||||
else:
|
||||
target_text = ""
|
||||
return {
|
||||
"rewound_count": result.get("rewound_count", 0),
|
||||
"turns_undone": turns_undone,
|
||||
"target_text": target_text,
|
||||
}
|
||||
|
||||
|
||||
def build_session_context(
|
||||
|
||||
+56
-21
@@ -2644,34 +2644,69 @@ class GatewaySlashCommandsMixin:
|
||||
# auto_continue / hidden); clients never count them as user turns.
|
||||
# Without this filter /retry rewrote the transcript around a marker
|
||||
# and re-sent opaque bookkeeping text (same class as the TUI ordinal).
|
||||
last_user_msg = None
|
||||
last_user_idx = None
|
||||
# is_user_originated_turn: excludes display_kind bookkeeping AND
|
||||
# compaction handoffs (durable role=user, sometimes without
|
||||
# display_kind on legacy sessions; #80622) — /retry must never
|
||||
# re-send a reference-only summary as if the user asked it.
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
# The canonical projection excludes bookkeeping and pure handoffs while
|
||||
# still recognizing a real ask embedded in a compaction carrier.
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
retryable_user_text,
|
||||
split_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
|
||||
for i in range(len(history) - 1, -1, -1):
|
||||
msg = history[i]
|
||||
if is_user_originated_turn(msg):
|
||||
last_user_msg = msg.get("content", "")
|
||||
if user_originated_turn_view(msg) is not None:
|
||||
last_user_idx = i
|
||||
break
|
||||
|
||||
if not last_user_msg:
|
||||
|
||||
if last_user_idx is None:
|
||||
return t("gateway.retry.no_previous")
|
||||
|
||||
# Truncate history to before the last user message and persist only the
|
||||
# live view. After in-place compaction the pre-compaction transcript
|
||||
# lives on as active=0/compacted=1 rows under this same session id, and
|
||||
# a bare rewrite (active_only=False) would DELETE them (same class as
|
||||
# #61145). /retry never intends to purge archived history, so avoid a
|
||||
# separate existence probe: it could fail open or race with the write.
|
||||
truncated = history[:last_user_idx]
|
||||
await self.async_session_store.rewrite_transcript(
|
||||
session_entry.session_id, truncated, active_only=True
|
||||
)
|
||||
|
||||
# Resolve the live text and the scaffold-preserving prefix before any
|
||||
# transcript write. Messaging retries cannot reconstruct attachments;
|
||||
# reject media/unknown content without truncating the session.
|
||||
try:
|
||||
truncated, live_view = history_before_user_originated_turn(
|
||||
history, last_user_idx
|
||||
)
|
||||
last_user_msg = retryable_user_text(live_view.get("content"))
|
||||
handoff, _ = split_user_originated_turn(history[last_user_idx])
|
||||
except ValueError as exc:
|
||||
return f"Cannot retry that message safely: {exc}"
|
||||
|
||||
if handoff is not None:
|
||||
# A composite carrier is one physical row containing both the
|
||||
# retained summary and the live ask. Let the carrier-aware rewind
|
||||
# archive that row/tail and insert its pure scaffold atomically.
|
||||
# Plain turns keep the existing rewrite path below; #84078 owns
|
||||
# its separate archive_dropped/prefix-CAS semantics.
|
||||
try:
|
||||
rewind_result = await self.async_session_store.rewind_session(
|
||||
session_entry.session_id,
|
||||
1,
|
||||
require_retryable_composite=True,
|
||||
)
|
||||
except ValueError as exc:
|
||||
return f"Cannot retry that message safely: {exc}"
|
||||
if rewind_result is None:
|
||||
return "Retry failed; transcript was not changed."
|
||||
# The store reselects and validates the latest carrier on the same
|
||||
# snapshot used by the atomic rewind. A concurrent newer turn can
|
||||
# therefore never be removed while this handler resends stale text.
|
||||
last_user_msg = rewind_result["target_text"]
|
||||
else:
|
||||
# After in-place compaction the pre-compaction transcript lives on
|
||||
# as active=0/compacted=1 rows under this session id. active_only
|
||||
# preserves that archive; a separate existence probe could fail
|
||||
# open or race with the write.
|
||||
if not await self.async_session_store.rewrite_transcript(
|
||||
session_entry.session_id,
|
||||
truncated,
|
||||
active_only=True,
|
||||
reject_active_turn_lease=True,
|
||||
):
|
||||
return "Retry failed; transcript was not changed."
|
||||
# Reset stored token count — transcript was truncated
|
||||
session_entry.last_prompt_tokens = 0
|
||||
|
||||
|
||||
@@ -642,14 +642,34 @@ async def get_session_messages(
|
||||
if result is None:
|
||||
raise HTTPException(status_code=404, detail="Session not found")
|
||||
sid, _limit, messages = result
|
||||
from agent.compaction_display import project_compaction_message_for_display
|
||||
from agent.context_compressor import is_compaction_summary_message
|
||||
|
||||
projected_messages = []
|
||||
for message in messages:
|
||||
if not is_compaction_summary_message(message):
|
||||
projected_messages.append(message)
|
||||
continue
|
||||
display_view = project_compaction_message_for_display(message)
|
||||
projected = message.copy()
|
||||
if display_view is None:
|
||||
if not projected.get("display_kind"):
|
||||
projected["display_kind"] = "hidden"
|
||||
else:
|
||||
# Keep the physical content for inspection/export compatibility;
|
||||
# Desktop consumes this display-only projection. A legacy hidden
|
||||
# wrapper must not hide a successfully recovered live ask.
|
||||
projected["display_content"] = display_view.get("content")
|
||||
projected.pop("display_kind", None)
|
||||
projected_messages.append(projected)
|
||||
return {
|
||||
"session_id": sid,
|
||||
"messages": messages,
|
||||
"messages": projected_messages,
|
||||
"pagination": {
|
||||
"limit": _limit,
|
||||
"offset": offset,
|
||||
"order": order or ("latest" if limit is None else "oldest"),
|
||||
"returned": len(messages),
|
||||
"returned": len(projected_messages),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
+404
-78
@@ -10346,13 +10346,16 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
compression_lock_holder: Optional[str],
|
||||
turn_lease_holder: Optional[str] = None,
|
||||
turn_lease_ttl_seconds: float = 300.0,
|
||||
reject_active_turn_lease: bool = False,
|
||||
reject_active_compression_lock: bool = False,
|
||||
) -> None:
|
||||
"""Transcript-append admission checks, run INSIDE the write txn.
|
||||
"""Transcript-write admission checks, run INSIDE the write txn.
|
||||
|
||||
Shared by :meth:`append_message` and :meth:`append_messages_batch` so
|
||||
the two writers can never diverge on these correctness invariants
|
||||
(this guard has already needed targeted fixes — see the #74478
|
||||
patience note below).
|
||||
(this guard has already needed targeted fixes — see the #74478 patience
|
||||
note below). User-initiated transcript mutations may opt in to rejecting
|
||||
an active unowned turn lease in that same transaction.
|
||||
"""
|
||||
# NOTE (#75316 redesign): appends do NOT check compression_locks.
|
||||
# The lock's job is to stop two COMPRESSIONS colliding, not to fence
|
||||
@@ -10364,33 +10367,78 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
# symptom family — turns dying as session_persistence_failed while a
|
||||
# slow provider summary held the lease (#74568, #77386), including
|
||||
# stale locks from dead PIDs blocking writes for the full TTL.
|
||||
if turn_lease_holder:
|
||||
# Destructive user mutations are different: a compressor that already
|
||||
# captured its watermark can otherwise publish the pre-rewind snapshot
|
||||
# after the mutation and resurrect the removed turn. Keep that narrow
|
||||
# fence opt-in so ordinary appends retain the watermark behavior.
|
||||
if reject_active_compression_lock:
|
||||
active_lock = conn.execute(
|
||||
"SELECT holder, expires_at FROM compression_locks "
|
||||
"WHERE session_id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if active_lock is not None:
|
||||
current_holder = active_lock["holder"]
|
||||
if (
|
||||
float(active_lock["expires_at"]) <= time.time()
|
||||
or _compression_lock_holder_process_is_dead(current_holder)
|
||||
):
|
||||
conn.execute(
|
||||
"DELETE FROM compression_locks "
|
||||
"WHERE session_id = ? AND holder = ?",
|
||||
(session_id, current_holder),
|
||||
)
|
||||
elif current_holder != compression_lock_holder:
|
||||
raise SessionCompressionInProgressError(
|
||||
f"Session {session_id!r} is being compressed by another writer"
|
||||
)
|
||||
if turn_lease_holder or reject_active_turn_lease:
|
||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||
lease = conn.execute(
|
||||
"SELECT holder, expires_at FROM session_turn_leases "
|
||||
"WHERE conversation_id = ?",
|
||||
(conversation_id,),
|
||||
).fetchone()
|
||||
if lease is None or lease["holder"] != turn_lease_holder:
|
||||
raise SessionTurnLeaseLostError(
|
||||
f"Session turn lease lost; refusing transcript write "
|
||||
f"for {session_id!r}"
|
||||
)
|
||||
now = time.time()
|
||||
if float(lease["expires_at"]) <= now:
|
||||
# Expiry makes the row reclaimable; it does not prove that a
|
||||
# takeover occurred. BEGIN IMMEDIATE serializes this renewal
|
||||
# with acquisition, so a still-matching owner can recover from
|
||||
# a starved refresher without weakening the foreign-holder fence.
|
||||
conn.execute(
|
||||
"UPDATE session_turn_leases SET expires_at = ? "
|
||||
"WHERE conversation_id = ? AND holder = ?",
|
||||
(
|
||||
now + max(0.1, float(turn_lease_ttl_seconds)),
|
||||
conversation_id,
|
||||
turn_lease_holder,
|
||||
),
|
||||
)
|
||||
if turn_lease_holder:
|
||||
if lease is None or lease["holder"] != turn_lease_holder:
|
||||
raise SessionTurnLeaseLostError(
|
||||
f"Session turn lease lost; refusing transcript write "
|
||||
f"for {session_id!r}"
|
||||
)
|
||||
if float(lease["expires_at"]) <= now:
|
||||
# Expiry makes the row reclaimable; it does not prove that a
|
||||
# takeover occurred. BEGIN IMMEDIATE serializes this renewal
|
||||
# with acquisition, so a still-matching owner can recover from
|
||||
# a starved refresher without weakening the foreign-holder fence.
|
||||
conn.execute(
|
||||
"UPDATE session_turn_leases SET expires_at = ? "
|
||||
"WHERE conversation_id = ? AND holder = ?",
|
||||
(
|
||||
now + max(0.1, float(turn_lease_ttl_seconds)),
|
||||
conversation_id,
|
||||
turn_lease_holder,
|
||||
),
|
||||
)
|
||||
elif lease is not None:
|
||||
current_holder = lease["holder"]
|
||||
if (
|
||||
float(lease["expires_at"]) <= now
|
||||
or _compression_lock_holder_process_is_dead(current_holder)
|
||||
):
|
||||
# Match acquisition semantics: an expired or provably dead
|
||||
# owner is reclaimable. Deleting it inside this BEGIN IMMEDIATE
|
||||
# transaction also fences a stale late flush after the mutation.
|
||||
conn.execute(
|
||||
"DELETE FROM session_turn_leases "
|
||||
"WHERE conversation_id = ? AND holder = ?",
|
||||
(conversation_id, current_holder),
|
||||
)
|
||||
else:
|
||||
raise SessionTurnLeaseLostError(
|
||||
f"Session has an active turn lease; refusing transcript "
|
||||
f"mutation for {session_id!r}"
|
||||
)
|
||||
session = conn.execute(
|
||||
"SELECT ended_at, end_reason FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
@@ -11013,6 +11061,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
messages: List[Dict[str, Any]],
|
||||
active_only: bool = False,
|
||||
archive_dropped: bool = False,
|
||||
reject_active_turn_lease: bool = False,
|
||||
) -> None:
|
||||
"""Atomically replace the stored messages for a session.
|
||||
|
||||
@@ -11047,21 +11096,35 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
fresh active rows exactly as in the destructive path, so the live view
|
||||
is identical either way; only the durability of the dropped turns
|
||||
differs.
|
||||
|
||||
Pass ``reject_active_turn_lease=True`` for user-initiated rewrites that
|
||||
do not already own the cross-process turn lease. The lease check and
|
||||
transcript mutation then share one write transaction, so a second
|
||||
process cannot archive or replace a turn that is still being produced.
|
||||
"""
|
||||
|
||||
active_clause = " AND active = 1" if active_only else ""
|
||||
|
||||
def _do(conn):
|
||||
session = conn.execute(
|
||||
"SELECT ended_at, end_reason FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
session is not None
|
||||
and session["ended_at"] is not None
|
||||
and session["end_reason"] == "compression"
|
||||
):
|
||||
raise CompressionSessionClosedError(session_id)
|
||||
if reject_active_turn_lease:
|
||||
self._check_transcript_write_guards(
|
||||
conn,
|
||||
session_id,
|
||||
None,
|
||||
reject_active_turn_lease=True,
|
||||
reject_active_compression_lock=True,
|
||||
)
|
||||
else:
|
||||
session = conn.execute(
|
||||
"SELECT ended_at, end_reason FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
session is not None
|
||||
and session["ended_at"] is not None
|
||||
and session["end_reason"] == "compression"
|
||||
):
|
||||
raise CompressionSessionClosedError(session_id)
|
||||
if archive_dropped:
|
||||
# Content-preserving UPDATE: the rows keep their FTS entries
|
||||
# (the messages_fts triggers fire on INSERT / DELETE / UPDATE
|
||||
@@ -11399,13 +11462,30 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
all_rows = cursor.fetchall()
|
||||
seen: dict = {}
|
||||
for row in all_rows:
|
||||
dedupe_content = row["content"]
|
||||
if row["role"] == "user":
|
||||
from agent.context_compressor import split_user_originated_turn
|
||||
|
||||
candidate = {
|
||||
"role": "user",
|
||||
"content": self._decode_content(row["content"]),
|
||||
"display_kind": row["display_kind"],
|
||||
"display_metadata": self._decode_display_metadata(
|
||||
row["display_metadata"]
|
||||
),
|
||||
}
|
||||
handoff, live_view = split_user_originated_turn(candidate)
|
||||
if handoff is not None and live_view is not None:
|
||||
dedupe_content = self._encode_content(
|
||||
live_view.get("content")
|
||||
)
|
||||
# Tool fields participate in the dedupe key: compaction copies
|
||||
# them verbatim, so identical tool messages across generations
|
||||
# still collapse, while distinct tool calls that happen to
|
||||
# share role/content/timestamp are never merged.
|
||||
key = (
|
||||
row["role"],
|
||||
row["content"],
|
||||
dedupe_content,
|
||||
row["timestamp"],
|
||||
row["tool_call_id"],
|
||||
row["tool_calls"],
|
||||
@@ -11730,6 +11810,11 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
desired session set / active state by the caller.
|
||||
"""
|
||||
messages = []
|
||||
# Watermark rotation column-clones concurrent tail rows into the child
|
||||
# after the new summary, so the copies need not be adjacent. Index the
|
||||
# exact durable clone identity while decoding instead of rescanning the
|
||||
# whole accumulated lineage for every user row.
|
||||
exact_user_clones: Dict[Tuple[Any, str], Dict[str, Any]] = {}
|
||||
for row in rows:
|
||||
content = self._decode_content(row["content"])
|
||||
if row["role"] in {"user", "assistant"} and isinstance(content, str):
|
||||
@@ -11820,9 +11905,47 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.warning("Failed to deserialize codex_message_items, falling back to None")
|
||||
msg["codex_message_items"] = None
|
||||
if include_ancestors and self._is_duplicate_replayed_user_message(messages, msg):
|
||||
continue
|
||||
if include_ancestors:
|
||||
canonical_content, _is_composite = (
|
||||
self._canonical_replayed_user_content(msg)
|
||||
)
|
||||
exact_clone_key = self._exact_replayed_user_clone_key(
|
||||
msg.get("timestamp"), canonical_content
|
||||
)
|
||||
previous_exact = (
|
||||
exact_user_clones.get(exact_clone_key)
|
||||
if exact_clone_key is not None
|
||||
else None
|
||||
)
|
||||
duplicate = None
|
||||
if previous_exact is not None:
|
||||
previous_index = next(
|
||||
(
|
||||
index
|
||||
for index, candidate in enumerate(messages)
|
||||
if candidate is previous_exact
|
||||
),
|
||||
None,
|
||||
)
|
||||
if previous_index is not None:
|
||||
duplicate = (previous_index, True)
|
||||
if duplicate is None:
|
||||
duplicate = self._find_duplicate_replayed_user_message(
|
||||
messages, msg
|
||||
)
|
||||
if duplicate is not None:
|
||||
duplicate_index, prefer_current = duplicate
|
||||
if prefer_current:
|
||||
# A rotated compression child can carry the same live
|
||||
# ask as the parent row plus the only surviving summary
|
||||
# scaffold. Keep the child carrier (and its durable row
|
||||
# id), not the simpler ancestor copy.
|
||||
messages.pop(duplicate_index)
|
||||
else:
|
||||
continue
|
||||
messages.append(msg)
|
||||
if include_ancestors and exact_clone_key is not None:
|
||||
exact_user_clones[exact_clone_key] = msg
|
||||
# DEFENSE-IN-DEPTH against background-review session pollution: a forked
|
||||
# skill/memory review that (in older builds, before the _persist_disabled
|
||||
# fix) shared the parent's session_id wrote its harness turn into this
|
||||
@@ -12033,15 +12156,28 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
"ORDER BY id",
|
||||
tuple(session_ids),
|
||||
).fetchall()
|
||||
ancestor_rows = [r for r in rows if r["session_id"] != session_id]
|
||||
if not ancestor_rows:
|
||||
ancestor_ids = {
|
||||
int(row["id"])
|
||||
for row in rows
|
||||
if row["session_id"] != session_id and row["id"] is not None
|
||||
}
|
||||
if not ancestor_ids:
|
||||
return []
|
||||
return self._rows_to_conversation(
|
||||
ancestor_rows,
|
||||
lineage = self._rows_to_conversation(
|
||||
rows,
|
||||
session_id=session_id,
|
||||
include_ancestors=True,
|
||||
repair_alternation=False,
|
||||
include_row_ids=True,
|
||||
)
|
||||
prefix: List[Dict[str, Any]] = []
|
||||
for message in lineage:
|
||||
if message.get("_row_id") not in ancestor_ids:
|
||||
continue
|
||||
projected = message.copy()
|
||||
projected.pop("_row_id", None)
|
||||
prefix.append(projected)
|
||||
return prefix
|
||||
|
||||
def _is_explicit_branch_session(self, session_id: str) -> bool:
|
||||
"""Return whether *session_id* is a copied user-facing branch.
|
||||
@@ -12107,25 +12243,134 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
return list(reversed(chain)) or [session_id]
|
||||
|
||||
@staticmethod
|
||||
def _is_duplicate_replayed_user_message(messages: List[Dict[str, Any]], msg: Dict[str, Any]) -> bool:
|
||||
def _canonical_replayed_user_content(
|
||||
msg: Dict[str, Any],
|
||||
) -> Tuple[Any, bool]:
|
||||
"""Return canonical live content and whether *msg* is composite."""
|
||||
if msg.get("role") != "user":
|
||||
return False
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, str) or not content:
|
||||
return False
|
||||
for prev in reversed(messages):
|
||||
if prev.get("role") == "user" and prev.get("content") == content:
|
||||
return True
|
||||
return None, False
|
||||
|
||||
from agent.context_compressor import split_user_originated_turn
|
||||
|
||||
handoff, live_view = split_user_originated_turn(msg)
|
||||
is_composite = handoff is not None and live_view is not None
|
||||
return (
|
||||
live_view.get("content")
|
||||
if is_composite and live_view is not None
|
||||
else msg.get("content"),
|
||||
is_composite,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _exact_replayed_user_clone_key(
|
||||
timestamp: Any, content: Any
|
||||
) -> Optional[Tuple[Any, str]]:
|
||||
"""Return a hashable key for a column-exact rotation clone."""
|
||||
if timestamp is None or content in (None, "", []):
|
||||
return None
|
||||
try:
|
||||
encoded = json.dumps(
|
||||
content,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return timestamp, encoded
|
||||
|
||||
@staticmethod
|
||||
def _find_duplicate_replayed_user_message(
|
||||
messages: List[Dict[str, Any]], msg: Dict[str, Any]
|
||||
) -> Optional[Tuple[int, bool]]:
|
||||
"""Return an adjacent replay duplicate and whether *msg* must win.
|
||||
|
||||
Compression rotation may persist the current ask once in the parent
|
||||
and again inside a composite child carrier. Compare the canonical live
|
||||
payload for that carrier, while retaining the historical exact-string
|
||||
dedupe for ordinary replayed users. The child carrier wins because it
|
||||
owns both the current durable row identity and the retained scaffold.
|
||||
"""
|
||||
if msg.get("role") != "user":
|
||||
return None
|
||||
|
||||
content, prefer_current = SessionDB._canonical_replayed_user_content(msg)
|
||||
if content in (None, "", []):
|
||||
return None
|
||||
|
||||
for index in range(len(messages) - 1, -1, -1):
|
||||
prev = messages[index]
|
||||
if prev.get("role") == "user":
|
||||
prev_content, prev_is_composite = (
|
||||
SessionDB._canonical_replayed_user_content(prev)
|
||||
)
|
||||
if prev_content == content and (
|
||||
prefer_current
|
||||
or prev_is_composite
|
||||
or isinstance(content, str)
|
||||
):
|
||||
return index, prefer_current
|
||||
if prev.get("role") == "assistant" and (prev.get("content") or prev.get("tool_calls")):
|
||||
return False
|
||||
return False
|
||||
return None
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _is_duplicate_replayed_user_message(
|
||||
messages: List[Dict[str, Any]], msg: Dict[str, Any]
|
||||
) -> bool:
|
||||
return SessionDB._find_duplicate_replayed_user_message(messages, msg) is not None
|
||||
|
||||
# =========================================================================
|
||||
# Rewind (soft-delete) — see /rewind slash command + issue #21910
|
||||
# =========================================================================
|
||||
|
||||
def get_active_message_ids(self, session_id: str) -> List[int]:
|
||||
"""Return the ordered physical ids pinned by rewind CAS checks.
|
||||
|
||||
Conversation projections intentionally omit legacy background-review
|
||||
harness rows. Destructive rewinds must nevertheless pin every active
|
||||
physical row so the caller snapshot matches the transaction-local
|
||||
comparison in :meth:`rewind_to_message`.
|
||||
"""
|
||||
with self._read_ctx() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT id FROM messages "
|
||||
"WHERE session_id = ? AND active = 1 ORDER BY id",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [int(row[0]) for row in rows]
|
||||
|
||||
@staticmethod
|
||||
def _active_transcript_counts(conn, session_id: str) -> tuple[int, int]:
|
||||
"""Return active message/tool-call counts inside the caller's txn."""
|
||||
rows = conn.execute(
|
||||
"SELECT tool_calls FROM messages "
|
||||
"WHERE session_id = ? AND active = 1",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
tool_call_count = 0
|
||||
for row in rows:
|
||||
raw = row[0]
|
||||
if not raw:
|
||||
continue
|
||||
try:
|
||||
decoded = json.loads(raw) if isinstance(raw, str) else raw
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
continue
|
||||
if isinstance(decoded, list):
|
||||
tool_call_count += len(decoded)
|
||||
elif decoded:
|
||||
tool_call_count += 1
|
||||
return len(rows), tool_call_count
|
||||
|
||||
def rewind_to_message(
|
||||
self, session_id: str, target_message_id: int
|
||||
self,
|
||||
session_id: str,
|
||||
target_message_id: int,
|
||||
*,
|
||||
preserve_compaction_handoff: bool = False,
|
||||
expected_active_ids: Optional[List[int]] = None,
|
||||
expected_target_content: Any = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Soft-delete all messages with id >= ``target_message_id`` in *session_id*.
|
||||
|
||||
@@ -12144,7 +12389,19 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
}
|
||||
|
||||
Raises ``ValueError`` if the target message does not exist in
|
||||
*session_id* or if its role is not ``"user"``.
|
||||
*session_id* or if its role is not ``"user"``. With
|
||||
``preserve_compaction_handoff=True``, a composite summary carrier is
|
||||
split inside the same write transaction: its original row is archived
|
||||
and its canonical hidden handoff scaffold is inserted as the new head.
|
||||
That opt-in result also contains ``replacement_message_id``.
|
||||
|
||||
``expected_active_ids`` optionally pins the ordered active row set.
|
||||
``expected_target_content`` additionally pins the selected canonical
|
||||
live-user payload. Both checks run inside the write transaction before
|
||||
any row or counter mutation. Presentation-only metadata changes (for
|
||||
example Desktop reactions) deliberately do not invalidate a rewind.
|
||||
A live cross-process turn lease always refuses the rewind; expired or
|
||||
provably dead holders are reclaimed inside the mutation transaction.
|
||||
|
||||
Always increments ``sessions.rewind_count`` — even when the
|
||||
target is already inactive — so the counter accurately reflects
|
||||
@@ -12153,29 +12410,78 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
target is a no-op on row state but still bumps the counter.
|
||||
"""
|
||||
|
||||
# 1) Validate target up-front (read-only, outside the write txn).
|
||||
with self._lock:
|
||||
row = self._conn.execute(
|
||||
def _do(conn):
|
||||
# Rewind changes the active transcript and must honor the same
|
||||
# compression/closed-parent and cross-process turn guards as
|
||||
# append writers.
|
||||
self._check_transcript_write_guards(
|
||||
conn,
|
||||
session_id,
|
||||
None,
|
||||
reject_active_turn_lease=True,
|
||||
reject_active_compression_lock=True,
|
||||
)
|
||||
|
||||
if expected_active_ids is not None:
|
||||
active_rows = conn.execute(
|
||||
"SELECT id FROM messages "
|
||||
"WHERE session_id = ? AND active = 1 ORDER BY id",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
active_ids = [int(active_row[0]) for active_row in active_rows]
|
||||
if active_ids != expected_active_ids:
|
||||
raise RuntimeError(
|
||||
"active transcript changed before the rewind could be persisted"
|
||||
)
|
||||
|
||||
row = conn.execute(
|
||||
"SELECT * FROM messages WHERE id = ? AND session_id = ?",
|
||||
(target_message_id, session_id),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
raise ValueError(
|
||||
f"message {target_message_id} not found in session {session_id}"
|
||||
)
|
||||
target_row = dict(row)
|
||||
if target_row.get("role") != "user":
|
||||
raise ValueError(
|
||||
f"rewind target must be a 'user' message (got role="
|
||||
f"{target_row.get('role')!r}, id={target_message_id})"
|
||||
)
|
||||
if row is None:
|
||||
raise ValueError(
|
||||
f"message {target_message_id} not found in session {session_id}"
|
||||
)
|
||||
target_row = dict(row)
|
||||
if target_row.get("role") != "user":
|
||||
raise ValueError(
|
||||
f"rewind target must be a 'user' message (got role="
|
||||
f"{target_row.get('role')!r}, id={target_message_id})"
|
||||
)
|
||||
|
||||
# Decode content for callers (prefill the prompt buffer).
|
||||
target_row["content"] = self._decode_content(target_row.get("content"))
|
||||
replacement_message_id: Optional[int] = None
|
||||
replacement: Optional[Dict[str, Any]] = None
|
||||
if preserve_compaction_handoff or expected_target_content is not None:
|
||||
if not target_row.get("active"):
|
||||
raise ValueError("rewind target is not active")
|
||||
from agent.context_compressor import split_user_originated_turn
|
||||
|
||||
rewound: List[int] = []
|
||||
split_target = target_row.copy()
|
||||
split_target["content"] = self._decode_content(
|
||||
split_target.get("content")
|
||||
)
|
||||
split_target["display_metadata"] = self._decode_display_metadata(
|
||||
split_target.get("display_metadata")
|
||||
)
|
||||
handoff, live_view = split_user_originated_turn(split_target)
|
||||
if live_view is None:
|
||||
raise ValueError("rewind target is not a user-originated turn")
|
||||
live_content = live_view.get("content")
|
||||
if isinstance(live_content, str):
|
||||
live_content = sanitize_context(live_content).strip()
|
||||
if (
|
||||
expected_target_content is not None
|
||||
and live_content != expected_target_content
|
||||
):
|
||||
raise RuntimeError(
|
||||
"rewind target changed before it could be persisted"
|
||||
)
|
||||
if preserve_compaction_handoff and handoff is None:
|
||||
raise ValueError(
|
||||
"preserve_compaction_handoff requires an active composite carrier"
|
||||
)
|
||||
replacement = handoff if preserve_compaction_handoff else None
|
||||
|
||||
def _do(conn):
|
||||
cursor = conn.execute(
|
||||
"SELECT id FROM messages "
|
||||
"WHERE session_id = ? AND id >= ? AND active = 1",
|
||||
@@ -12188,28 +12494,48 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
f"UPDATE messages SET active = 0 WHERE id IN ({placeholders})",
|
||||
ids,
|
||||
)
|
||||
if replacement is not None:
|
||||
self._insert_message_rows(conn, session_id, [replacement])
|
||||
inserted = conn.execute("SELECT last_insert_rowid()").fetchone()
|
||||
replacement_message_id = int(inserted[0])
|
||||
conn.execute(
|
||||
"UPDATE sessions SET rewind_count = COALESCE(rewind_count, 0) + 1 "
|
||||
"WHERE id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
return ids
|
||||
|
||||
rewound = self._execute_write(_do)
|
||||
|
||||
# 2) Compute new head id (largest still-active row id in session).
|
||||
with self._lock:
|
||||
head_row = self._conn.execute(
|
||||
message_count, tool_call_count = self._active_transcript_counts(
|
||||
conn, session_id
|
||||
)
|
||||
conn.execute(
|
||||
"UPDATE sessions SET message_count = ?, tool_call_count = ? "
|
||||
"WHERE id = ?",
|
||||
(message_count, tool_call_count, session_id),
|
||||
)
|
||||
head_row = conn.execute(
|
||||
"SELECT MAX(id) FROM messages WHERE session_id = ? AND active = 1",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
new_head_id = head_row[0] if head_row and head_row[0] is not None else None
|
||||
new_head_id = (
|
||||
head_row[0] if head_row and head_row[0] is not None else None
|
||||
)
|
||||
return target_row, ids, new_head_id, replacement_message_id
|
||||
|
||||
return {
|
||||
target_row, rewound, new_head_id, replacement_message_id = (
|
||||
self._execute_write(_do)
|
||||
)
|
||||
|
||||
# Decode content for callers (prefill the prompt buffer) without a
|
||||
# second fallible database operation after the transaction commits.
|
||||
target_row["content"] = self._decode_content(target_row.get("content"))
|
||||
|
||||
result = {
|
||||
"rewound_count": len(rewound),
|
||||
"target_message": target_row,
|
||||
"new_head_id": new_head_id,
|
||||
}
|
||||
if preserve_compaction_handoff:
|
||||
result["replacement_message_id"] = replacement_message_id
|
||||
return result
|
||||
|
||||
def restore_rewound(self, session_id: str, since_message_id: int) -> int:
|
||||
"""Mark inactive messages with id >= *since_message_id* active again.
|
||||
|
||||
@@ -162,6 +162,7 @@ from agent.usage_pricing import normalize_usage
|
||||
from agent.context_compressor import ( # noqa: F401
|
||||
COMPRESSED_SUMMARY_METADATA_KEY,
|
||||
ContextCompressor,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from agent.retry_utils import jittered_backoff # noqa: F401
|
||||
from agent.prompt_builder import ( # noqa: F401 # re-exported via _ra() / mock.patch("run_agent.<name>") / from run_agent import <name>
|
||||
@@ -2372,6 +2373,7 @@ class AIAgent:
|
||||
"hidden"
|
||||
if (
|
||||
msg.get(COMPRESSED_SUMMARY_METADATA_KEY)
|
||||
and user_originated_turn_view(msg) is None
|
||||
and (
|
||||
ContextCompressor.classify_summary_content(
|
||||
msg.get("content")
|
||||
|
||||
@@ -10,6 +10,8 @@ the already-completed work.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_compressor import (
|
||||
COMPRESSED_SUMMARY_HAS_USER_TURN_KEY,
|
||||
COMPRESSED_SUMMARY_METADATA_KEY,
|
||||
@@ -20,13 +22,18 @@ from agent.context_compressor import (
|
||||
_MERGED_PRIOR_CONTEXT_HEADER,
|
||||
_MERGED_SUMMARY_DELIMITER,
|
||||
_SUMMARY_END_MARKER,
|
||||
history_before_user_originated_turn,
|
||||
is_compaction_summary_message,
|
||||
is_user_originated_turn,
|
||||
reference_handoff_would_drive_next_model_call,
|
||||
retryable_user_text,
|
||||
split_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from agent.conversation_loop import (
|
||||
_should_skip_model_call_for_reference_handoff,
|
||||
)
|
||||
from agent.agent_runtime_helpers import repair_message_sequence
|
||||
from agent.turn_context import reanchor_current_turn_user_idx
|
||||
|
||||
|
||||
@@ -42,6 +49,12 @@ def _standalone_handoff(task: str = "finish the already-done refactor") -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _composite_handoff(ask: str = "REAL ASK") -> dict:
|
||||
handoff = _standalone_handoff()
|
||||
handoff["content"] += f"\n\n{ask}"
|
||||
return handoff
|
||||
|
||||
|
||||
def _merged_assistant_carrier(
|
||||
*, finish_reason: str = "stop", tool_calls: list[dict] | None = None
|
||||
) -> dict:
|
||||
@@ -230,6 +243,52 @@ class TestUserOriginatedTurnPredicate:
|
||||
def test_plain_user_is_originated(self):
|
||||
assert is_user_originated_turn({"role": "user", "content": "hello"}) is True
|
||||
|
||||
def test_force_user_leading_carrier_is_originated_via_live_view(self):
|
||||
carrier = _composite_handoff()
|
||||
|
||||
handoff, live = split_user_originated_turn(carrier)
|
||||
|
||||
assert is_user_originated_turn(carrier) is True
|
||||
assert user_originated_turn_view(carrier)["content"] == "REAL ASK"
|
||||
assert live["content"] == "REAL ASK"
|
||||
assert handoff["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in handoff["content"]
|
||||
|
||||
def test_hidden_legacy_carrier_still_projects_but_typed_carrier_does_not(self):
|
||||
hidden = {**_composite_handoff(), "display_kind": "hidden"}
|
||||
typed = {**_composite_handoff(), "display_kind": "auto_continue"}
|
||||
|
||||
assert user_originated_turn_view(hidden)["content"] == "REAL ASK"
|
||||
assert user_originated_turn_view(typed) is None
|
||||
|
||||
def test_rewind_prefix_preserves_only_the_carrier_scaffold(self):
|
||||
messages = [
|
||||
{"role": "assistant", "content": "prefix"},
|
||||
_composite_handoff(),
|
||||
{"role": "assistant", "content": "failed"},
|
||||
]
|
||||
|
||||
prefix, live = history_before_user_originated_turn(messages, 1)
|
||||
|
||||
assert live["content"] == "REAL ASK"
|
||||
assert prefix[0] == messages[0]
|
||||
assert prefix[1]["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in prefix[1]["content"]
|
||||
assert messages[1]["content"].endswith("REAL ASK")
|
||||
|
||||
def test_live_projection_keeps_reactions_but_drops_wrapper_metadata(self):
|
||||
carrier = _composite_handoff()
|
||||
carrier["display_metadata"] = {
|
||||
"reactions": [{"emoji": "👍", "author": "user"}],
|
||||
"synthetic_only": "do not project",
|
||||
}
|
||||
|
||||
live = user_originated_turn_view(carrier)
|
||||
|
||||
assert live["display_metadata"] == {
|
||||
"reactions": [{"emoji": "👍", "author": "user"}]
|
||||
}
|
||||
|
||||
def test_display_kind_hidden_not_originated(self):
|
||||
assert (
|
||||
is_user_originated_turn(
|
||||
@@ -266,6 +325,56 @@ class TestReanchorSkipsHandoffFallback:
|
||||
# Exact content rewritten by merge — fall back must not land on handoff.
|
||||
assert reanchor_current_turn_user_idx(messages, "rewritten ask") == 0
|
||||
|
||||
def test_exact_live_projection_reanchors_to_composite_carrier(self):
|
||||
messages = [
|
||||
{"role": "user", "content": "older ask"},
|
||||
_composite_handoff("rewritten ask"),
|
||||
]
|
||||
|
||||
assert reanchor_current_turn_user_idx(messages, "rewritten ask") == 1
|
||||
|
||||
|
||||
class TestRetryableUserText:
|
||||
def test_losslessly_flattens_text_only_parts(self):
|
||||
assert retryable_user_text(
|
||||
[
|
||||
{"type": "text", "text": "hello "},
|
||||
{"type": "input_text", "text": "world"},
|
||||
]
|
||||
) == "hello world"
|
||||
|
||||
def test_rejects_structured_media_parts(self):
|
||||
with pytest.raises(ValueError, match="media or unknown"):
|
||||
retryable_user_text(
|
||||
[{"type": "image_url", "image_url": {"url": "data:image/png;base64,x"}}]
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text",
|
||||
[
|
||||
"Explain why  renders",
|
||||
"Review @file:README.md and this data:image/png example",
|
||||
"Explain the literal [screenshot] marker",
|
||||
"Compare [file: README.md] with [image|ybres:RID]",
|
||||
],
|
||||
)
|
||||
def test_keeps_ordinary_text_that_uses_media_like_syntax(self, text):
|
||||
assert retryable_user_text(text) == text
|
||||
|
||||
|
||||
class TestCarrierAlternationRepair:
|
||||
def test_fresh_user_does_not_mutate_persisted_carrier_dict(self):
|
||||
carrier = _composite_handoff()
|
||||
original = carrier.copy()
|
||||
fresh = {"role": "user", "content": "NEXT ASK"}
|
||||
messages = [carrier, fresh]
|
||||
|
||||
assert repair_message_sequence(None, messages) == 0
|
||||
|
||||
assert messages == [carrier, fresh]
|
||||
assert messages[0] is carrier
|
||||
assert carrier == original
|
||||
|
||||
|
||||
class TestNoToolCallsWithoutLaterRealUser:
|
||||
def test_historical_snapshot_alone_is_not_actionable(self):
|
||||
|
||||
+330
-1
@@ -1,10 +1,42 @@
|
||||
"""Regression tests for CLI /retry history replacement semantics."""
|
||||
"""Regression tests for CLI /retry and carrier-aware rewind semantics."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
from tests.cli.test_cli_init import _make_cli
|
||||
|
||||
|
||||
def _composite_carrier(ask="REAL ASK"):
|
||||
return {
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}\n\n{ask}"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _message_rows(db, session_id):
|
||||
rows = db._conn.execute(
|
||||
"SELECT id, content, active FROM messages "
|
||||
"WHERE session_id = ? ORDER BY id",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [tuple(row) for row in rows]
|
||||
|
||||
|
||||
def test_retry_last_truncates_history_before_requeueing_message():
|
||||
cli = _make_cli()
|
||||
cli._session_db = None
|
||||
cli.conversation_history = [
|
||||
{"role": "user", "content": "first"},
|
||||
{"role": "assistant", "content": "one"},
|
||||
@@ -31,6 +63,7 @@ def test_retry_last_truncates_history_before_requeueing_message():
|
||||
|
||||
def test_process_command_retry_requeues_original_message_not_retry_command():
|
||||
cli = _make_cli()
|
||||
cli._session_db = None
|
||||
queued = []
|
||||
|
||||
class _Queue:
|
||||
@@ -47,3 +80,299 @@ def test_process_command_retry_requeues_original_message_not_retry_command():
|
||||
|
||||
assert queued == ["retry me"]
|
||||
assert cli.conversation_history == []
|
||||
|
||||
|
||||
def test_retry_fails_closed_when_warm_and_durable_targets_differ(tmp_path):
|
||||
cli = _make_cli()
|
||||
cli._session_db.close()
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
cli._session_db = db
|
||||
cli.session_id = "cli-target-mismatch"
|
||||
db.create_session(cli.session_id, source="cli")
|
||||
db.append_message(cli.session_id, "user", "DURABLE ASK")
|
||||
db.append_message(cli.session_id, "assistant", "old answer")
|
||||
|
||||
history = [
|
||||
{"role": "user", "content": "WARM ASK"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
cli.conversation_history = history
|
||||
cli._pending_input = MagicMock()
|
||||
before_rows = _message_rows(db, cli.session_id)
|
||||
|
||||
cli.process_command("/retry")
|
||||
|
||||
cli._pending_input.put.assert_not_called()
|
||||
assert cli.conversation_history is history
|
||||
assert _message_rows(db, cli.session_id) == before_rows
|
||||
db.close()
|
||||
|
||||
|
||||
def test_retry_fails_closed_when_transcript_changes_after_snapshot(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
cli = _make_cli()
|
||||
cli._session_db.close()
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
sibling = SessionDB(db_path=db.db_path)
|
||||
cli._session_db = db
|
||||
cli.session_id = "cli-cas-race"
|
||||
db.create_session(cli.session_id, source="cli")
|
||||
db.append_message(cli.session_id, "user", "RETRY ME")
|
||||
db.append_message(cli.session_id, "assistant", "failed answer")
|
||||
history = db.get_messages_as_conversation(cli.session_id)
|
||||
cli.conversation_history = history
|
||||
original_rewind = db.rewind_to_message
|
||||
|
||||
def _append_then_rewind(*args, **kwargs):
|
||||
sibling.append_message(cli.session_id, "assistant", "concurrent tail")
|
||||
return original_rewind(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(db, "rewind_to_message", _append_then_rewind)
|
||||
|
||||
assert cli.retry_last() is None
|
||||
|
||||
assert cli.conversation_history is history
|
||||
rows = db._conn.execute(
|
||||
"SELECT content, active FROM messages "
|
||||
"WHERE session_id = ? ORDER BY id",
|
||||
(cli.session_id,),
|
||||
).fetchall()
|
||||
assert [tuple(row) for row in rows] == [
|
||||
("RETRY ME", 1),
|
||||
("failed answer", 1),
|
||||
("concurrent tail", 1),
|
||||
]
|
||||
sibling.close()
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", ["retry", "undo"])
|
||||
def test_rewind_matches_warm_raw_carrier_to_durable_sanitized_sidecar(
|
||||
tmp_path, command
|
||||
):
|
||||
from agent.memory_manager import sanitize_context
|
||||
|
||||
cli = _make_cli()
|
||||
cli._session_db.close()
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
cli._session_db = db
|
||||
cli.session_id = f"cli-sanitized-{command}"
|
||||
db.create_session(cli.session_id, source="cli")
|
||||
raw_carrier = _composite_carrier(
|
||||
" REAL ASK\n\n<memory-context>\nprivate\n</memory-context> "
|
||||
)["content"]
|
||||
db.append_message(
|
||||
cli.session_id,
|
||||
"user",
|
||||
sanitize_context(raw_carrier).strip(),
|
||||
api_content=raw_carrier,
|
||||
)
|
||||
db.append_message(cli.session_id, "assistant", "failed answer")
|
||||
durable = db.get_messages_as_conversation(cli.session_id)
|
||||
cli.conversation_history = [
|
||||
{"role": "user", "content": raw_carrier},
|
||||
durable[1],
|
||||
]
|
||||
cli._pending_input = MagicMock()
|
||||
cli._prefill_input_buffer = MagicMock()
|
||||
|
||||
if command == "retry":
|
||||
cli.process_command("/retry")
|
||||
cli._pending_input.put.assert_called_once_with("REAL ASK")
|
||||
else:
|
||||
cli.undo_last()
|
||||
cli._prefill_input_buffer.assert_called_once_with("REAL ASK")
|
||||
|
||||
assert len(cli.conversation_history) == 1
|
||||
scaffold = cli.conversation_history[0]
|
||||
assert scaffold["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in scaffold["content"]
|
||||
active = db.get_messages_as_conversation(cli.session_id, include_row_ids=True)
|
||||
assert active[0]["_row_id"] == scaffold["_row_id"]
|
||||
assert active[0]["content"] == scaffold["content"]
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", ["retry", "undo"])
|
||||
@pytest.mark.parametrize("prefix_kind", ["buried_ephemeral", "old_media"])
|
||||
def test_rewind_keeps_the_richer_warm_prefix_after_validating_the_target(
|
||||
tmp_path, command, prefix_kind
|
||||
):
|
||||
cli = _make_cli()
|
||||
cli._session_db.close()
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
cli._session_db = db
|
||||
cli.session_id = f"cli-projection-{prefix_kind}-{command}"
|
||||
db.create_session(cli.session_id, source="cli")
|
||||
|
||||
if prefix_kind == "buried_ephemeral":
|
||||
history = [
|
||||
{"role": "user", "content": "OLDER ASK"},
|
||||
{"role": "assistant", "content": "candidate answer"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "[System: verify before stopping]",
|
||||
"_verification_stop_synthetic": True,
|
||||
},
|
||||
{"role": "assistant", "content": "verified answer"},
|
||||
{"role": "user", "content": "PLAIN TARGET"},
|
||||
{"role": "assistant", "content": "failed answer"},
|
||||
]
|
||||
durable_prefix = [
|
||||
("user", "OLDER ASK"),
|
||||
("assistant", "candidate answer"),
|
||||
("assistant", "verified answer"),
|
||||
]
|
||||
expected_prefix = ["OLDER ASK", "candidate answer", "verified answer"]
|
||||
expected_active = [1, 1, 1, 0, 0]
|
||||
else:
|
||||
media_content = [
|
||||
{"type": "text", "text": "OLDER ASK"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png,AA"}},
|
||||
]
|
||||
history = [
|
||||
{"role": "user", "content": media_content},
|
||||
{"role": "assistant", "content": "older answer"},
|
||||
{"role": "user", "content": "PLAIN TARGET"},
|
||||
{"role": "assistant", "content": "failed answer"},
|
||||
]
|
||||
durable_prefix = [
|
||||
("user", "OLDER ASK\n[screenshot]"),
|
||||
("assistant", "older answer"),
|
||||
]
|
||||
expected_prefix = [media_content, "older answer"]
|
||||
expected_active = [1, 1, 0, 0]
|
||||
|
||||
for role, content in durable_prefix:
|
||||
db.append_message(cli.session_id, role, content)
|
||||
db.append_message(cli.session_id, "user", "PLAIN TARGET")
|
||||
db.append_message(cli.session_id, "assistant", "failed answer")
|
||||
cli.conversation_history = history
|
||||
cli._pending_input = MagicMock()
|
||||
cli._prefill_input_buffer = MagicMock()
|
||||
|
||||
if command == "retry":
|
||||
cli.process_command("/retry")
|
||||
cli._pending_input.put.assert_called_once_with("PLAIN TARGET")
|
||||
else:
|
||||
cli.undo_last()
|
||||
cli._prefill_input_buffer.assert_called_once_with("PLAIN TARGET")
|
||||
|
||||
assert [message.get("content") for message in cli.conversation_history] == (
|
||||
expected_prefix
|
||||
)
|
||||
assert [row[2] for row in _message_rows(db, cli.session_id)] == expected_active
|
||||
db.close()
|
||||
|
||||
|
||||
def test_retry_last_durably_preserves_composite_carrier_scaffold(tmp_path):
|
||||
cli = _make_cli()
|
||||
cli._session_db.close()
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
cli._session_db = db
|
||||
cli.session_id = "cli-carrier-retry"
|
||||
db.create_session(cli.session_id, source="cli")
|
||||
db.append_message(cli.session_id, "user", _composite_carrier()["content"])
|
||||
db.append_message(cli.session_id, "assistant", "failed answer")
|
||||
cli.conversation_history = db.get_messages_as_conversation(cli.session_id)
|
||||
old_history = cli.conversation_history
|
||||
cli.agent = SimpleNamespace(
|
||||
_session_messages=old_history,
|
||||
_last_flushed_db_idx=len(old_history),
|
||||
_db_flush_scan_prefix=list(old_history),
|
||||
)
|
||||
|
||||
retry_msg = cli.retry_last()
|
||||
|
||||
assert retry_msg == "REAL ASK"
|
||||
assert len(cli.conversation_history) == 1
|
||||
scaffold = cli.conversation_history[0]
|
||||
assert scaffold["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in scaffold["content"]
|
||||
assert scaffold["_db_persisted"] is True
|
||||
active = db.get_messages_as_conversation(cli.session_id, include_row_ids=True)
|
||||
assert len(active) == 1
|
||||
assert active[0]["content"] == scaffold["content"]
|
||||
assert active[0]["_row_id"] == scaffold["_row_id"]
|
||||
assert cli.agent._session_messages is cli.conversation_history
|
||||
assert cli.agent._last_flushed_db_idx == 1
|
||||
assert cli.agent._db_flush_scan_prefix == cli.conversation_history
|
||||
db.close()
|
||||
|
||||
|
||||
def test_retry_last_rejects_media_before_db_or_memory_mutation():
|
||||
cli = _make_cli()
|
||||
db = MagicMock()
|
||||
cli._session_db = db
|
||||
history = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "look again"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,AA"}},
|
||||
],
|
||||
},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
cli.conversation_history = history
|
||||
|
||||
assert cli.retry_last() is None
|
||||
assert cli.conversation_history is history
|
||||
db.get_messages_as_conversation.assert_not_called()
|
||||
db.rewind_to_message.assert_not_called()
|
||||
|
||||
|
||||
def test_retry_last_db_failure_leaves_warm_history_unchanged():
|
||||
cli = _make_cli()
|
||||
db = MagicMock()
|
||||
db.get_messages_as_conversation.side_effect = OSError("db unavailable")
|
||||
cli._session_db = db
|
||||
history = [
|
||||
{"role": "user", "content": "retry me"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
cli.conversation_history = history
|
||||
|
||||
assert cli.retry_last() is None
|
||||
assert cli.conversation_history is history
|
||||
|
||||
|
||||
def test_undo_last_prefills_live_text_and_retains_durable_scaffold(tmp_path):
|
||||
cli = _make_cli()
|
||||
cli._session_db.close()
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
cli._session_db = db
|
||||
cli.session_id = "cli-carrier-undo"
|
||||
db.create_session(cli.session_id, source="cli")
|
||||
db.append_message(cli.session_id, "user", "older ask")
|
||||
db.append_message(cli.session_id, "assistant", "older answer")
|
||||
db.append_message(cli.session_id, "user", _composite_carrier()["content"])
|
||||
db.append_message(cli.session_id, "assistant", "failed answer")
|
||||
cli.conversation_history = db.get_messages_as_conversation(cli.session_id)
|
||||
cli._prefill_input_buffer = MagicMock()
|
||||
cli.agent = SimpleNamespace(
|
||||
_session_messages=cli.conversation_history,
|
||||
_last_flushed_db_idx=len(cli.conversation_history),
|
||||
_db_flush_scan_prefix=list(cli.conversation_history),
|
||||
_invalidate_system_prompt=MagicMock(),
|
||||
_memory_manager=None,
|
||||
)
|
||||
|
||||
cli.undo_last()
|
||||
|
||||
cli._prefill_input_buffer.assert_called_once_with("REAL ASK")
|
||||
assert [m.get("content") for m in cli.conversation_history[:2]] == [
|
||||
"older ask",
|
||||
"older answer",
|
||||
]
|
||||
scaffold = cli.conversation_history[2]
|
||||
assert scaffold["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in scaffold["content"]
|
||||
assert scaffold["_db_persisted"] is True
|
||||
active = db.get_messages_as_conversation(cli.session_id, include_row_ids=True)
|
||||
assert active[2]["_row_id"] == scaffold["_row_id"]
|
||||
assert active[2]["content"] == scaffold["content"]
|
||||
assert cli.agent._session_messages is cli.conversation_history
|
||||
assert cli.agent._last_flushed_db_idx == 3
|
||||
db.close()
|
||||
|
||||
@@ -1,15 +1,226 @@
|
||||
"""Regression tests for /retry replacement semantics."""
|
||||
"""Regression tests for /retry replacement and carrier-aware undo semantics."""
|
||||
|
||||
import os
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from gateway.config import GatewayConfig
|
||||
from gateway.platforms.base import MessageEvent, MessageType
|
||||
from gateway.run import GatewayRunner
|
||||
from gateway.session import SessionStore
|
||||
|
||||
|
||||
def _composite_carrier(ask="REAL ASK"):
|
||||
return {
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}\n\n{ask}"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _seed_pending_recovery(store, session_id):
|
||||
pending = {"role": "assistant", "content": "pending recovery answer"}
|
||||
store._dirty_transcripts[session_id] = [dict(pending)]
|
||||
store._transcript_append_failures[session_id] = 3
|
||||
return pending
|
||||
|
||||
|
||||
def test_rewrite_transcript_keeps_pending_recovery_state_when_lease_rejects(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
import hermes_state
|
||||
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig())
|
||||
session_id = "rewrite-pending-lease"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(session_id, "user", "old ask")
|
||||
pending = _seed_pending_recovery(store, session_id)
|
||||
before = store._db.get_messages(session_id, include_inactive=True)
|
||||
holder = f"pid={os.getpid()}:turn=foreign"
|
||||
assert store._db.try_acquire_session_turn_lease(
|
||||
session_id, holder, ttl_seconds=60
|
||||
)
|
||||
|
||||
assert not store.rewrite_transcript(
|
||||
session_id,
|
||||
[{"role": "user", "content": "replacement ask"}],
|
||||
active_only=True,
|
||||
reject_active_turn_lease=True,
|
||||
)
|
||||
|
||||
assert store._db.get_messages(session_id, include_inactive=True) == before
|
||||
assert store._dirty_transcripts[session_id] == [pending]
|
||||
assert store._transcript_append_failures[session_id] == 3
|
||||
|
||||
store._db.release_session_turn_lease(session_id, holder)
|
||||
assert store.rewrite_transcript(
|
||||
session_id,
|
||||
[{"role": "user", "content": "replacement ask"}],
|
||||
active_only=True,
|
||||
reject_active_turn_lease=True,
|
||||
)
|
||||
assert session_id not in store._dirty_transcripts
|
||||
assert session_id not in store._transcript_append_failures
|
||||
|
||||
|
||||
def test_rewind_session_keeps_pending_recovery_state_when_lease_rejects(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
import hermes_state
|
||||
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig())
|
||||
session_id = "rewind-pending-lease"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(session_id, "user", _composite_carrier()["content"])
|
||||
store._db.append_message(session_id, "assistant", "old answer")
|
||||
pending = _seed_pending_recovery(store, session_id)
|
||||
before = store._db.get_messages(session_id, include_inactive=True)
|
||||
holder = f"pid={os.getpid()}:turn=foreign"
|
||||
assert store._db.try_acquire_session_turn_lease(
|
||||
session_id, holder, ttl_seconds=60
|
||||
)
|
||||
|
||||
assert (
|
||||
store.rewind_session(session_id, require_retryable_composite=True) is None
|
||||
)
|
||||
|
||||
assert store._db.get_messages(session_id, include_inactive=True) == before
|
||||
assert store._dirty_transcripts[session_id] == [pending]
|
||||
assert store._transcript_append_failures[session_id] == 3
|
||||
|
||||
store._db.release_session_turn_lease(session_id, holder)
|
||||
result = store.rewind_session(
|
||||
session_id, require_retryable_composite=True
|
||||
)
|
||||
assert result is not None
|
||||
assert result["target_text"] == "REAL ASK"
|
||||
assert session_id not in store._dirty_transcripts
|
||||
assert session_id not in store._transcript_append_failures
|
||||
|
||||
|
||||
def test_rewind_session_surfaces_unretryable_media_before_mutation(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
import hermes_state
|
||||
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig())
|
||||
session_id = "rewind-composite-media"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(
|
||||
session_id,
|
||||
"user",
|
||||
[
|
||||
{"type": "text", "text": _composite_carrier()["content"]},
|
||||
{"type": "image_url", "image_url": {"url": "image"}},
|
||||
],
|
||||
)
|
||||
store._db.append_message(session_id, "assistant", "old answer")
|
||||
before = store._db.get_messages(session_id, include_inactive=True)
|
||||
|
||||
with pytest.raises(ValueError, match="media or unknown content"):
|
||||
store.rewind_session(session_id, require_retryable_composite=True)
|
||||
|
||||
assert store._db.get_messages(session_id, include_inactive=True) == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("operation", ["rewrite", "rewind"])
|
||||
def test_transcript_mutation_serializes_pending_queue_drain(
|
||||
operation, tmp_path, monkeypatch
|
||||
):
|
||||
import hermes_state
|
||||
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig())
|
||||
session_id = f"serialized-{operation}"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(session_id, "user", _composite_carrier()["content"])
|
||||
store._db.append_message(session_id, "assistant", "old answer")
|
||||
_seed_pending_recovery(store, session_id)
|
||||
|
||||
mutation_entered = threading.Event()
|
||||
release_mutation = threading.Event()
|
||||
append_started = threading.Event()
|
||||
append_done = threading.Event()
|
||||
errors = []
|
||||
|
||||
if operation == "rewrite":
|
||||
original_mutation = store._db.replace_messages
|
||||
|
||||
def gated_mutation(*args, **kwargs):
|
||||
mutation_entered.set()
|
||||
assert release_mutation.wait(timeout=5)
|
||||
return original_mutation(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(store._db, "replace_messages", gated_mutation)
|
||||
|
||||
def mutate():
|
||||
assert store.rewrite_transcript(
|
||||
session_id,
|
||||
[{"role": "user", "content": "replacement ask"}],
|
||||
active_only=True,
|
||||
)
|
||||
|
||||
else:
|
||||
original_mutation = store._db.rewind_to_message
|
||||
|
||||
def gated_mutation(*args, **kwargs):
|
||||
mutation_entered.set()
|
||||
assert release_mutation.wait(timeout=5)
|
||||
return original_mutation(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(store._db, "rewind_to_message", gated_mutation)
|
||||
|
||||
def mutate():
|
||||
assert store.rewind_session(session_id) is not None
|
||||
|
||||
def run_mutation():
|
||||
try:
|
||||
mutate()
|
||||
except BaseException as exc: # surface worker failures in the test thread
|
||||
errors.append(exc)
|
||||
|
||||
def append_after_mutation_starts():
|
||||
append_started.set()
|
||||
try:
|
||||
store.append_to_transcript(
|
||||
session_id,
|
||||
{"role": "assistant", "content": "concurrent answer"},
|
||||
)
|
||||
except BaseException as exc: # surface worker failures in the test thread
|
||||
errors.append(exc)
|
||||
finally:
|
||||
append_done.set()
|
||||
|
||||
mutation_thread = threading.Thread(target=run_mutation)
|
||||
mutation_thread.start()
|
||||
assert mutation_entered.wait(timeout=5)
|
||||
append_thread = threading.Thread(target=append_after_mutation_starts)
|
||||
append_thread.start()
|
||||
assert append_started.wait(timeout=5)
|
||||
assert not append_done.wait(timeout=0.1)
|
||||
|
||||
release_mutation.set()
|
||||
mutation_thread.join(timeout=5)
|
||||
append_thread.join(timeout=5)
|
||||
assert not mutation_thread.is_alive()
|
||||
assert not append_thread.is_alive()
|
||||
assert errors == []
|
||||
assert store.load_transcript(session_id)[-1]["content"] == "concurrent answer"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_replaces_last_user_turn_in_transcript(tmp_path, monkeypatch):
|
||||
# Pin DEFAULT_DB_PATH so SessionDB() doesn't write to the real ~/.hermes/state.db.
|
||||
@@ -67,6 +278,228 @@ async def test_gateway_retry_replaces_last_user_turn_in_transcript(tmp_path, mon
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_redispatches_live_carrier_text_and_keeps_scaffold(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
import hermes_state
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
|
||||
config = GatewayConfig()
|
||||
store = SessionStore(sessions_dir=tmp_path, config=config)
|
||||
session_id = "retry-carrier-session"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(session_id, "user", "older ask")
|
||||
store._db.append_message(session_id, "assistant", "older answer")
|
||||
store._db.append_message(session_id, "user", _composite_carrier()["content"])
|
||||
store._db.append_message(session_id, "assistant", "failed answer")
|
||||
|
||||
gw = GatewayRunner.__new__(GatewayRunner)
|
||||
gw.config = config
|
||||
gw.session_store = store
|
||||
session_entry = MagicMock(session_id=session_id, last_prompt_tokens=123)
|
||||
gw.session_store.get_or_create_session = MagicMock(return_value=session_entry)
|
||||
|
||||
async def fake_handle_message(event):
|
||||
assert event.text == "REAL ASK"
|
||||
active = store.load_transcript(session_id)
|
||||
assert [m.get("content") for m in active[:2]] == ["older ask", "older answer"]
|
||||
scaffold = active[2]
|
||||
assert scaffold["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in scaffold["content"]
|
||||
return "new answer"
|
||||
|
||||
gw._handle_message = AsyncMock(side_effect=fake_handle_message)
|
||||
|
||||
result = await gw._handle_retry_command(
|
||||
MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock())
|
||||
)
|
||||
|
||||
assert result == "new answer"
|
||||
assert session_entry.last_prompt_tokens == 0
|
||||
gw._handle_message.assert_awaited_once()
|
||||
archived = [
|
||||
row
|
||||
for row in store._db.get_messages(session_id, include_inactive=True)
|
||||
if not row["active"]
|
||||
]
|
||||
assert [row["content"] for row in archived] == [
|
||||
_composite_carrier()["content"],
|
||||
"failed answer",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_does_not_rewind_a_newer_plain_turn(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
"""The carrier selected for retry must still be latest at commit time."""
|
||||
import hermes_state
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
|
||||
config = GatewayConfig()
|
||||
store = SessionStore(sessions_dir=tmp_path, config=config)
|
||||
session_id = "retry-carrier-race-session"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(session_id, "user", _composite_carrier()["content"])
|
||||
store._db.append_message(session_id, "assistant", "failed answer")
|
||||
|
||||
gw = GatewayRunner.__new__(GatewayRunner)
|
||||
gw.config = config
|
||||
gw.session_store = store
|
||||
session_entry = MagicMock(session_id=session_id, last_prompt_tokens=123)
|
||||
gw.session_store.get_or_create_session = MagicMock(return_value=session_entry)
|
||||
original_rewind = store.rewind_session
|
||||
|
||||
def append_newer_turn_then_rewind(*args, **kwargs):
|
||||
store._db.append_message(session_id, "user", "newer ask")
|
||||
store._db.append_message(session_id, "assistant", "newer answer")
|
||||
return original_rewind(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(store, "rewind_session", append_newer_turn_then_rewind)
|
||||
gw._handle_message = AsyncMock()
|
||||
|
||||
result = await gw._handle_retry_command(
|
||||
MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock())
|
||||
)
|
||||
|
||||
assert result.startswith("Retry failed;")
|
||||
assert session_entry.last_prompt_tokens == 123
|
||||
gw._handle_message.assert_not_awaited()
|
||||
assert [
|
||||
message.get("content")
|
||||
for message in store.load_transcript(session_id)
|
||||
if message.get("role") == "user"
|
||||
] == [_composite_carrier()["content"], "newer ask"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_rejects_media_before_redispatch_or_token_reset():
|
||||
gw = GatewayRunner.__new__(GatewayRunner)
|
||||
backing_store = MagicMock()
|
||||
gw.session_store = backing_store
|
||||
session_entry = SimpleNamespace(session_id="sid", last_prompt_tokens=123)
|
||||
facade = SimpleNamespace(
|
||||
_store=backing_store,
|
||||
get_or_create_session=AsyncMock(return_value=session_entry),
|
||||
load_transcript=AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "look again"},
|
||||
{"type": "image_url", "image_url": {"url": "image"}},
|
||||
],
|
||||
},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
),
|
||||
rewrite_transcript=AsyncMock(return_value=True),
|
||||
)
|
||||
gw._async_session_store = facade
|
||||
gw._handle_message = AsyncMock()
|
||||
|
||||
result = await gw._handle_retry_command(
|
||||
MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock())
|
||||
)
|
||||
|
||||
assert result.startswith("Cannot retry that message safely:")
|
||||
assert session_entry.last_prompt_tokens == 123
|
||||
gw._handle_message.assert_not_awaited()
|
||||
facade.rewrite_transcript.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_preserves_composite_media_diagnostic_from_store():
|
||||
gw = GatewayRunner.__new__(GatewayRunner)
|
||||
backing_store = MagicMock()
|
||||
gw.session_store = backing_store
|
||||
session_entry = SimpleNamespace(session_id="sid", last_prompt_tokens=123)
|
||||
facade = SimpleNamespace(
|
||||
_store=backing_store,
|
||||
get_or_create_session=AsyncMock(return_value=session_entry),
|
||||
load_transcript=AsyncMock(
|
||||
return_value=[
|
||||
_composite_carrier(),
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
),
|
||||
rewind_session=AsyncMock(
|
||||
side_effect=ValueError("retry does not support media content")
|
||||
),
|
||||
)
|
||||
gw._async_session_store = facade
|
||||
gw._handle_message = AsyncMock()
|
||||
|
||||
result = await gw._handle_retry_command(
|
||||
MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock())
|
||||
)
|
||||
|
||||
assert result == (
|
||||
"Cannot retry that message safely: retry does not support media content"
|
||||
)
|
||||
assert session_entry.last_prompt_tokens == 123
|
||||
gw._handle_message.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_stops_when_transcript_rewrite_fails():
|
||||
gw = GatewayRunner.__new__(GatewayRunner)
|
||||
backing_store = MagicMock()
|
||||
gw.session_store = backing_store
|
||||
session_entry = SimpleNamespace(session_id="sid", last_prompt_tokens=123)
|
||||
facade = SimpleNamespace(
|
||||
_store=backing_store,
|
||||
get_or_create_session=AsyncMock(return_value=session_entry),
|
||||
load_transcript=AsyncMock(
|
||||
return_value=[
|
||||
{"role": "user", "content": "retry me"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
),
|
||||
rewrite_transcript=AsyncMock(return_value=False),
|
||||
)
|
||||
gw._async_session_store = facade
|
||||
gw._handle_message = AsyncMock()
|
||||
|
||||
result = await gw._handle_retry_command(
|
||||
MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock())
|
||||
)
|
||||
|
||||
assert result.startswith("Retry failed;")
|
||||
assert session_entry.last_prompt_tokens == 123
|
||||
gw._handle_message.assert_not_awaited()
|
||||
facade.rewrite_transcript.assert_awaited_once()
|
||||
assert (
|
||||
facade.rewrite_transcript.await_args.kwargs["reject_active_turn_lease"]
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_gateway_undo_prefills_live_carrier_text_and_keeps_scaffold(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
import hermes_state
|
||||
monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db")
|
||||
|
||||
store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig())
|
||||
session_id = "undo-carrier-session"
|
||||
store._db.create_session(session_id=session_id, source="test")
|
||||
store._db.append_message(session_id, "user", _composite_carrier()["content"])
|
||||
store._db.append_message(session_id, "assistant", "failed answer")
|
||||
|
||||
result = store.rewind_session(session_id)
|
||||
|
||||
assert result["target_text"] == "REAL ASK"
|
||||
assert result["rewound_count"] == 2
|
||||
active = store._db.get_messages_as_conversation(
|
||||
session_id, include_row_ids=True
|
||||
)
|
||||
assert len(active) == 1
|
||||
assert active[0]["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in active[0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_retry_preserves_archived_compaction_rows_when_probe_fails(
|
||||
tmp_path, monkeypatch
|
||||
|
||||
@@ -54,3 +54,93 @@ def test_rewind_n_turns(store):
|
||||
assert len(store.load_transcript(sid)) == 2 # q1,a1
|
||||
|
||||
|
||||
def test_rewind_pins_raw_active_ids_when_projection_hides_review_harness(store):
|
||||
sid = _seed(store, "gw-review-harness", turns=2)
|
||||
store._db.append_message(
|
||||
sid,
|
||||
"user",
|
||||
"Review the conversation above and update the skill library safely",
|
||||
)
|
||||
store._db.append_message(sid, "assistant", "curator-only reply")
|
||||
|
||||
# Legacy background-review rows are intentionally absent from replay, but
|
||||
# they remain physical active rows that the rewind CAS must pin.
|
||||
assert [message["content"] for message in store.load_transcript(sid)] == [
|
||||
"q1",
|
||||
"a1",
|
||||
"q2",
|
||||
"a2",
|
||||
]
|
||||
|
||||
result = store.rewind_session(sid)
|
||||
|
||||
assert result is not None
|
||||
assert result["target_text"] == "q2"
|
||||
assert result["rewound_count"] == 4
|
||||
assert [message["content"] for message in store.load_transcript(sid)] == [
|
||||
"q1",
|
||||
"a1",
|
||||
]
|
||||
|
||||
|
||||
def test_rewind_fails_closed_when_transcript_changes_after_snapshot(
|
||||
store, monkeypatch
|
||||
):
|
||||
sid = _seed(store, "gw-cas", turns=2)
|
||||
sibling = SessionDB(db_path=store._db.db_path)
|
||||
original_rewind = store._db.rewind_to_message
|
||||
|
||||
def _append_then_rewind(*args, **kwargs):
|
||||
sibling.append_message(sid, "assistant", "concurrent tail")
|
||||
return original_rewind(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(store._db, "rewind_to_message", _append_then_rewind)
|
||||
|
||||
assert store.rewind_session(sid) is None
|
||||
|
||||
rows = store._db._conn.execute(
|
||||
"SELECT content, active FROM messages "
|
||||
"WHERE session_id = ? ORDER BY id",
|
||||
(sid,),
|
||||
).fetchall()
|
||||
assert [tuple(row) for row in rows] == [
|
||||
("q1", 1),
|
||||
("a1", 1),
|
||||
("q2", 1),
|
||||
("a2", 1),
|
||||
("concurrent tail", 1),
|
||||
]
|
||||
sibling.close()
|
||||
|
||||
|
||||
def test_rewind_fails_closed_when_new_turn_lands_after_id_snapshot(
|
||||
store, monkeypatch
|
||||
):
|
||||
sid = _seed(store, "gw-snapshot-order", turns=2)
|
||||
sibling = SessionDB(db_path=store._db.db_path)
|
||||
original_load = store._db.get_messages_as_conversation
|
||||
|
||||
def _load_then_append(*args, **kwargs):
|
||||
snapshot = original_load(*args, **kwargs)
|
||||
sibling.append_message(sid, "user", "q3-from-other-process")
|
||||
sibling.append_message(sid, "assistant", "a3-from-other-process")
|
||||
return snapshot
|
||||
|
||||
monkeypatch.setattr(store._db, "get_messages_as_conversation", _load_then_append)
|
||||
|
||||
assert store.rewind_session(sid) is None
|
||||
|
||||
rows = store._db._conn.execute(
|
||||
"SELECT content, active FROM messages "
|
||||
"WHERE session_id = ? ORDER BY id",
|
||||
(sid,),
|
||||
).fetchall()
|
||||
assert [tuple(row) for row in rows] == [
|
||||
("q1", 1),
|
||||
("a1", 1),
|
||||
("q2", 1),
|
||||
("a2", 1),
|
||||
("q3-from-other-process", 1),
|
||||
("a3-from-other-process", 1),
|
||||
]
|
||||
sibling.close()
|
||||
|
||||
@@ -2098,6 +2098,71 @@ class TestWebServerEndpoints:
|
||||
contents = [m["content"] for m in resp.json()["messages"]]
|
||||
assert contents == ["old q", "old a", "summary", "live q", "live a"]
|
||||
|
||||
def test_get_session_messages_projects_and_dedupes_composite_carrier(self):
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_MERGED_PRIOR_CONTEXT_HEADER,
|
||||
_MERGED_SUMMARY_DELIMITER,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
handoff = (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}"
|
||||
)
|
||||
carrier = f"{handoff}\n\nREAL ASK"
|
||||
assistant_carrier = (
|
||||
f"{_MERGED_PRIOR_CONTEXT_HEADER}\n"
|
||||
"real completed answer\n\n"
|
||||
f"{_MERGED_SUMMARY_DELIMITER}\n\n{handoff}"
|
||||
)
|
||||
db = SessionDB()
|
||||
try:
|
||||
db.create_session(session_id="compacted-carrier-display", source="desktop")
|
||||
db.append_message(
|
||||
"compacted-carrier-display",
|
||||
"user",
|
||||
"REAL ASK",
|
||||
timestamp=123.0,
|
||||
)
|
||||
db.archive_and_compact(
|
||||
"compacted-carrier-display",
|
||||
[{"role": "user", "content": carrier, "timestamp": 123.0}],
|
||||
)
|
||||
db.append_message(
|
||||
"compacted-carrier-display",
|
||||
"user",
|
||||
handoff,
|
||||
timestamp=124.0,
|
||||
)
|
||||
db.append_message(
|
||||
"compacted-carrier-display",
|
||||
"assistant",
|
||||
assistant_carrier,
|
||||
timestamp=125.0,
|
||||
)
|
||||
active_id = db.get_messages("compacted-carrier-display")[0]["id"]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
resp = self.client.get(
|
||||
"/api/sessions/compacted-carrier-display/messages"
|
||||
"?include_compacted=true"
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
messages = resp.json()["messages"]
|
||||
assert len(messages) == 3
|
||||
assert messages[0]["id"] == active_id
|
||||
assert messages[0]["content"] == carrier
|
||||
assert messages[0]["display_content"] == "REAL ASK"
|
||||
assert not messages[0].get("display_kind")
|
||||
assert messages[1]["content"] == handoff
|
||||
assert messages[1]["display_kind"] == "hidden"
|
||||
assert messages[2]["content"] == assistant_carrier
|
||||
assert messages[2]["display_content"] == "real completed answer"
|
||||
|
||||
def test_get_session_messages_latest_page_with_compacted_rows(self):
|
||||
"""The desktop's real read path (getLatestSessionMessages: limit +
|
||||
order=latest + include_compacted=true) pages back from the newest
|
||||
|
||||
@@ -0,0 +1,381 @@
|
||||
"""Transactional persistence contracts for composite compaction carriers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_MERGED_SUMMARY_DELIMITER,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from hermes_state import (
|
||||
CompressionSessionClosedError,
|
||||
SessionCompressionInProgressError,
|
||||
SessionDB,
|
||||
SessionTurnLeaseLostError,
|
||||
)
|
||||
|
||||
|
||||
def _carrier(ask: str = "REAL ASK") -> str:
|
||||
return (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}\n\n{ask}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
state = SessionDB(db_path=tmp_path / "state.db")
|
||||
yield state
|
||||
state.close()
|
||||
|
||||
|
||||
def _session_counts(db: SessionDB, session_id: str) -> tuple[int, int, int]:
|
||||
row = db._conn.execute(
|
||||
"SELECT message_count, tool_call_count, rewind_count "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
return row["message_count"], row["tool_call_count"], row["rewind_count"]
|
||||
|
||||
|
||||
def _row_state(db: SessionDB, session_id: str) -> list[tuple]:
|
||||
return [
|
||||
tuple(row)
|
||||
for row in db._conn.execute(
|
||||
"SELECT id, role, content, active, display_kind "
|
||||
"FROM messages WHERE session_id = ? ORDER BY id",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
]
|
||||
|
||||
|
||||
def _active_ids(db: SessionDB, session_id: str) -> list[int]:
|
||||
return [
|
||||
int(message["_row_id"])
|
||||
for message in db.get_messages_as_conversation(
|
||||
session_id, include_row_ids=True
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def test_composite_rewind_archives_tail_and_inserts_its_hidden_scaffold(db):
|
||||
sid = "carrier-rewind"
|
||||
db.create_session(sid, source="tui")
|
||||
db.append_message(sid, "user", "older ask")
|
||||
db.append_message(
|
||||
sid,
|
||||
"assistant",
|
||||
None,
|
||||
tool_calls=[{"id": "call-1", "function": {"name": "terminal"}}],
|
||||
)
|
||||
db.append_message(sid, "tool", "ok", tool_call_id="call-1")
|
||||
target_id = db.append_message(sid, "user", _carrier())
|
||||
db.append_message(sid, "assistant", "failed")
|
||||
expected_active_ids = _active_ids(db, sid)
|
||||
|
||||
result = db.rewind_to_message(
|
||||
sid,
|
||||
target_id,
|
||||
preserve_compaction_handoff=True,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content="REAL ASK",
|
||||
)
|
||||
|
||||
assert result["rewound_count"] == 2
|
||||
assert result["replacement_message_id"] == result["new_head_id"]
|
||||
active = db.get_messages_as_conversation(sid, include_row_ids=True)
|
||||
assert len(active) == 4
|
||||
assert active[-1]["_row_id"] == result["replacement_message_id"]
|
||||
assert active[-1]["display_kind"] == "hidden"
|
||||
assert SUMMARY_PREFIX in active[-1]["content"]
|
||||
assert "REAL ASK" not in active[-1]["content"]
|
||||
archived = db._conn.execute(
|
||||
"SELECT active FROM messages WHERE id IN (?, ?) ORDER BY id",
|
||||
(target_id, target_id + 1),
|
||||
).fetchall()
|
||||
assert [row[0] for row in archived] == [0, 0]
|
||||
assert _session_counts(db, sid) == (4, 1, 1)
|
||||
|
||||
|
||||
def test_lineage_display_prefers_tip_carrier_over_replayed_parent_ask(db):
|
||||
parent = "carrier-parent"
|
||||
child = "carrier-child"
|
||||
db.create_session(parent, source="tui")
|
||||
db.append_message(parent, "user", "REAL ASK")
|
||||
db.end_session(parent, "compression")
|
||||
db.create_session(child, source="tui", parent_session_id=parent)
|
||||
carrier_id = db.append_message(child, "user", _carrier())
|
||||
|
||||
model_history, display_history = db.get_resume_conversations(child)
|
||||
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
|
||||
visible_users = [
|
||||
user_originated_turn_view(message)
|
||||
for message in display_history
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
assert [message["content"] for message in visible_users] == ["REAL ASK"]
|
||||
assert display_history[-1]["_row_id"] == carrier_id
|
||||
assert model_history[-1]["_row_id"] == carrier_id
|
||||
assert db.get_ancestor_display_prefix(child) == []
|
||||
|
||||
|
||||
def test_lineage_display_dedupes_multimodal_ask_in_tip_carrier(db):
|
||||
parent = "media-carrier-parent"
|
||||
child = "media-carrier-child"
|
||||
ask = [
|
||||
{"type": "text", "text": "inspect this"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,AA=="}},
|
||||
]
|
||||
carrier = [
|
||||
{"type": "text", "text": f"{_carrier('')}\n"},
|
||||
*ask,
|
||||
]
|
||||
db.create_session(parent, source="tui")
|
||||
db.append_message(parent, "user", ask)
|
||||
db.end_session(parent, "compression")
|
||||
db.create_session(child, source="tui", parent_session_id=parent)
|
||||
carrier_id = db.append_message(child, "user", carrier)
|
||||
|
||||
_, display_history = db.get_resume_conversations(child)
|
||||
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
|
||||
visible_users = [
|
||||
user_originated_turn_view(message)
|
||||
for message in display_history
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
assert [message["content"] for message in visible_users] == [ask]
|
||||
assert display_history[-1]["_row_id"] == carrier_id
|
||||
|
||||
|
||||
def test_default_rewind_return_shape_and_active_counters_remain_compatible(db):
|
||||
sid = "default-rewind"
|
||||
db.create_session(sid, source="cli")
|
||||
db.append_message(sid, "user", "first")
|
||||
db.append_message(sid, "assistant", "answer")
|
||||
target_id = db.append_message(sid, "user", "second")
|
||||
db.append_message(
|
||||
sid,
|
||||
"assistant",
|
||||
None,
|
||||
tool_calls=[{"id": "call-2", "function": {"name": "terminal"}}],
|
||||
)
|
||||
|
||||
result = db.rewind_to_message(sid, target_id)
|
||||
|
||||
assert set(result) == {"rewound_count", "target_message", "new_head_id"}
|
||||
assert result["rewound_count"] == 2
|
||||
assert _session_counts(db, sid) == (2, 0, 1)
|
||||
|
||||
|
||||
def test_guarded_composite_rewind_rejects_append_without_inserting_scaffold(db):
|
||||
sid = "guarded-rewind-append"
|
||||
db.create_session(sid, source="cli")
|
||||
db.append_message(sid, "user", "first")
|
||||
db.append_message(sid, "assistant", "answer")
|
||||
target_id = db.append_message(sid, "user", _carrier())
|
||||
db.append_message(sid, "assistant", "failed")
|
||||
snapshot = db.get_messages_as_conversation(sid, include_row_ids=True)
|
||||
expected_active_ids = [int(message["_row_id"]) for message in snapshot]
|
||||
assert snapshot[-2]["_row_id"] == target_id
|
||||
|
||||
# Deterministic validation -> write race: a sibling writer commits after
|
||||
# the snapshot but before rewind_to_message begins its write transaction.
|
||||
sibling = SessionDB(db_path=db.db_path)
|
||||
sibling.append_message(sid, "assistant", "concurrent append")
|
||||
sibling.close()
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(RuntimeError, match="active transcript changed"):
|
||||
db.rewind_to_message(
|
||||
sid,
|
||||
target_id,
|
||||
preserve_compaction_handoff=True,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content=_carrier(),
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
|
||||
def test_guarded_rewind_rejects_selected_target_content_change(db):
|
||||
sid = "guarded-rewind-in-place"
|
||||
db.create_session(sid, source="cli")
|
||||
db.append_message(sid, "user", "first")
|
||||
db.append_message(sid, "assistant", "answer")
|
||||
target_id = db.append_message(sid, "user", "second")
|
||||
db.append_message(sid, "assistant", "failed")
|
||||
expected_active_ids = _active_ids(db, sid)
|
||||
|
||||
sibling = SessionDB(db_path=db.db_path)
|
||||
sibling._execute_write(
|
||||
lambda conn: conn.execute(
|
||||
"UPDATE messages SET content = ? WHERE id = ?",
|
||||
("changed second", target_id),
|
||||
)
|
||||
)
|
||||
sibling.close()
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(RuntimeError, match="rewind target changed"):
|
||||
db.rewind_to_message(
|
||||
sid,
|
||||
target_id,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content="second",
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
|
||||
|
||||
def test_guarded_rewind_ignores_reaction_metadata_change(db):
|
||||
sid = "guarded-rewind-reaction"
|
||||
db.create_session(sid, source="cli")
|
||||
target_id = db.append_message(sid, "user", "second")
|
||||
db.append_message(sid, "assistant", "failed")
|
||||
expected_active_ids = _active_ids(db, sid)
|
||||
|
||||
assert db.set_message_reaction(sid, target_id + 1, "👍", author="user")
|
||||
|
||||
result = db.rewind_to_message(
|
||||
sid,
|
||||
target_id,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content="second",
|
||||
)
|
||||
|
||||
assert result["rewound_count"] == 2
|
||||
assert db.get_messages_as_conversation(sid) == []
|
||||
|
||||
|
||||
def test_rewind_guard_rejects_foreign_live_compression_without_any_change(db):
|
||||
sid = "locked-rewind"
|
||||
db.create_session(sid, source="tui")
|
||||
target_id = db.append_message(sid, "user", _carrier())
|
||||
db.append_message(sid, "assistant", "failed")
|
||||
assert db.try_acquire_compression_lock(sid, "foreign-writer", ttl_seconds=60)
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(SessionCompressionInProgressError):
|
||||
db.rewind_to_message(
|
||||
sid, target_id, preserve_compaction_handoff=True
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
|
||||
|
||||
def test_rewind_guard_rejects_foreign_turn_lease_without_any_change(db):
|
||||
sid = "leased-rewind"
|
||||
db.create_session(sid, source="tui")
|
||||
target_id = db.append_message(sid, "user", _carrier())
|
||||
expected_active_ids = _active_ids(db, sid)
|
||||
holder = f"pid={os.getpid()}:turn=active"
|
||||
assert db.try_acquire_session_turn_lease(sid, holder, ttl_seconds=60)
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(SessionTurnLeaseLostError, match="active turn lease"):
|
||||
db.rewind_to_message(
|
||||
sid,
|
||||
target_id,
|
||||
preserve_compaction_handoff=True,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content="REAL ASK",
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
|
||||
db.release_session_turn_lease(sid, holder)
|
||||
result = db.rewind_to_message(
|
||||
sid,
|
||||
target_id,
|
||||
preserve_compaction_handoff=True,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content="REAL ASK",
|
||||
)
|
||||
assert result["rewound_count"] == 1
|
||||
|
||||
|
||||
def test_guarded_replace_rejects_foreign_turn_lease_without_any_change(db):
|
||||
sid = "leased-replace"
|
||||
db.create_session(sid, source="tui")
|
||||
db.append_message(sid, "user", "old ask")
|
||||
holder = f"pid={os.getpid()}:turn=active"
|
||||
assert db.try_acquire_session_turn_lease(sid, holder, ttl_seconds=60)
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(SessionTurnLeaseLostError, match="active turn lease"):
|
||||
db.replace_messages(
|
||||
sid,
|
||||
[{"role": "user", "content": "replacement"}],
|
||||
active_only=True,
|
||||
archive_dropped=True,
|
||||
reject_active_turn_lease=True,
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
|
||||
db.release_session_turn_lease(sid, holder)
|
||||
db.replace_messages(
|
||||
sid,
|
||||
[{"role": "user", "content": "replacement"}],
|
||||
active_only=True,
|
||||
archive_dropped=True,
|
||||
reject_active_turn_lease=True,
|
||||
)
|
||||
assert [m[2] for m in _row_state(db, sid) if m[3] == 1] == ["replacement"]
|
||||
|
||||
|
||||
def test_guarded_replace_rejects_foreign_live_compression_without_any_change(db):
|
||||
sid = "compression-locked-replace"
|
||||
db.create_session(sid, source="tui")
|
||||
db.append_message(sid, "user", "old ask")
|
||||
assert db.try_acquire_compression_lock(sid, "foreign-writer", ttl_seconds=60)
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(SessionCompressionInProgressError):
|
||||
db.replace_messages(
|
||||
sid,
|
||||
[{"role": "user", "content": "replacement"}],
|
||||
active_only=True,
|
||||
archive_dropped=True,
|
||||
reject_active_turn_lease=True,
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
|
||||
|
||||
def test_rewind_guard_rejects_compression_ended_parent_without_any_change(db):
|
||||
sid = "closed-rewind"
|
||||
db.create_session(sid, source="tui")
|
||||
target_id = db.append_message(sid, "user", _carrier())
|
||||
db.append_message(sid, "assistant", "failed")
|
||||
db.end_session(sid, "compression")
|
||||
before_rows = _row_state(db, sid)
|
||||
before_counts = _session_counts(db, sid)
|
||||
|
||||
with pytest.raises(CompressionSessionClosedError):
|
||||
db.rewind_to_message(
|
||||
sid, target_id, preserve_compaction_handoff=True
|
||||
)
|
||||
|
||||
assert _row_state(db, sid) == before_rows
|
||||
assert _session_counts(db, sid) == before_counts
|
||||
@@ -31,6 +31,55 @@ def _contents(db, session_id=SESSION_ID):
|
||||
|
||||
|
||||
class TestIdentityFlush:
|
||||
def test_summary_flush_hides_pure_handoff_but_not_composite_live_ask(self):
|
||||
from agent.context_compressor import (
|
||||
COMPRESSED_SUMMARY_METADATA_KEY,
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "t.db")
|
||||
try:
|
||||
agent = _make_agent(db)
|
||||
scaffold = (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold\n\n"
|
||||
f"{_SUMMARY_END_MARKER}"
|
||||
)
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": scaffold,
|
||||
COMPRESSED_SUMMARY_METADATA_KEY: True,
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": scaffold + "\n\nREAL ASK",
|
||||
COMPRESSED_SUMMARY_METADATA_KEY: True,
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": scaffold,
|
||||
COMPRESSED_SUMMARY_METADATA_KEY: True,
|
||||
},
|
||||
]
|
||||
|
||||
agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
rows = db._conn.execute(
|
||||
"SELECT content, display_kind FROM messages "
|
||||
"WHERE session_id = ? ORDER BY id",
|
||||
(SESSION_ID,),
|
||||
).fetchall()
|
||||
assert rows[0]["display_kind"] == "hidden"
|
||||
assert rows[1]["display_kind"] is None
|
||||
assert rows[1]["content"].endswith("REAL ASK")
|
||||
assert rows[2]["display_kind"] == "hidden"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def test_repair_shrunk_messages_below_history_length_still_persists_assistant(self):
|
||||
"""When repair shortens messages below conversation_history, don't slice empty."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
@@ -93,6 +93,17 @@ class TestDropThinkingOnlyAndMergeUsers:
|
||||
# Should return the original list untouched (identity) when no changes.
|
||||
assert out is msgs
|
||||
|
||||
def test_adjacent_users_merge_even_when_no_thinking_row_was_dropped(self):
|
||||
scaffold = {"role": "user", "content": "SUMMARY SCAFFOLD"}
|
||||
live_ask = {"role": "user", "content": "REAL ASK"}
|
||||
msgs = [scaffold, live_ask]
|
||||
|
||||
out = AIAgent._drop_thinking_only_and_merge_users(msgs)
|
||||
|
||||
assert out == [{"role": "user", "content": "SUMMARY SCAFFOLD\n\nREAL ASK"}]
|
||||
assert scaffold["content"] == "SUMMARY SCAFFOLD"
|
||||
assert live_ask["content"] == "REAL ASK"
|
||||
|
||||
|
||||
def test_preserves_alternation_after_drop(self):
|
||||
msgs = [
|
||||
|
||||
@@ -233,11 +233,25 @@ class TestRotationPathWatermark:
|
||||
the concurrent tail must follow the rotation instead of stranding in
|
||||
the closed parent."""
|
||||
|
||||
def test_tail_clones_into_the_child(self, db: SessionDB) -> None:
|
||||
@pytest.mark.parametrize(
|
||||
"tail_content",
|
||||
[
|
||||
"mid-rotation steer",
|
||||
[
|
||||
{"type": "text", "text": "inspect this"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,AA=="},
|
||||
},
|
||||
],
|
||||
],
|
||||
ids=["text", "multimodal"],
|
||||
)
|
||||
def test_tail_clones_into_the_child(self, db: SessionDB, tail_content) -> None:
|
||||
_seed(db)
|
||||
watermark = db.get_active_message_watermark("sess1")
|
||||
assert db.try_acquire_compression_lock("sess1", "rotator") is True
|
||||
db.append_message("sess1", role="user", content="mid-rotation steer")
|
||||
db.append_message("sess1", role="user", content=tail_content)
|
||||
# Ceiling captured AFTER the foreign append, BEFORE the rotation
|
||||
# path's own pre-publish flush (which this test has none of).
|
||||
ceiling = db.get_active_message_watermark("sess1")
|
||||
@@ -257,8 +271,20 @@ class TestRotationPathWatermark:
|
||||
assert [m["content"] for m in child] == [
|
||||
SUMMARY[0]["content"],
|
||||
SUMMARY[1]["content"],
|
||||
"mid-rotation steer",
|
||||
tail_content,
|
||||
]
|
||||
model_history, display_history = db.get_resume_conversations("child1")
|
||||
visible_steers = [
|
||||
message
|
||||
for message in display_history
|
||||
if message.get("content") == tail_content
|
||||
]
|
||||
assert len(visible_steers) == 1
|
||||
assert visible_steers[0]["_row_id"] == model_history[-1]["_row_id"]
|
||||
assert all(
|
||||
message.get("content") != tail_content
|
||||
for message in db.get_ancestor_display_prefix("child1")
|
||||
)
|
||||
info = db.get_session("child1")
|
||||
assert info["message_count"] == 3
|
||||
# Parent keeps its copy for lineage recovery; parent is closed.
|
||||
|
||||
@@ -5876,7 +5876,14 @@ def test_prompt_submit_rejects_negative_truncate_ordinal(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6013,7 +6020,14 @@ def test_prompt_submit_refuses_unconfirmed_nonempty_truncation(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6071,7 +6085,14 @@ def test_prompt_submit_truncates_by_message_id(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6121,7 +6142,14 @@ def test_prompt_submit_truncation_falls_back_to_sid_when_session_key_null(monkey
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6202,7 +6230,14 @@ def test_prompt_submit_refuses_ordinal_only_when_history_has_row_ids(monkeypatch
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6298,7 +6333,14 @@ def test_prompt_submit_truncates_by_row_id(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6343,7 +6385,14 @@ def test_prompt_submit_truncates_by_string_row_id(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6497,7 +6546,14 @@ def test_prompt_submit_refuses_empty_truncation_without_confirm(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -6586,7 +6642,14 @@ def test_prompt_submit_empty_truncation_allowed_with_confirm(monkeypatch):
|
||||
self._target()
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -10916,6 +10979,7 @@ def test_rollback_restore_truncates_from_real_user_turn_not_marker(monkeypatch):
|
||||
server._sessions["sid"] = _session(
|
||||
agent=types.SimpleNamespace(_checkpoint_mgr=_Mgr()),
|
||||
history=list(history),
|
||||
session_key="",
|
||||
)
|
||||
try:
|
||||
resp = server.handle_request(
|
||||
@@ -10976,6 +11040,7 @@ def test_rollback_restore_skips_legacy_compaction_handoff(monkeypatch):
|
||||
server._sessions["sid"] = _session(
|
||||
agent=types.SimpleNamespace(_checkpoint_mgr=_Mgr()),
|
||||
history=list(history),
|
||||
session_key="",
|
||||
)
|
||||
try:
|
||||
resp = server.handle_request(
|
||||
@@ -10998,6 +11063,85 @@ def test_rollback_restore_skips_legacy_compaction_handoff(monkeypatch):
|
||||
# ── session.steer ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_rollback_restore_preserves_composite_carrier_scaffold(monkeypatch, tmp_path):
|
||||
"""A checkpoint restore drops the live ask but keeps compacted context."""
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
class _Mgr:
|
||||
enabled = True
|
||||
|
||||
def list_checkpoints(self, cwd):
|
||||
return [{"hash": "abc123"}]
|
||||
|
||||
def restore(self, cwd, target, file_path=None):
|
||||
return {"success": True, "message": "restored"}
|
||||
|
||||
carrier = {
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}\n\nREAL ASK"
|
||||
),
|
||||
}
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
db.create_session("rollback-carrier", source="tui")
|
||||
db.append_message("rollback-carrier", "user", carrier["content"])
|
||||
db.append_message("rollback-carrier", "assistant", "answer")
|
||||
durable = db.get_messages_as_conversation("rollback-carrier")
|
||||
agent = types.SimpleNamespace(
|
||||
_checkpoint_mgr=_Mgr(),
|
||||
_session_messages=list(durable),
|
||||
_last_flushed_db_idx=len(durable),
|
||||
_db_flush_scan_prefix=list(durable),
|
||||
)
|
||||
server._sessions["sid"] = _session(
|
||||
agent=agent,
|
||||
history=list(durable),
|
||||
session_key="rollback-carrier",
|
||||
)
|
||||
monkeypatch.setattr(server, "_get_db", lambda: db)
|
||||
try:
|
||||
resp = server.handle_request(
|
||||
{
|
||||
"id": "1",
|
||||
"method": "rollback.restore",
|
||||
"params": {"session_id": "sid", "hash": "abc123"},
|
||||
}
|
||||
)
|
||||
|
||||
assert "result" in resp, resp
|
||||
assert resp["result"]["success"] is True
|
||||
assert resp["result"]["history_removed"] == 2
|
||||
remaining = server._sessions["sid"]["history"]
|
||||
assert len(remaining) == 1
|
||||
assert remaining[0]["display_kind"] == "hidden"
|
||||
assert SUMMARY_PREFIX in remaining[0]["content"]
|
||||
assert "REAL ASK" not in remaining[0]["content"]
|
||||
cold = db.get_messages_as_conversation(
|
||||
"rollback-carrier", include_row_ids=True
|
||||
)
|
||||
assert len(cold) == 1
|
||||
assert cold[0]["content"] == remaining[0]["content"]
|
||||
assert cold[0]["display_kind"] == "hidden"
|
||||
assert cold[0]["_row_id"] == remaining[0]["_row_id"]
|
||||
assert agent._session_messages == remaining
|
||||
assert agent._last_flushed_db_idx == 1
|
||||
assert agent._db_flush_scan_prefix == remaining
|
||||
inactive = db.get_messages_as_conversation(
|
||||
"rollback-carrier", include_inactive=True
|
||||
)
|
||||
assert any("REAL ASK" in str(message.get("content")) for message in inactive)
|
||||
assert any(message.get("content") == "answer" for message in inactive)
|
||||
finally:
|
||||
server._sessions.pop("sid", None)
|
||||
db.close()
|
||||
|
||||
|
||||
def test_session_steer_calls_agent_steer_when_agent_supports_it():
|
||||
"""The TUI RPC method must call agent.steer(text) and return a
|
||||
queued status without touching interrupt state.
|
||||
@@ -11417,6 +11561,7 @@ def test_session_undo_allowed_when_idle():
|
||||
"""Regression guard: when not running, /undo still works."""
|
||||
server._sessions["sid"] = _session(
|
||||
running=False,
|
||||
session_key="",
|
||||
history=[
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
@@ -11917,7 +12062,14 @@ def test_prompt_submit_can_truncate_before_user_ordinal(monkeypatch):
|
||||
def get_messages_as_conversation(self, *_args, **_kwargs):
|
||||
return []
|
||||
|
||||
def replace_messages(self, session_id, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
session_id,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
self.replaced.append((session_id, list(messages)))
|
||||
|
||||
stub_db = _StubDb()
|
||||
@@ -11977,7 +12129,14 @@ def test_prompt_submit_refuses_turn_when_truncate_persist_fails(monkeypatch):
|
||||
def get_messages_as_conversation(self, *_args, **_kwargs):
|
||||
return []
|
||||
|
||||
def replace_messages(self, session_id, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
session_id,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
raise OSError("disk full")
|
||||
|
||||
monkeypatch.setattr(server, "_get_db", lambda: _FailDb())
|
||||
@@ -12072,7 +12231,14 @@ def test_prompt_submit_truncate_ordinal_skips_display_kind_rows(monkeypatch):
|
||||
def get_messages_as_conversation(self, *_args, **_kwargs):
|
||||
return []
|
||||
|
||||
def replace_messages(self, session_id, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
session_id,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
self.replaced.append((session_id, list(messages)))
|
||||
|
||||
stub_db = _StubDb()
|
||||
@@ -12175,7 +12341,15 @@ def test_prompt_submit_truncate_translates_display_prefix_ordinal(monkeypatch):
|
||||
# may take the ordinal-only path past the durability gate.
|
||||
return []
|
||||
|
||||
def replace_messages(self, session_id, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
session_id,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
assert reject_active_turn_lease is True
|
||||
self.replaced.append((session_id, list(messages)))
|
||||
|
||||
stub_db = _StubDb()
|
||||
@@ -12260,6 +12434,12 @@ def test_prompt_submit_row_id_accepts_full_lineage_ordinal(monkeypatch):
|
||||
reconcile cross-check must treat `tip_ordinal + prefix_user_count` as
|
||||
agreement, not #82756 drift — and the cut stays aimed by the row id.
|
||||
"""
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
|
||||
tip_history = [
|
||||
{"_row_id": 501, "role": "user", "content": "post-compress A"},
|
||||
{"_row_id": 502, "role": "assistant", "content": "reply A"},
|
||||
@@ -12269,6 +12449,16 @@ def test_prompt_submit_row_id_accepts_full_lineage_ordinal(monkeypatch):
|
||||
display_prefix = [
|
||||
{"role": "user", "content": "pre-compress 1"},
|
||||
{"role": "assistant", "content": "pre reply 1"},
|
||||
{
|
||||
# Legacy pure handoffs did not always carry display_kind=hidden.
|
||||
# They are physically user rows but not visible/user-originated
|
||||
# turns, so they must not shift the Desktop lineage ordinal.
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}"
|
||||
),
|
||||
},
|
||||
{"role": "user", "content": "pre-compress 2"},
|
||||
{"role": "assistant", "content": "pre reply 2"},
|
||||
]
|
||||
@@ -12276,7 +12466,15 @@ def test_prompt_submit_row_id_accepts_full_lineage_ordinal(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
assert reject_active_turn_lease is True
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
sess = _session(history=list(tip_history), display_history_prefix=display_prefix)
|
||||
@@ -19642,7 +19840,14 @@ def test_personality_marker_does_not_shift_truncate_ordinal(monkeypatch):
|
||||
def get_messages_as_conversation(self, *_args, **_kwargs):
|
||||
return []
|
||||
|
||||
def replace_messages(self, session_id, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
session_id,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
self.replaced.append((session_id, list(messages)))
|
||||
|
||||
session = _session(
|
||||
@@ -19753,10 +19958,16 @@ def test_prompt_submit_truncation_archives_instead_of_deleting(monkeypatch):
|
||||
return []
|
||||
|
||||
def replace_messages(
|
||||
self, session_id, messages, active_only=False, archive_dropped=False
|
||||
self,
|
||||
session_id,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
captured["active_only"] = active_only
|
||||
captured["archive_dropped"] = archive_dropped
|
||||
captured["reject_active_turn_lease"] = reject_active_turn_lease
|
||||
|
||||
server._sessions["archive-trunc-sid"] = _session(
|
||||
agent=_Agent(),
|
||||
@@ -19795,6 +20006,7 @@ def test_prompt_submit_truncation_archives_instead_of_deleting(monkeypatch):
|
||||
)
|
||||
# #80216: still must not touch rows archived by an earlier compaction.
|
||||
assert captured.get("active_only") is True
|
||||
assert captured.get("reject_active_turn_lease") is True
|
||||
finally:
|
||||
server._sessions.pop("archive-trunc-sid", None)
|
||||
|
||||
@@ -19816,7 +20028,14 @@ def test_prompt_submit_unmatched_row_id_refuses_even_with_ordinal(monkeypatch):
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
history = [
|
||||
@@ -19860,7 +20079,14 @@ def test_prompt_submit_unmatched_message_id_refuses_even_with_ordinal(monkeypatc
|
||||
replaced = []
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
# Production-shaped history: no renderer "id" keys on user dicts.
|
||||
@@ -19918,7 +20144,14 @@ def test_prompt_submit_row_id_resolves_via_db_when_memory_lacks_stamps(monkeypat
|
||||
]
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
def get_messages_as_conversation(self, key, repair_alternation=False, include_row_ids=False):
|
||||
@@ -20103,6 +20336,7 @@ def test_prompt_submit_row_id_misaligned_memory_refuses_content_swap(
|
||||
db._insert_message_rows(db._conn, session_key, msgs)
|
||||
db._conn.commit()
|
||||
rid_b = msgs[2]["_row_id"]
|
||||
original_row_ids = [message["_row_id"] for message in msgs]
|
||||
|
||||
# Same length + same role pattern, but content positions swapped: a
|
||||
# positional stamp would mark live "B" with durable A's row id and the
|
||||
@@ -20134,6 +20368,7 @@ def test_prompt_submit_row_id_misaligned_memory_refuses_content_swap(
|
||||
"session_id": sid,
|
||||
"text": "rewind B",
|
||||
"truncate_before_row_id": rid_b,
|
||||
"rebind_survivor_row_ids": [*original_row_ids, 999_999],
|
||||
"confirm_truncate": True,
|
||||
},
|
||||
}
|
||||
@@ -20172,6 +20407,7 @@ def test_prompt_submit_row_id_misaligned_memory_role_shift_targets_real_turn(
|
||||
db._insert_message_rows(db._conn, session_key, msgs)
|
||||
db._conn.commit()
|
||||
rid_b = msgs[2]["_row_id"]
|
||||
original_row_ids = [message["_row_id"] for message in msgs]
|
||||
|
||||
live_history = [
|
||||
{"role": "user", "content": "A"},
|
||||
@@ -20195,6 +20431,7 @@ def test_prompt_submit_row_id_misaligned_memory_role_shift_targets_real_turn(
|
||||
"session_id": sid,
|
||||
"text": "rewind B",
|
||||
"truncate_before_row_id": rid_b,
|
||||
"rebind_survivor_row_ids": [*original_row_ids, 999_999],
|
||||
"confirm_truncate": True,
|
||||
},
|
||||
}
|
||||
@@ -20211,6 +20448,13 @@ def test_prompt_submit_row_id_misaligned_memory_role_shift_targets_real_turn(
|
||||
("assistant", "ra"),
|
||||
("assistant", "rb"),
|
||||
]
|
||||
# The live list was too misaligned to bind old survivors to their new
|
||||
# physical rows safely. All requested IDs known to the pre-write active
|
||||
# transcript are therefore cleared; an unrelated archived/ancestor ID
|
||||
# remains absent from the bounded map and keeps its identity.
|
||||
assert resp["result"]["survivor_row_id_map"] == {
|
||||
str(row_id): None for row_id in original_row_ids
|
||||
}
|
||||
finally:
|
||||
server._sessions.pop(sid, None)
|
||||
|
||||
@@ -20244,7 +20488,14 @@ def test_prompt_submit_row_id_db_fallback_ordinal_mapping_verifies_content(
|
||||
]
|
||||
|
||||
class _FakeDB:
|
||||
def replace_messages(self, key, messages, active_only=False, archive_dropped=False):
|
||||
def replace_messages(
|
||||
self,
|
||||
key,
|
||||
messages,
|
||||
active_only=False,
|
||||
archive_dropped=False,
|
||||
reject_active_turn_lease=False,
|
||||
):
|
||||
replaced.append((key, list(messages)))
|
||||
|
||||
def get_messages_as_conversation(self, key, repair_alternation=False, include_row_ids=False):
|
||||
@@ -20284,8 +20535,9 @@ def test_prompt_submit_row_id_db_fallback_ordinal_mapping_verifies_content(
|
||||
server._sessions.pop(sid, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("turn_isolation", [False, True])
|
||||
def test_prompt_submit_consecutive_rewinds_with_returned_survivor_row_ids(
|
||||
monkeypatch, tmp_path
|
||||
monkeypatch, tmp_path, turn_isolation
|
||||
):
|
||||
"""#83202 review (consecutive-rewind staleness): replace_messages re-inserts
|
||||
the surviving prefix as NEW rows, so the pre-rewind client row ids die on
|
||||
@@ -20314,6 +20566,16 @@ def test_prompt_submit_consecutive_rewinds_with_returned_survivor_row_ids(
|
||||
sid = "real-db-consec-rewind-sid"
|
||||
server._sessions[sid] = sess
|
||||
monkeypatch.setattr(server, "_get_db", lambda: db)
|
||||
monkeypatch.setattr(
|
||||
server,
|
||||
"_load_cfg",
|
||||
lambda: {"dashboard": {"turn_isolation": turn_isolation}},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
server,
|
||||
"_submit_prompt_to_compute_host",
|
||||
lambda *_args, **_kwargs: server._ok("host", {"status": "streaming"}),
|
||||
)
|
||||
monkeypatch.setattr(server, "_start_agent_build", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_start_inflight_turn", lambda *a, **k: None)
|
||||
|
||||
@@ -20329,17 +20591,30 @@ def test_prompt_submit_consecutive_rewinds_with_returned_survivor_row_ids(
|
||||
"text": "rewound third",
|
||||
"truncate_before_row_id": original_row_ids[4],
|
||||
"truncate_before_user_ordinal": 2,
|
||||
"rebind_survivor_row_ids": [*original_row_ids, 999_999],
|
||||
"confirm_truncate": True,
|
||||
},
|
||||
}
|
||||
)
|
||||
assert resp1.get("error") is None, resp1
|
||||
survivors = resp1["result"].get("survivor_user_row_ids")
|
||||
# Fresh ids for the two surviving user turns, in visible-user order.
|
||||
assert isinstance(survivors, list) and len(survivors) == 2
|
||||
assert all(isinstance(r, int) for r in survivors)
|
||||
assert "survivor_user_row_ids" not in resp1["result"]
|
||||
row_id_map = resp1["result"].get("survivor_row_id_map")
|
||||
assert isinstance(row_id_map, dict)
|
||||
survivors = [
|
||||
row_id_map[str(original_row_ids[0])],
|
||||
row_id_map[str(original_row_ids[2])],
|
||||
]
|
||||
# They must be NEW rows — the old ids are archived (active=0) now.
|
||||
assert set(survivors).isdisjoint(set(original_row_ids))
|
||||
assert row_id_map == {
|
||||
str(original_row_ids[0]): survivors[0],
|
||||
str(original_row_ids[1]): sess["history"][1]["_row_id"],
|
||||
str(original_row_ids[2]): survivors[1],
|
||||
str(original_row_ids[3]): sess["history"][3]["_row_id"],
|
||||
str(original_row_ids[4]): None,
|
||||
str(original_row_ids[5]): None,
|
||||
}
|
||||
assert "999999" not in row_id_map
|
||||
sess["running"] = False
|
||||
|
||||
# Rewind 2a: the STALE pre-rewind id for "second" must fail closed.
|
||||
@@ -20386,6 +20661,66 @@ def test_prompt_submit_consecutive_rewinds_with_returned_survivor_row_ids(
|
||||
server._sessions.pop(sid, None)
|
||||
|
||||
|
||||
def test_prompt_submit_rebind_map_clears_active_row_hidden_by_sequence_repair(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
"""The bounded map classifies physical active IDs before user;user repair."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "rowid-repaired-wedge.db")
|
||||
session_key = "real-db-rowid-repaired-wedge"
|
||||
db.create_session(session_key, "cli")
|
||||
physical = [
|
||||
{"role": "user", "content": "first fragment"},
|
||||
{"role": "user", "content": "second fragment"},
|
||||
{"role": "assistant", "content": "combined reply"},
|
||||
{"role": "user", "content": "target"},
|
||||
{"role": "assistant", "content": "target reply"},
|
||||
]
|
||||
with db._lock:
|
||||
db._insert_message_rows(db._conn, session_key, physical)
|
||||
db._conn.commit()
|
||||
physical_ids = [message["_row_id"] for message in physical]
|
||||
repaired = db.get_messages_as_conversation(
|
||||
session_key, repair_alternation=True, include_row_ids=True
|
||||
)
|
||||
# Provider repair merges the wedge and necessarily drops the second
|
||||
# physical user's row identity from the replay view.
|
||||
assert physical_ids[1] not in {
|
||||
server._message_row_id(message) for message in repaired
|
||||
}
|
||||
|
||||
sess = _session(
|
||||
history=[dict(message) for message in repaired], session_key=session_key
|
||||
)
|
||||
sid = "rowid-repaired-wedge-sid"
|
||||
server._sessions[sid] = sess
|
||||
monkeypatch.setattr(server, "_get_db", lambda: db)
|
||||
monkeypatch.setattr(server, "_start_agent_build", lambda *a, **k: None)
|
||||
monkeypatch.setattr(server, "_start_inflight_turn", lambda *a, **k: None)
|
||||
|
||||
try:
|
||||
response = server.handle_request(
|
||||
{
|
||||
"id": "1",
|
||||
"method": "prompt.submit",
|
||||
"params": {
|
||||
"session_id": sid,
|
||||
"text": "retry target",
|
||||
"truncate_before_row_id": physical_ids[3],
|
||||
"rebind_survivor_row_ids": [*physical_ids, 999_999],
|
||||
"confirm_truncate": True,
|
||||
},
|
||||
}
|
||||
)
|
||||
assert response.get("error") is None, response
|
||||
row_id_map = response["result"]["survivor_row_id_map"]
|
||||
assert row_id_map[str(physical_ids[1])] is None
|
||||
assert "999999" not in row_id_map
|
||||
finally:
|
||||
server._sessions.pop(sid, None)
|
||||
|
||||
|
||||
def test_prompt_submit_unconfirmed_truncation_refuses_before_target_resolution(
|
||||
monkeypatch,
|
||||
):
|
||||
|
||||
@@ -0,0 +1,512 @@
|
||||
"""Regression contracts for TUI rewind of live compaction carriers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_compressor import (
|
||||
HISTORICAL_TASK_HEADING,
|
||||
SUMMARY_PREFIX,
|
||||
_SUMMARY_END_MARKER,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
from tui_gateway import server
|
||||
|
||||
|
||||
def _composite_carrier() -> dict:
|
||||
return {
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n"
|
||||
f"{_SUMMARY_END_MARKER}\n\nREAL ASK"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def carrier_session(tmp_path):
|
||||
old_db = server._db
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
installed_ids: list[str] = []
|
||||
|
||||
def install(history: list[dict]):
|
||||
sid = f"carrier-sid-{len(installed_ids)}"
|
||||
session_key = f"carrier-session-{len(installed_ids)}"
|
||||
installed_ids.append(sid)
|
||||
db.create_session(session_key, source="tui")
|
||||
for message in history:
|
||||
db.append_message(
|
||||
session_key,
|
||||
message["role"],
|
||||
message.get("content"),
|
||||
)
|
||||
durable = db.get_messages_as_conversation(session_key)
|
||||
agent = SimpleNamespace(
|
||||
_session_messages=list(durable),
|
||||
_last_flushed_db_idx=len(durable),
|
||||
_db_flush_scan_prefix=list(durable),
|
||||
)
|
||||
session = {
|
||||
"agent": agent,
|
||||
"attached_images": [],
|
||||
"history": list(durable),
|
||||
"history_lock": threading.Lock(),
|
||||
"history_version": 0,
|
||||
"running": False,
|
||||
"session_key": session_key,
|
||||
}
|
||||
server._sessions[sid] = session
|
||||
return sid, session_key, session
|
||||
|
||||
server._db = db
|
||||
yield db, install
|
||||
for sid in installed_ids:
|
||||
server._sessions.pop(sid, None)
|
||||
server._db = old_db
|
||||
db.close()
|
||||
|
||||
|
||||
def _dispatch(sid: str, name: str) -> dict:
|
||||
return server._methods["command.dispatch"](
|
||||
"request-id",
|
||||
{"session_id": sid, "name": name, "arg": ""},
|
||||
)
|
||||
|
||||
|
||||
def _session_undo(sid: str) -> dict:
|
||||
return server._methods["session.undo"](
|
||||
"request-id",
|
||||
{"session_id": sid},
|
||||
)
|
||||
|
||||
|
||||
def _assert_scaffold_preserved(
|
||||
db: SessionDB,
|
||||
session_key: str,
|
||||
session: dict,
|
||||
*,
|
||||
prefix_len: int = 0,
|
||||
) -> None:
|
||||
active = db.get_messages_as_conversation(session_key, include_row_ids=True)
|
||||
assert len(active) == prefix_len + 1
|
||||
scaffold = active[prefix_len]
|
||||
assert scaffold["role"] == "user"
|
||||
assert scaffold["display_kind"] == "hidden"
|
||||
assert SUMMARY_PREFIX in scaffold["content"]
|
||||
assert "REAL ASK" not in scaffold["content"]
|
||||
assert session["history"][prefix_len]["content"] == scaffold["content"]
|
||||
assert session["history"][prefix_len]["display_kind"] == "hidden"
|
||||
|
||||
|
||||
def test_retry_selects_the_live_ask_inside_a_force_user_leading_carrier(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[_composite_carrier(), {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["result"] == {"type": "send", "message": "REAL ASK"}
|
||||
_assert_scaffold_preserved(db, session_key, session)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", ["retry", "undo"])
|
||||
def test_rewind_matches_cold_sanitized_carrier_to_unchanged_warm_ask(
|
||||
carrier_session, command
|
||||
):
|
||||
db, install = carrier_session
|
||||
carrier = _composite_carrier()
|
||||
carrier["content"] = carrier["content"].replace(
|
||||
"REAL ASK",
|
||||
" REAL ASK\n\n<memory-context>\nprivate\n</memory-context> ",
|
||||
)
|
||||
sid, session_key, session = install(
|
||||
[carrier, {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
db._conn.execute(
|
||||
"UPDATE messages SET api_content = ? "
|
||||
"WHERE session_id = ? AND role = 'user'",
|
||||
(carrier["content"], session_key),
|
||||
)
|
||||
db._conn.commit()
|
||||
# A live agent still has the raw wire form while a cold DB projection has
|
||||
# already applied the role-aware sanitize_context(...).strip() rule and
|
||||
# retained the raw provider wire in api_content.
|
||||
session["history"][0] = carrier.copy()
|
||||
session["agent"]._session_messages = list(session["history"])
|
||||
|
||||
response = _dispatch(sid, command)
|
||||
|
||||
assert response["result"]["message"] == "REAL ASK"
|
||||
_assert_scaffold_preserved(db, session_key, session)
|
||||
|
||||
|
||||
def test_retry_fails_closed_when_transcript_changes_after_snapshot(
|
||||
carrier_session, monkeypatch
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[_composite_carrier(), {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
sibling = SessionDB(db_path=db.db_path)
|
||||
original_rewind = db.rewind_to_message
|
||||
|
||||
def _append_then_rewind(*args, **kwargs):
|
||||
sibling.append_message(session_key, "assistant", "concurrent tail")
|
||||
return original_rewind(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(db, "rewind_to_message", _append_then_rewind)
|
||||
before_history = [dict(message) for message in session["history"]]
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["error"]["code"] == 5008
|
||||
assert "active transcript changed" in response["error"]["message"]
|
||||
assert session["history"] == before_history
|
||||
rows = db._conn.execute(
|
||||
"SELECT content, active, display_kind FROM messages "
|
||||
"WHERE session_id = ? ORDER BY id",
|
||||
(session_key,),
|
||||
).fetchall()
|
||||
assert [tuple(row) for row in rows] == [
|
||||
(_composite_carrier()["content"], 1, None),
|
||||
("failed", 1, None),
|
||||
("concurrent tail", 1, None),
|
||||
]
|
||||
sibling.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", ["retry", "undo"])
|
||||
def test_rewind_allows_database_only_reaction_metadata_change(
|
||||
carrier_session, command
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[
|
||||
{"role": "user", "content": "OLDER ASK"},
|
||||
{"role": "assistant", "content": "older answer"},
|
||||
_composite_carrier(),
|
||||
{"role": "assistant", "content": "failed"},
|
||||
]
|
||||
)
|
||||
older_answer = next(
|
||||
row
|
||||
for row in db.get_messages(session_key)
|
||||
if row["role"] == "assistant" and row["content"] == "older answer"
|
||||
)
|
||||
assert db.set_message_reaction(
|
||||
session_key, older_answer["id"], "👍", author="user"
|
||||
)
|
||||
|
||||
response = _dispatch(sid, command)
|
||||
|
||||
assert response["result"]["message"] == "REAL ASK"
|
||||
_assert_scaffold_preserved(db, session_key, session, prefix_len=2)
|
||||
|
||||
|
||||
def test_retry_ignores_buried_ephemeral_scaffolding_missing_from_db(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[
|
||||
{"role": "user", "content": "OLDER ASK"},
|
||||
{"role": "assistant", "content": "older answer"},
|
||||
_composite_carrier(),
|
||||
{"role": "assistant", "content": "failed"},
|
||||
]
|
||||
)
|
||||
session["history"].insert(
|
||||
2,
|
||||
{
|
||||
"role": "user",
|
||||
"content": "internal recovery nudge",
|
||||
"_dropped_toolcall_nudge": True,
|
||||
},
|
||||
)
|
||||
session["agent"]._session_messages = list(session["history"])
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["result"] == {"type": "send", "message": "REAL ASK"}
|
||||
_assert_scaffold_preserved(db, session_key, session, prefix_len=2)
|
||||
|
||||
|
||||
def test_retry_drops_buried_ephemeral_scaffolding_from_the_warm_prefix(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[
|
||||
{"role": "user", "content": "OLDER ASK"},
|
||||
{"role": "assistant", "content": "candidate answer"},
|
||||
{"role": "assistant", "content": "verified answer"},
|
||||
_composite_carrier(),
|
||||
{"role": "assistant", "content": "failed"},
|
||||
]
|
||||
)
|
||||
session["history"].insert(
|
||||
2,
|
||||
{
|
||||
"role": "user",
|
||||
"content": "[System: verify before stopping]",
|
||||
"_verification_stop_synthetic": True,
|
||||
},
|
||||
)
|
||||
session["agent"]._session_messages = list(session["history"])
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["result"] == {"type": "send", "message": "REAL ASK"}
|
||||
assert [message.get("content") for message in session["history"][:3]] == [
|
||||
"OLDER ASK",
|
||||
"candidate answer",
|
||||
"verified answer",
|
||||
]
|
||||
active = db.get_messages_as_conversation(session_key, include_row_ids=True)
|
||||
# Alternation repair is a model/memory projection; the two durable source
|
||||
# rows remain independently recoverable ahead of the inserted scaffold.
|
||||
assert [message.get("content") for message in active[:3]] == [
|
||||
"OLDER ASK",
|
||||
"candidate answer",
|
||||
"verified answer",
|
||||
]
|
||||
assert active[3]["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in active[3]["content"]
|
||||
assert session["history"][3]["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in session["history"][3]["content"]
|
||||
|
||||
|
||||
def test_retry_preserves_older_warm_media_while_targeting_plain_ask(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[
|
||||
{"role": "user", "content": "look\n[screenshot]"},
|
||||
{"role": "assistant", "content": "seen"},
|
||||
_composite_carrier(),
|
||||
{"role": "assistant", "content": "failed"},
|
||||
]
|
||||
)
|
||||
session["history"][0]["content"] = [
|
||||
{"type": "text", "text": "look"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,x"}},
|
||||
]
|
||||
session["agent"]._session_messages = list(session["history"])
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["result"] == {"type": "send", "message": "REAL ASK"}
|
||||
assert isinstance(session["history"][0]["content"], list)
|
||||
_assert_scaffold_preserved(db, session_key, session, prefix_len=2)
|
||||
|
||||
|
||||
def test_undo_targets_the_composite_carrier_not_an_older_user_turn(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[
|
||||
{"role": "user", "content": "OLDER ASK"},
|
||||
{"role": "assistant", "content": "older answer"},
|
||||
_composite_carrier(),
|
||||
{"role": "assistant", "content": "failed"},
|
||||
]
|
||||
)
|
||||
|
||||
response = _dispatch(sid, "undo")
|
||||
|
||||
assert response["result"]["type"] == "prefill"
|
||||
assert response["result"]["message"] == "REAL ASK"
|
||||
active = db.get_messages_as_conversation(session_key, include_row_ids=True)
|
||||
assert [message.get("content") for message in active[:2]] == [
|
||||
"OLDER ASK",
|
||||
"older answer",
|
||||
]
|
||||
_assert_scaffold_preserved(db, session_key, session, prefix_len=2)
|
||||
|
||||
|
||||
def test_undo_rewinds_media_placeholder_without_treating_it_as_retry(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
carrier = _composite_carrier()
|
||||
carrier["content"] = carrier["content"].replace(
|
||||
"REAL ASK", "look\n[screenshot]"
|
||||
)
|
||||
sid, session_key, session = install(
|
||||
[carrier, {"role": "assistant", "content": "seen"}]
|
||||
)
|
||||
|
||||
response = _dispatch(sid, "undo")
|
||||
|
||||
assert response["result"]["type"] == "prefill"
|
||||
assert response["result"]["message"] == "look\n[screenshot]"
|
||||
active = db.get_messages_as_conversation(session_key)
|
||||
assert len(active) == 1
|
||||
assert active[0]["display_kind"] == "hidden"
|
||||
assert "look\n[screenshot]" not in active[0]["content"]
|
||||
assert len(session["history"]) == 1
|
||||
assert session["history"][0]["content"] == active[0]["content"]
|
||||
assert session["history"][0]["display_kind"] == "hidden"
|
||||
|
||||
|
||||
def test_session_undo_preserves_the_composite_carriers_scaffold(carrier_session):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[_composite_carrier(), {"role": "assistant", "content": "answer"}]
|
||||
)
|
||||
|
||||
response = _session_undo(sid)
|
||||
|
||||
assert response["result"]["removed"] == 2
|
||||
_assert_scaffold_preserved(db, session_key, session)
|
||||
|
||||
|
||||
def test_history_projection_unwraps_composite_and_hides_sole_handoff():
|
||||
composite = {**_composite_carrier(), "_row_id": 7}
|
||||
sole_handoff = {
|
||||
**composite,
|
||||
"content": composite["content"].split("\n\nREAL ASK", 1)[0],
|
||||
}
|
||||
|
||||
assert server._history_to_messages(
|
||||
[composite, {"role": "user", "content": "newer ask", "_row_id": 9}]
|
||||
) == [
|
||||
{"role": "user", "text": "REAL ASK", "row_id": 7},
|
||||
{"role": "user", "text": "newer ask", "row_id": 9},
|
||||
]
|
||||
assert server._history_to_messages([sole_handoff]) == []
|
||||
|
||||
|
||||
def test_retry_preserves_literal_media_like_text(carrier_session):
|
||||
db, install = carrier_session
|
||||
carrier = _composite_carrier()
|
||||
carrier["content"] = carrier["content"].replace(
|
||||
"REAL ASK", "inspect [image|ybres:RID]"
|
||||
)
|
||||
sid, session_key, session = install(
|
||||
[carrier, {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["result"] == {
|
||||
"type": "send",
|
||||
"message": "inspect [image|ybres:RID]",
|
||||
}
|
||||
_assert_scaffold_preserved(db, session_key, session)
|
||||
|
||||
|
||||
def test_retry_rejects_durable_media_before_rewind_when_warm_view_is_text(
|
||||
carrier_session,
|
||||
):
|
||||
db, install = carrier_session
|
||||
carrier = _composite_carrier()
|
||||
handoff = carrier["content"].rsplit("\n\nREAL ASK", 1)[0]
|
||||
durable_carrier = carrier.copy()
|
||||
durable_carrier["content"] = [
|
||||
{"type": "text", "text": handoff},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,x"}},
|
||||
]
|
||||
sid, session_key, session = install(
|
||||
[durable_carrier, {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
# The warm projection can be a degraded text-only view that compares equal
|
||||
# to the durable media payload. Durable retryability must still be checked
|
||||
# before the physical carrier and tail are archived.
|
||||
warm_carrier = carrier.copy()
|
||||
warm_carrier["content"] = handoff + "\n\n[screenshot]"
|
||||
session["history"][0] = warm_carrier
|
||||
session["agent"]._session_messages = list(session["history"])
|
||||
before_history = [message.copy() for message in session["history"]]
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["error"]["code"] == 4018
|
||||
assert session["history"] == before_history
|
||||
assert len(db.get_messages_as_conversation(session_key)) == 2
|
||||
|
||||
|
||||
def test_retry_rejects_pending_attachments_before_mutating_history(carrier_session):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[_composite_carrier(), {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
session["attached_images"] = ["/tmp/pending.png"]
|
||||
before_memory = list(session["history"])
|
||||
|
||||
response = _dispatch(sid, "retry")
|
||||
|
||||
assert response["error"]["code"] == 4018
|
||||
assert session["history"] == before_memory
|
||||
assert len(db.get_messages_as_conversation(session_key)) == 2
|
||||
|
||||
def test_prompt_row_id_rewind_preserves_scaffold_before_regeneration(
|
||||
carrier_session, monkeypatch
|
||||
):
|
||||
db, install = carrier_session
|
||||
sid, session_key, session = install(
|
||||
[_composite_carrier(), {"role": "assistant", "content": "failed"}]
|
||||
)
|
||||
target_row_id = db.get_messages_as_conversation(
|
||||
session_key, include_row_ids=True
|
||||
)[0]["_row_id"]
|
||||
seen = {}
|
||||
|
||||
class _Agent:
|
||||
_session_messages = list(session["history"])
|
||||
_last_flushed_db_idx = len(_session_messages)
|
||||
_db_flush_scan_prefix = list(_session_messages)
|
||||
|
||||
def run_conversation(
|
||||
self, prompt, conversation_history=None, stream_callback=None, **_kwargs
|
||||
):
|
||||
seen["prompt"] = prompt
|
||||
seen["history"] = list(conversation_history or [])
|
||||
return {
|
||||
"final_response": "regenerated",
|
||||
"messages": [
|
||||
*(conversation_history or []),
|
||||
{"role": "user", "content": prompt},
|
||||
{"role": "assistant", "content": "regenerated"},
|
||||
],
|
||||
}
|
||||
|
||||
class _ImmediateThread:
|
||||
def __init__(self, target=None, daemon=None):
|
||||
self._target = target
|
||||
|
||||
def start(self):
|
||||
self._target()
|
||||
|
||||
session["agent"] = _Agent()
|
||||
monkeypatch.setattr(server.threading, "Thread", _ImmediateThread)
|
||||
monkeypatch.setattr(server, "_get_usage", lambda _agent: {})
|
||||
monkeypatch.setattr(server, "render_message", lambda *_args: "")
|
||||
monkeypatch.setattr(server, "_emit", lambda *_args: None)
|
||||
|
||||
response = server._methods["prompt.submit"](
|
||||
"request-id",
|
||||
{
|
||||
"session_id": sid,
|
||||
"text": "EDITED ASK",
|
||||
"truncate_before_row_id": target_row_id,
|
||||
"truncate_before_user_ordinal": 0,
|
||||
"confirm_truncate": True,
|
||||
},
|
||||
)
|
||||
|
||||
assert response["result"]["status"] == "streaming"
|
||||
assert seen["prompt"] == "EDITED ASK"
|
||||
assert len(seen["history"]) == 1
|
||||
assert seen["history"][0]["display_kind"] == "hidden"
|
||||
assert "REAL ASK" not in seen["history"][0]["content"]
|
||||
active = db.get_messages_as_conversation(session_key, include_row_ids=True)
|
||||
assert len(active) == 1
|
||||
assert active[0]["display_kind"] == "hidden"
|
||||
+121
-22
@@ -14,11 +14,13 @@ _profile_scoped = _registry.profile_scoped
|
||||
|
||||
|
||||
def _history_user_indices(history: list) -> list:
|
||||
"""Indices of model-visible user turns (excludes display_kind timeline markers)."""
|
||||
"""Indices of canonical live-user turns, including composite carriers."""
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
|
||||
return [
|
||||
i
|
||||
for i, m in enumerate(history)
|
||||
if m.get("role") == "user" and not m.get("display_kind")
|
||||
if user_originated_turn_view(m) is not None
|
||||
]
|
||||
|
||||
|
||||
@@ -50,17 +52,28 @@ def _mem_db_pair_agrees(mem, db_msg) -> bool:
|
||||
return False
|
||||
if mem.get("role") != db_msg.get("role"):
|
||||
return False
|
||||
if mem.get("role") == "user":
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
from agent.memory_manager import sanitize_context
|
||||
|
||||
mem_view = user_originated_turn_view(mem)
|
||||
db_view = user_originated_turn_view(db_msg)
|
||||
if (mem_view is None) != (db_view is None):
|
||||
return False
|
||||
if mem_view is None:
|
||||
return bool(mem.get("display_kind")) == bool(
|
||||
db_msg.get("display_kind")
|
||||
)
|
||||
mem_content = mem_view.get("content")
|
||||
db_content = db_view.get("content")
|
||||
if isinstance(mem_content, str) and isinstance(db_content, str):
|
||||
if sanitize_context(mem_content).strip() != sanitize_context(
|
||||
db_content
|
||||
).strip():
|
||||
return False
|
||||
return True
|
||||
if bool(mem.get("display_kind")) != bool(db_msg.get("display_kind")):
|
||||
return False
|
||||
if mem.get("role") == "user" and not mem.get("display_kind"):
|
||||
mem_content = mem.get("content")
|
||||
db_content = db_msg.get("content")
|
||||
if (
|
||||
isinstance(mem_content, str)
|
||||
and isinstance(db_content, str)
|
||||
and mem_content.strip() != db_content.strip()
|
||||
):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
@@ -72,7 +85,11 @@ def _find_user_turn_by_row_id(history: list, target_row_id: int):
|
||||
return None
|
||||
|
||||
|
||||
def _load_durable_truncation_history(session: dict, fallback_sid: str = ""):
|
||||
def _load_durable_truncation_history(
|
||||
session: dict,
|
||||
fallback_sid: str = "",
|
||||
repair_alternation: bool = True,
|
||||
):
|
||||
"""Load the durable live-replay transcript, or None when it cannot be proven safe."""
|
||||
session_key = str(session.get("session_key") or fallback_sid or "")
|
||||
if not session_key:
|
||||
@@ -83,7 +100,9 @@ def _load_durable_truncation_history(session: dict, fallback_sid: str = ""):
|
||||
if not callable(get_conv):
|
||||
return None
|
||||
history = get_conv(
|
||||
session_key, repair_alternation=True, include_row_ids=True
|
||||
session_key,
|
||||
repair_alternation=repair_alternation,
|
||||
include_row_ids=True,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug(
|
||||
@@ -368,6 +387,17 @@ def _(rid, params: dict) -> dict:
|
||||
# the fresh post-rewrite row ids of the surviving user turns, for client
|
||||
# rowId rebinding (see comment at the assignment site).
|
||||
survivor_user_row_ids = None
|
||||
survivor_row_id_map = None
|
||||
raw_rebind_ids = params.get("rebind_survivor_row_ids")
|
||||
requested_rebind_ids = (
|
||||
{
|
||||
row_id
|
||||
for row_id in raw_rebind_ids
|
||||
if isinstance(row_id, int) and not isinstance(row_id, bool)
|
||||
}
|
||||
if isinstance(raw_rebind_ids, list)
|
||||
else None
|
||||
)
|
||||
with session["history_lock"]:
|
||||
# A watch session's run lives in the PARENT turn, so its own running
|
||||
# flag is False — without this, typing mid-run builds a second agent
|
||||
@@ -394,7 +424,9 @@ def _(rid, params: dict) -> dict:
|
||||
or truncate_message_id is not None
|
||||
or truncate_row_id is not None
|
||||
):
|
||||
history = session.get("history", [])
|
||||
history = _history_without_ephemeral_scaffolding(
|
||||
session.get("history", [])
|
||||
)
|
||||
|
||||
# Malformed params refuse first (4004), regardless of consent —
|
||||
# the historical ordinal-path precedence.
|
||||
@@ -451,12 +483,10 @@ def _(rid, params: dict) -> dict:
|
||||
# ordinal and a tip-relative ordinal below can translate, instead
|
||||
# of loading ancestors into the tip (which would duplicate
|
||||
# compressed history on later resumes).
|
||||
prefix_user_count = sum(
|
||||
1
|
||||
for message in session.get("display_history_prefix") or []
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "user"
|
||||
and not message.get("display_kind")
|
||||
prefix_user_count = len(
|
||||
_history_user_indices(
|
||||
session.get("display_history_prefix") or []
|
||||
)
|
||||
)
|
||||
|
||||
user_indices = _history_user_indices(history)
|
||||
@@ -602,7 +632,11 @@ def _(rid, params: dict) -> dict:
|
||||
"target user message is no longer in session history",
|
||||
data=_stale_target_data(resolved_ordinal=ordinal),
|
||||
)
|
||||
truncated = history[: user_indices[ordinal]]
|
||||
from agent.context_compressor import history_before_user_originated_turn
|
||||
|
||||
truncated, _live_view = history_before_user_originated_turn(
|
||||
history, user_indices[ordinal]
|
||||
)
|
||||
# Second gate, on top of confirm_truncate: ordinal 0 resolves to
|
||||
# history[:0] == [] and replace_messages() DELETEs every durable
|
||||
# row. A confirmed rewind that happens to erase the whole
|
||||
@@ -679,11 +713,48 @@ def _(rid, params: dict) -> dict:
|
||||
# default fix have no key, and replace_messages(None)
|
||||
# triggers an FK violation.
|
||||
truncation_key = session.get("session_key") or sid
|
||||
old_active_row_ids = {
|
||||
row_id
|
||||
for message in history
|
||||
if isinstance(
|
||||
(row_id := _message_row_id(message)), int
|
||||
)
|
||||
}
|
||||
if requested_rebind_ids is not None:
|
||||
# Row-id fallback can resolve a durable target even
|
||||
# when the live list is too misaligned to stamp safely,
|
||||
# and alternation repair can merge a physical user;user
|
||||
# pair while preserving only the first row id. Read the
|
||||
# authoritative un-repaired pre-write active-id set so
|
||||
# a rewritten row is never mistaken for an untouched
|
||||
# archived/ancestor row by the bounded client map.
|
||||
durable_rebind_history = (
|
||||
_load_durable_truncation_history(
|
||||
session,
|
||||
truncation_key,
|
||||
repair_alternation=False,
|
||||
)
|
||||
)
|
||||
if durable_rebind_history is None:
|
||||
raise RuntimeError(
|
||||
"could not load durable row identities for truncation"
|
||||
)
|
||||
old_active_row_ids.update(
|
||||
row_id
|
||||
for message in durable_rebind_history
|
||||
if isinstance(
|
||||
(row_id := _message_row_id(message)), int
|
||||
)
|
||||
)
|
||||
old_survivor_row_ids = [
|
||||
_message_row_id(message) for message in truncated
|
||||
]
|
||||
db.replace_messages(
|
||||
truncation_key,
|
||||
truncated,
|
||||
active_only=True,
|
||||
archive_dropped=True,
|
||||
reject_active_turn_lease=True,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
@@ -716,6 +787,26 @@ def _(rid, params: dict) -> dict:
|
||||
_message_row_id(truncated[i])
|
||||
for i in _history_user_indices(truncated)
|
||||
]
|
||||
if requested_rebind_ids is not None:
|
||||
survivor_row_id_map = {
|
||||
str(old_row_id): new_row_id
|
||||
for old_row_id, new_row_id in zip(
|
||||
old_survivor_row_ids,
|
||||
(
|
||||
_message_row_id(message)
|
||||
for message in truncated
|
||||
),
|
||||
)
|
||||
if isinstance(old_row_id, int)
|
||||
and isinstance(new_row_id, int)
|
||||
and old_row_id in requested_rebind_ids
|
||||
}
|
||||
for dropped_row_id in requested_rebind_ids.intersection(
|
||||
old_active_row_ids
|
||||
):
|
||||
survivor_row_id_map.setdefault(
|
||||
str(dropped_row_id), None
|
||||
)
|
||||
session["history"] = truncated
|
||||
session["history_version"] = int(session.get("history_version", 0)) + 1
|
||||
session["running"] = True
|
||||
@@ -728,13 +819,15 @@ def _(rid, params: dict) -> dict:
|
||||
rid, sid, session, text, display_kind=display_kind
|
||||
)
|
||||
if not isolated_response.get("error"):
|
||||
if survivor_user_row_ids is not None:
|
||||
if survivor_user_row_ids is not None and requested_rebind_ids is None:
|
||||
# The truncation already happened inline above (memory + DB),
|
||||
# before compute-host dispatch — the rebind payload applies to
|
||||
# this path exactly as it does to the inline one.
|
||||
isolated_response["result"][
|
||||
"survivor_user_row_ids"
|
||||
] = survivor_user_row_ids
|
||||
if survivor_row_id_map is not None:
|
||||
isolated_response["result"]["survivor_row_id_map"] = survivor_row_id_map
|
||||
return isolated_response
|
||||
logger.warning(
|
||||
"compute-host dispatch failed for session %s; falling back inline: %s",
|
||||
@@ -829,6 +922,12 @@ def _(rid, params: dict) -> dict:
|
||||
**(
|
||||
{"survivor_user_row_ids": survivor_user_row_ids}
|
||||
if survivor_user_row_ids is not None
|
||||
and requested_rebind_ids is None
|
||||
else {}
|
||||
),
|
||||
**(
|
||||
{"survivor_row_id_map": survivor_row_id_map}
|
||||
if survivor_row_id_map is not None
|
||||
else {}
|
||||
),
|
||||
},
|
||||
|
||||
@@ -2821,25 +2821,34 @@ def _(rid, params: dict) -> dict:
|
||||
)
|
||||
removed = 0
|
||||
with session["history_lock"]:
|
||||
history = session.get("history", [])
|
||||
if session.get("running"):
|
||||
return _err(
|
||||
rid, 4009, "session busy — /interrupt the current turn before /undo"
|
||||
)
|
||||
history = _history_without_ephemeral_scaffolding(
|
||||
session.get("history", [])
|
||||
)
|
||||
# Truncate from the last *real* user turn. Popping only trailing
|
||||
# assistant/tool then one user left timeline markers
|
||||
# (async_delegation_complete, model_switch, …) or compaction
|
||||
# handoffs as the undo target — so session.undo removed
|
||||
# bookkeeping instead of the last exchange (#80622).
|
||||
# Match list_recent_user_messages / CLI turn counting.
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
# Match user_originated_turn_view / CLI turn counting.
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
|
||||
last_user_idx = None
|
||||
for i in range(len(history) - 1, -1, -1):
|
||||
msg = history[i]
|
||||
if is_user_originated_turn(msg):
|
||||
last_user_idx = i
|
||||
break
|
||||
if last_user_idx is not None:
|
||||
removed = len(history) - last_user_idx
|
||||
del history[last_user_idx:]
|
||||
session["history_version"] = int(session.get("history_version", 0)) + 1
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(history)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if user_indices:
|
||||
try:
|
||||
_installed, _live_view, rewound_count = (
|
||||
_rewind_active_session_history(session, len(user_indices) - 1)
|
||||
)
|
||||
removed = rewound_count
|
||||
except Exception as exc:
|
||||
return _err(rid, 5008, f"undo: {exc}")
|
||||
return _ok(rid, {"removed": removed})
|
||||
|
||||
|
||||
|
||||
+112
-115
@@ -706,41 +706,55 @@ def _(rid, params: dict) -> dict:
|
||||
return _err(
|
||||
rid, 4009, "session busy — /interrupt the current turn before /retry"
|
||||
)
|
||||
history = session.get("history", [])
|
||||
if not history:
|
||||
return _err(rid, 4018, "no previous user message to retry")
|
||||
# Walk backwards to the last *real* user turn. Timeline bookkeeping
|
||||
# rows (display_kind set) and compaction handoffs are durable
|
||||
# role=user but must not count as user-originated asks — same
|
||||
# predicate as CLI resume/count and the prompt.submit ordinal fix.
|
||||
# Without this, /retry re-sends opaque markers (model_switch /
|
||||
# async_delegation_complete / auto_continue / CONTEXT COMPACTION
|
||||
# handoffs) and truncates only the marker instead of the failed
|
||||
# exchange (#80622).
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
retryable_user_text,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
|
||||
last_user_idx = None
|
||||
for i in range(len(history) - 1, -1, -1):
|
||||
msg = history[i]
|
||||
if is_user_originated_turn(msg):
|
||||
last_user_idx = i
|
||||
break
|
||||
if last_user_idx is None:
|
||||
return _err(rid, 4018, "no previous user message to retry")
|
||||
content = history[last_user_idx].get("content", "")
|
||||
if isinstance(content, list):
|
||||
content = " ".join(
|
||||
p.get("text", "")
|
||||
for p in content
|
||||
if isinstance(p, dict) and p.get("type") == "text"
|
||||
)
|
||||
if not content:
|
||||
return _err(rid, 4018, "last user message is empty")
|
||||
# Truncate history: remove everything from the last user message onward
|
||||
# (mirrors CLI retry_last() which strips the failed exchange)
|
||||
with session["history_lock"]:
|
||||
session["history"] = history[:last_user_idx]
|
||||
session["history_version"] = int(session.get("history_version", 0)) + 1
|
||||
if session.get("running"):
|
||||
return _err(
|
||||
rid,
|
||||
4009,
|
||||
"session busy — /interrupt the current turn before /retry",
|
||||
)
|
||||
if session.get("attached_images"):
|
||||
return _err(
|
||||
rid,
|
||||
4018,
|
||||
"retry cannot safely reconstruct or combine attached media",
|
||||
)
|
||||
history = _history_without_ephemeral_scaffolding(
|
||||
session.get("history", [])
|
||||
)
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(history)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if not user_indices:
|
||||
return _err(rid, 4018, "no previous user message to retry")
|
||||
_prefix, live_view = history_before_user_originated_turn(
|
||||
history, user_indices[-1]
|
||||
)
|
||||
try:
|
||||
content = retryable_user_text(live_view.get("content"))
|
||||
except ValueError as exc:
|
||||
return _err(rid, 4018, str(exc))
|
||||
try:
|
||||
_active, durable_live_view, _rewound_count = (
|
||||
_rewind_active_session_history(
|
||||
session,
|
||||
len(user_indices) - 1,
|
||||
require_retryable=True,
|
||||
)
|
||||
)
|
||||
except ValueError as exc:
|
||||
return _err(rid, 4018, str(exc))
|
||||
except Exception as exc:
|
||||
return _err(rid, 5008, f"retry: failed to persist history: {exc}")
|
||||
content = retryable_user_text(durable_live_view.get("content"))
|
||||
return _ok(rid, {"type": "send", "message": content})
|
||||
|
||||
if name == "steer":
|
||||
@@ -887,58 +901,52 @@ def _(rid, params: dict) -> dict:
|
||||
return _err(
|
||||
rid, 4009, "session busy — /interrupt the current turn before /undo"
|
||||
)
|
||||
# _session_db, not _get_db(): every read and write below is scoped by
|
||||
# session id, and a profile session (app-global remote mode) keeps its
|
||||
# rows in its own profile's state.db. Against the launch handle
|
||||
# list_recent_user_messages finds nothing, so /undo fails closed with
|
||||
# 4018 for the whole session instead of rewinding anything.
|
||||
with _session_db(session) as db:
|
||||
if db is None:
|
||||
return _db_unavailable_error(rid, code=5008)
|
||||
session_key = session.get("session_key", "")
|
||||
if not session_key:
|
||||
return _err(rid, 4001, "no session key for undo")
|
||||
# Parse the optional count argument (e.g. "/undo 3" → 3).
|
||||
session_key = session.get("session_key", "")
|
||||
if not session_key:
|
||||
return _err(rid, 4001, "no session key for undo")
|
||||
# Parse the optional count argument (e.g. "/undo 3" → 3).
|
||||
n = 1
|
||||
arg_str = (arg or "").strip()
|
||||
if arg_str:
|
||||
try:
|
||||
n = int(arg_str.split()[0])
|
||||
except (ValueError, IndexError):
|
||||
return _err(rid, 4004, f"undo: invalid count {arg_str!r} — use /undo or /undo N")
|
||||
if n < 1:
|
||||
n = 1
|
||||
arg_str = (arg or "").strip()
|
||||
if arg_str:
|
||||
try:
|
||||
n = int(arg_str.split()[0])
|
||||
except (ValueError, IndexError):
|
||||
return _err(rid, 4004, f"undo: invalid count {arg_str!r} — use /undo or /undo N")
|
||||
if n < 1:
|
||||
n = 1
|
||||
try:
|
||||
recents = db.list_recent_user_messages(session_key, limit=max(n, 10))
|
||||
except Exception as e:
|
||||
return _err(rid, 5008, f"undo: failed to load history: {e}")
|
||||
if not recents:
|
||||
return _err(rid, 4018, "no user messages to undo")
|
||||
# recents[0] is the most-recent user turn; pick the Nth-from-last.
|
||||
# If N exceeds the number of user turns, back up to the oldest.
|
||||
target_idx = min(n - 1, len(recents) - 1)
|
||||
target_id = recents[target_idx]["id"]
|
||||
try:
|
||||
result = db.rewind_to_message(session_key, target_id)
|
||||
except ValueError as e:
|
||||
return _err(rid, 4004, f"undo: {e}")
|
||||
except Exception as e:
|
||||
return _err(rid, 5008, f"undo: {e}")
|
||||
# Reload the active-only transcript into the in-memory session
|
||||
# history so subsequent turns see the truncated view.
|
||||
# repair_alternation: this reload feeds LIVE REPLAY — session["history"]
|
||||
# is the working conversation for subsequent turns, and a rewind that
|
||||
# lands on a durable user;user pair would otherwise re-fire the
|
||||
# pre-request repair on every request from here on.
|
||||
try:
|
||||
active = db.get_messages_as_conversation(
|
||||
session_key, repair_alternation=True, include_row_ids=True
|
||||
)
|
||||
except Exception:
|
||||
active = []
|
||||
from agent.context_compressor import (
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from agent.message_content import flatten_message_text
|
||||
|
||||
with session["history_lock"]:
|
||||
session["history"] = list(active)
|
||||
session["history_version"] = int(session.get("history_version", 0)) + 1
|
||||
if session.get("running"):
|
||||
return _err(
|
||||
rid,
|
||||
4009,
|
||||
"session busy — /interrupt the current turn before /undo",
|
||||
)
|
||||
history = _history_without_ephemeral_scaffolding(
|
||||
session.get("history", [])
|
||||
)
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(history)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if not user_indices:
|
||||
return _err(rid, 4018, "no user messages to undo")
|
||||
turns_undone = min(n, len(user_indices))
|
||||
target_position = len(user_indices) - turns_undone
|
||||
try:
|
||||
active, live_view, rewound_count = _rewind_active_session_history(
|
||||
session, target_position
|
||||
)
|
||||
except ValueError as exc:
|
||||
return _err(rid, 4004, f"undo: {exc}")
|
||||
except Exception as exc:
|
||||
return _err(rid, 5008, f"undo: {exc}")
|
||||
target_text = flatten_message_text(live_view.get("content"))
|
||||
# Notify memory providers — same hook /branch fires, plus the
|
||||
# rewound flag so providers caching per-turn document state
|
||||
# know to invalidate. See #6672 + #21910.
|
||||
@@ -965,18 +973,6 @@ def _(rid, params: dict) -> dict:
|
||||
agent._last_flushed_db_idx = len(active)
|
||||
except Exception:
|
||||
pass
|
||||
target_msg = result.get("target_message") or {}
|
||||
target_text = target_msg.get("content") or ""
|
||||
if isinstance(target_text, list):
|
||||
parts = [
|
||||
p.get("text", "") for p in target_text
|
||||
if isinstance(p, dict) and p.get("type") == "text"
|
||||
]
|
||||
target_text = "\n".join(t for t in parts if t)
|
||||
if not isinstance(target_text, str):
|
||||
target_text = ""
|
||||
rewound_count = result.get("rewound_count", 0)
|
||||
turns_undone = target_idx + 1
|
||||
turn_word = "turn" if turns_undone == 1 else "turns"
|
||||
notice = (
|
||||
f"↶ Undid {turns_undone} {turn_word} ({rewound_count} message(s)). "
|
||||
@@ -1359,27 +1355,28 @@ def _(rid, params: dict) -> dict:
|
||||
if result.get("success") and not file_path:
|
||||
removed = 0
|
||||
with session["history_lock"]:
|
||||
history = session.get("history", [])
|
||||
# Truncate from the last *real* user turn. Same predicate
|
||||
# as list_recent_user_messages / /undo / /retry —
|
||||
# is_user_originated_turn also excludes compaction
|
||||
# handoffs (durable role=user, sometimes without
|
||||
# display_kind on legacy sessions; #80622).
|
||||
from agent.context_compressor import is_user_originated_turn
|
||||
history = _history_without_ephemeral_scaffolding(
|
||||
session.get("history", [])
|
||||
)
|
||||
from agent.context_compressor import user_originated_turn_view
|
||||
|
||||
last_user_idx = None
|
||||
for i in range(len(history) - 1, -1, -1):
|
||||
msg = history[i]
|
||||
if is_user_originated_turn(msg):
|
||||
last_user_idx = i
|
||||
break
|
||||
if last_user_idx is not None:
|
||||
removed = len(history) - last_user_idx
|
||||
del history[last_user_idx:]
|
||||
if removed:
|
||||
session["history_version"] = (
|
||||
int(session.get("history_version", 0)) + 1
|
||||
)
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(history)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if user_indices:
|
||||
try:
|
||||
_active, _live_view, removed = (
|
||||
_rewind_active_session_history(
|
||||
session, len(user_indices) - 1
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
raise RuntimeError(
|
||||
"checkpoint restored, but session history rewind "
|
||||
f"failed: {exc}"
|
||||
) from exc
|
||||
result["history_removed"] = removed
|
||||
return result
|
||||
|
||||
|
||||
@@ -3508,6 +3508,159 @@ def _session_db(session: dict):
|
||||
db.close()
|
||||
|
||||
|
||||
def _rewind_active_session_history(
|
||||
session: dict,
|
||||
user_ordinal: int,
|
||||
*,
|
||||
require_retryable: bool = False,
|
||||
) -> tuple[list[dict], dict, int]:
|
||||
"""Rewind one canonical user turn while retaining carrier scaffolding.
|
||||
|
||||
The caller holds ``history_lock``. Persistent sessions archive the target
|
||||
and tail; a composite carrier's own hidden handoff is inserted in that same
|
||||
transaction. Memory is installed only after the durable commit and is
|
||||
built from the already-validated prefix plus the returned scaffold row id,
|
||||
so there is no fallible post-commit reload.
|
||||
"""
|
||||
from agent.context_compressor import (
|
||||
history_before_user_originated_turn,
|
||||
retryable_user_text,
|
||||
split_user_originated_turn,
|
||||
user_originated_turn_view,
|
||||
)
|
||||
from agent.memory_manager import sanitize_context
|
||||
from agent.tool_dispatch_helpers import (
|
||||
_is_multimodal_tool_result,
|
||||
_multimodal_text_summary,
|
||||
)
|
||||
|
||||
def _comparison_content(message: dict) -> Any:
|
||||
content = message.get("content")
|
||||
if _is_multimodal_tool_result(content):
|
||||
content = _multimodal_text_summary(content)
|
||||
elif isinstance(content, list):
|
||||
text_parts = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
text_parts.append(str(part.get("text", "")))
|
||||
elif (
|
||||
isinstance(part, dict)
|
||||
and part.get("type") in {"image", "image_url", "input_image"}
|
||||
):
|
||||
text_parts.append("[screenshot]")
|
||||
content = "\n".join(text_parts) if text_parts else None
|
||||
if message.get("role") in {"user", "assistant"} and isinstance(content, str):
|
||||
return sanitize_context(content).strip()
|
||||
return content
|
||||
|
||||
history = _history_without_ephemeral_scaffolding(session.get("history", []))
|
||||
user_indices = [
|
||||
index
|
||||
for index, message in enumerate(history)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if user_ordinal < 0 or user_ordinal >= len(user_indices):
|
||||
raise ValueError("target user message is no longer in session history")
|
||||
target_index = user_indices[user_ordinal]
|
||||
installed, live_view = history_before_user_originated_turn(history, target_index)
|
||||
rewound_count = len(history) - target_index
|
||||
|
||||
session_key = str(session.get("session_key") or "").strip()
|
||||
persisted = False
|
||||
if session_key:
|
||||
with _session_db(session) as db:
|
||||
if db is None:
|
||||
raise RuntimeError("session database is unavailable")
|
||||
expected_active_ids = db.get_active_message_ids(session_key)
|
||||
durable = db.get_messages_as_conversation(
|
||||
session_key,
|
||||
include_row_ids=True,
|
||||
)
|
||||
durable_user_indices = [
|
||||
index
|
||||
for index, message in enumerate(durable)
|
||||
if user_originated_turn_view(message) is not None
|
||||
]
|
||||
if len(durable_user_indices) != len(user_indices):
|
||||
raise RuntimeError(
|
||||
"session history changed before the rewind could be persisted"
|
||||
)
|
||||
durable_target_index = durable_user_indices[user_ordinal]
|
||||
durable_target = durable[durable_target_index]
|
||||
durable_prefix, durable_live_view = history_before_user_originated_turn(
|
||||
durable, durable_target_index
|
||||
)
|
||||
if _comparison_content(durable_live_view) != _comparison_content(live_view):
|
||||
raise RuntimeError(
|
||||
"session history changed before the rewind could be persisted"
|
||||
)
|
||||
target_row_id = durable_target.get("_row_id")
|
||||
if not isinstance(target_row_id, int):
|
||||
raise RuntimeError("rewind target has no durable row identity")
|
||||
if require_retryable:
|
||||
retryable_user_text(durable_live_view.get("content"))
|
||||
scaffold, _ = split_user_originated_turn(durable_target)
|
||||
result = db.rewind_to_message(
|
||||
session_key,
|
||||
target_row_id,
|
||||
preserve_compaction_handoff=scaffold is not None,
|
||||
expected_active_ids=expected_active_ids,
|
||||
expected_target_content=durable_live_view.get("content"),
|
||||
)
|
||||
if scaffold is not None:
|
||||
replacement_id = result.get("replacement_message_id")
|
||||
if not isinstance(replacement_id, int):
|
||||
raise RuntimeError(
|
||||
"rewind commit did not return the replacement scaffold id"
|
||||
)
|
||||
durable_prefix[-1]["_row_id"] = replacement_id
|
||||
durable_prefix[-1]["_db_persisted"] = True
|
||||
installed[-1] = durable_prefix[-1]
|
||||
# Current clients address destructive follow-ups by durable row id.
|
||||
# Preserve the richer warm content (for example image parts), but
|
||||
# copy row identities when the retained warm/durable shapes align.
|
||||
if len(installed) == len(durable_prefix) and all(
|
||||
warm.get("role") == durable_message.get("role")
|
||||
and bool(warm.get("display_kind"))
|
||||
== bool(durable_message.get("display_kind"))
|
||||
and _comparison_content(warm)
|
||||
== _comparison_content(durable_message)
|
||||
for warm, durable_message in zip(installed, durable_prefix)
|
||||
):
|
||||
for warm, durable_message in zip(installed, durable_prefix):
|
||||
row_id = durable_message.get("_row_id")
|
||||
if isinstance(row_id, int):
|
||||
warm["_row_id"] = row_id
|
||||
live_view = durable_live_view
|
||||
rewound_count = int(result.get("rewound_count", 0))
|
||||
persisted = True
|
||||
elif require_retryable:
|
||||
retryable_user_text(live_view.get("content"))
|
||||
|
||||
installed = [message.copy() for message in installed]
|
||||
session["history"] = installed
|
||||
session["history_version"] = int(session.get("history_version", 0)) + 1
|
||||
agent = session.get("agent")
|
||||
if agent is not None:
|
||||
agent._session_messages = installed
|
||||
if hasattr(agent, "_last_flushed_db_idx"):
|
||||
agent._last_flushed_db_idx = len(installed) if persisted else 0
|
||||
if hasattr(agent, "_db_flush_scan_prefix"):
|
||||
agent._db_flush_scan_prefix = installed[:] if persisted else None
|
||||
return installed, live_view, rewound_count
|
||||
|
||||
|
||||
def _history_without_ephemeral_scaffolding(history: list[dict]) -> list[dict]:
|
||||
"""Return the durable transcript shape without transient recovery rows."""
|
||||
from run_agent import _is_ephemeral_scaffolding
|
||||
|
||||
return [
|
||||
message.copy()
|
||||
for message in history
|
||||
if not _is_ephemeral_scaffolding(message)
|
||||
]
|
||||
|
||||
|
||||
def _persist_session_git_meta(session: dict, cwd: str, generation: int) -> None:
|
||||
"""Resolve + persist a session's git branch / repo root WITHOUT blocking.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user