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:
Teknium
2026-08-24 03:31:29 -07:00
committed by GitHub
31 changed files with 4028 additions and 446 deletions
+16 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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,
+2
View File
@@ -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
+216 -51
View File
@@ -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
View File
@@ -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
View File
@@ -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
+22 -2
View File
@@ -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
View File
@@ -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.
+2
View File
@@ -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 ![alt](https://example.test/a.png) 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
View File
@@ -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()
+434 -1
View File
@@ -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
+90
View File
@@ -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()
+65
View File
@@ -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
+49
View File
@@ -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 = [
+29 -3
View File
@@ -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.
+360 -25
View File
@@ -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
View File
@@ -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 {}
),
},
+22 -13
View File
@@ -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
View File
@@ -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
+153
View File
@@ -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.