diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 1c89288330..ccd9d95c8b 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -755,6 +755,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) @@ -1425,8 +1438,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]] = [] @@ -1476,6 +1487,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)", diff --git a/agent/context_compressor.py b/agent/context_compressor.py index 4c48794fb2..47bd2fd4ad 100644 --- a/agent/context_compressor.py +++ b/agent/context_compressor.py @@ -16,6 +16,7 @@ Improvements over v2: - Richer tool call/result detail in summarizer input """ +import copy import hashlib import json import logging @@ -7935,6 +7936,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. @@ -7943,21 +8156,16 @@ def _handoff_carries_live_user_content(message: Any) -> bool: ask, leaving a non-empty remainder after ``_SUMMARY_END_MARKER``. Either 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. + Delegates to the canonical live-user projection. That projection uses + ``_strip_context_summary_handoff_message`` for string and multimodal + carriers, then excludes typed synthetic rows and empty/non-actionable + remainders. Callers must still pre-filter with + ``is_compaction_summary_message`` because an ordinary human row naturally + has a live-user projection too. """ if not isinstance(message, dict): return False - return ( - ContextCompressor._strip_context_summary_handoff_message(message) - is not None - ) + return user_originated_turn_view(message) is not None def reference_handoff_would_drive_next_model_call( @@ -8013,15 +8221,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 diff --git a/agent/turn_context.py b/agent/turn_context.py index df8969a666..e63372d4b7 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -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 diff --git a/apps/desktop/src/app/chat/session-tile-actions.ts b/apps/desktop/src/app/chat/session-tile-actions.ts index 5541976160..528497f3ea 100644 --- a/apps/desktop/src/app/chat/session-tile-actions.ts +++ b/apps/desktop/src/app/chat/session-tile-actions.ts @@ -37,6 +37,7 @@ import { applyBranchVisibility, applyReloadOptimistic, applyRewindOptimistic, + durableRowIdsForRebind, finalizeInterruptedMessages, planEdit, planReload, @@ -410,7 +411,8 @@ export function useSessionTileActions({ runtimeId, scope, storedSessionId }: Ses interruptFirst: boolean, truncateMessageId?: string, truncateRowId?: number, - sourceText?: string + sourceText?: string, + rebindRowIds?: readonly number[] ) => runRewindSubmit( requestGateway, @@ -426,7 +428,8 @@ export function useSessionTileActions({ runtimeId, scope, storedSessionId }: Ses } }, truncateRowId, - sourceText + sourceText, + rebindRowIds ), [requestGateway] ) @@ -474,7 +477,8 @@ export function useSessionTileActions({ runtimeId, scope, storedSessionId }: Ses false, plan.truncateMessageId, plan.truncateRowId, - plan.sourceText + plan.sourceText, + durableRowIdsForRebind(state.messages) ) ) } catch (err) { @@ -510,7 +514,8 @@ export function useSessionTileActions({ runtimeId, scope, storedSessionId }: Ses interruptFirst, plan.truncateMessageId, plan.truncateRowId, - plan.sourceText + plan.sourceText, + durableRowIdsForRebind(messages) ) ) } catch (err) { @@ -558,7 +563,8 @@ export function useSessionTileActions({ runtimeId, scope, storedSessionId }: Ses interruptFirst, plan.truncateMessageId, plan.truncateRowId, - plan.sourceText + plan.sourceText, + durableRowIdsForRebind(messages) ) ) } catch (err) { diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts index 6324176f0f..90f0fb86fa 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts @@ -56,6 +56,7 @@ import { applyBranchVisibility, applyReloadOptimistic, applyRewindOptimistic, + durableRowIdsForRebind, finalizeInterruptedMessages, planEdit, planReload, @@ -834,7 +835,8 @@ export function usePromptActions({ truncateMessageId: string | undefined, interruptFirst: boolean, truncateRowId?: number, - sourceText?: string + sourceText?: string, + rebindRowIds?: readonly number[] ) => runRewindSubmit( requestGateway, @@ -851,7 +853,8 @@ export function usePromptActions({ } }, truncateRowId, - sourceText + sourceText, + rebindRowIds ), [activeSessionIdRef, requestGateway, selectedStoredSessionIdRef] ) @@ -866,7 +869,8 @@ export function usePromptActions({ return } - const plan = planReload($messages.get(), parentId) + const messages = $messages.get() + const plan = planReload(messages, parentId) if (!plan) { return @@ -883,7 +887,8 @@ export function usePromptActions({ plan.truncateMessageId, false, plan.truncateRowId, - plan.sourceText + plan.sourceText, + durableRowIdsForRebind(messages) ) applySurvivorRowIds(sessionId, survivorRowIds) @@ -951,7 +956,8 @@ export function usePromptActions({ plan.truncateMessageId, interruptFirst, plan.truncateRowId, - plan.sourceText + plan.sourceText, + durableRowIdsForRebind(messages) ) applySurvivorRowIds(sessionId, survivorRowIds) @@ -1036,7 +1042,8 @@ export function usePromptActions({ plan.truncateMessageId, interruptFirst, plan.truncateRowId, - plan.sourceText + plan.sourceText, + durableRowIdsForRebind(messages) ) applySurvivorRowIds(sessionId, survivorRowIds) @@ -1069,7 +1076,8 @@ export function usePromptActions({ retryPlan.truncateMessageId, false, retryPlan.truncateRowId, - retryPlan.sourceText + retryPlan.sourceText, + durableRowIdsForRebind(refreshed) ) applySurvivorRowIds(sessionId, survivorRowIds) diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.test.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.test.ts index d5f13f7240..3a42802f61 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.test.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.test.ts @@ -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 () => { diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.ts index 6ca5483cd9..72b25801d3 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/rewind.ts @@ -35,24 +35,41 @@ import { type RequestGateway = (method: string, params?: Record, timeoutMs?: number) => Promise /** - * 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 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() + + 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 { // 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 ) diff --git a/apps/desktop/src/lib/chat-messages.test.ts b/apps/desktop/src/lib/chat-messages.test.ts index 38bc2aca88..c6bc0b28b1 100644 --- a/apps/desktop/src/lib/chat-messages.test.ts +++ b/apps/desktop/src/lib/chat-messages.test.ts @@ -323,6 +323,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 }, diff --git a/apps/desktop/src/lib/chat-messages/hydration.ts b/apps/desktop/src/lib/chat-messages/hydration.ts index 13fef62d32..69bc09a174 100644 --- a/apps/desktop/src/lib/chat-messages/hydration.ts +++ b/apps/desktop/src/lib/chat-messages/hydration.ts @@ -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, diff --git a/apps/desktop/src/types/hermes.ts b/apps/desktop/src/types/hermes.ts index 783ee945e5..994393f524 100644 --- a/apps/desktop/src/types/hermes.ts +++ b/apps/desktop/src/types/hermes.ts @@ -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 diff --git a/cli.py b/cli.py index b53906386a..0d4a5a7386 100644 --- a/cli.py +++ b/cli.py @@ -10133,6 +10133,121 @@ 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 + + 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") + expected_active_ids = [ + int(message["_row_id"]) + for message in durable + if isinstance(message.get("_row_id"), int) + ] + + 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. @@ -10150,22 +10265,68 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # CLI resume counting and list_recent_user_messages. 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 @@ -10204,59 +10365,62 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # messages (exclude display_kind timeline rows and compaction # handoffs — same predicate as list_recent_user_messages, 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. @@ -10271,6 +10435,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: diff --git a/gateway/session.py b/gateway/session.py index 3ba5b10c48..17a9f0a761 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -3587,17 +3587,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 = {} @@ -3923,6 +3927,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. @@ -3943,16 +3948,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. @@ -4004,7 +4019,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), @@ -4015,47 +4036,89 @@ 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: + durable = self._db.get_messages_as_conversation( + session_id, + include_row_ids=True, + ) + expected_active_ids = [ + int(message["_row_id"]) + for message in durable + if isinstance(message.get("_row_id"), int) + ] + 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: + if handoff is None: + return None + target_text = retryable_user_text(target_view.get("content")) + except Exception as e: + logger.debug("rewind_session: failed to resolve canonical target: %s", e) + return None + 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( diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index cd54164bd0..f679724034 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -2644,34 +2644,66 @@ 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. + rewind_result = await self.async_session_store.rewind_session( + session_entry.session_id, + 1, + require_retryable_composite=True, + ) + 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 diff --git a/hermes_cli/web_routers/sessions.py b/hermes_cli/web_routers/sessions.py index 8b37f9fa90..2e88a3ea8e 100644 --- a/hermes_cli/web_routers/sessions.py +++ b/hermes_cli/web_routers/sessions.py @@ -642,14 +642,33 @@ async def get_session_messages( if result is None: raise HTTPException(status_code=404, detail="Session not found") sid, _limit, messages = result + from agent.context_compressor import split_user_originated_turn + + projected_messages = [] + for message in messages: + handoff, live_view = split_user_originated_turn(message) + if handoff is None: + projected_messages.append(message) + continue + projected = message.copy() + if live_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"] = live_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), }, } diff --git a/hermes_state.py b/hermes_state.py index 9fdc820035..0cbef7c254 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -9163,13 +9163,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 @@ -9181,33 +9184,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,), @@ -9830,6 +9878,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. @@ -9864,21 +9913,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 @@ -10216,13 +10279,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"], @@ -10547,6 +10627,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): @@ -10625,9 +10710,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 @@ -10838,15 +10961,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. @@ -10912,25 +11048,118 @@ 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 # ========================================================================= + @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*. @@ -10949,7 +11178,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 @@ -10958,29 +11199,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", @@ -10993,28 +11283,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. diff --git a/run_agent.py b/run_agent.py index f5c4de1274..09214549dd 100644 --- a/run_agent.py +++ b/run_agent.py @@ -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.") / from run_agent import @@ -2340,6 +2341,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") diff --git a/tests/agent/test_reference_handoff_active_turn.py b/tests/agent/test_reference_handoff_active_turn.py index 4e62b0f9b8..45df7c76b6 100644 --- a/tests/agent/test_reference_handoff_active_turn.py +++ b/tests/agent/test_reference_handoff_active_turn.py @@ -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, @@ -17,13 +19,18 @@ from agent.context_compressor import ( HISTORICAL_TASK_HEADING, SUMMARY_PREFIX, _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 @@ -39,6 +46,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 + + class TestReferenceHandoffWouldDriveNextModelCall: def test_standalone_handoff_alone_drives(self): messages = [_standalone_handoff()] @@ -136,6 +149,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( @@ -172,6 +231,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): diff --git a/tests/cli/test_cli_retry.py b/tests/cli/test_cli_retry.py index b287b45754..8386184345 100644 --- a/tests/cli/test_cli_retry.py +++ b/tests/cli/test_cli_retry.py @@ -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\nprivate\n " + )["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() diff --git a/tests/gateway/test_retry_replacement.py b/tests/gateway/test_retry_replacement.py index c83fafbe46..b3f2e8d96f 100644 --- a/tests/gateway/test_retry_replacement.py +++ b/tests/gateway/test_retry_replacement.py @@ -1,15 +1,200 @@ -"""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 + + +@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 +252,195 @@ 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_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 diff --git a/tests/gateway/test_undo_rewind_session.py b/tests/gateway/test_undo_rewind_session.py index b95f34622a..10df30a431 100644 --- a/tests/gateway/test_undo_rewind_session.py +++ b/tests/gateway/test_undo_rewind_session.py @@ -54,3 +54,32 @@ def test_rewind_n_turns(store): assert len(store.load_transcript(sid)) == 2 # 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() + diff --git a/tests/hermes_cli/test_web_server.py b/tests/hermes_cli/test_web_server.py index 5c1fb1e985..b9820c127c 100644 --- a/tests/hermes_cli/test_web_server.py +++ b/tests/hermes_cli/test_web_server.py @@ -2044,6 +2044,56 @@ 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, + _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" + 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, + ) + 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) == 2 + 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" + 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 diff --git a/tests/hermes_state/test_composite_carrier_rewind.py b/tests/hermes_state/test_composite_carrier_rewind.py new file mode 100644 index 0000000000..b086e463ee --- /dev/null +++ b/tests/hermes_state/test_composite_carrier_rewind.py @@ -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 diff --git a/tests/run_agent/test_identity_flush.py b/tests/run_agent/test_identity_flush.py index 5c8a699ead..8cb1e54d6d 100644 --- a/tests/run_agent/test_identity_flush.py +++ b/tests/run_agent/test_identity_flush.py @@ -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 diff --git a/tests/run_agent/test_thinking_only_sanitizer.py b/tests/run_agent/test_thinking_only_sanitizer.py index d6875084ea..f6190fab72 100644 --- a/tests/run_agent/test_thinking_only_sanitizer.py +++ b/tests/run_agent/test_thinking_only_sanitizer.py @@ -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 = [ diff --git a/tests/test_compression_watermark_commit.py b/tests/test_compression_watermark_commit.py index c7e677fbc2..c35fda2d31 100644 --- a/tests/test_compression_watermark_commit.py +++ b/tests/test_compression_watermark_commit.py @@ -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. diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 2f15a65015..0a79a3035d 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -5097,7 +5097,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 = [ @@ -5234,7 +5241,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 = [ @@ -5292,7 +5306,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 = [ @@ -5342,7 +5363,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 = [ @@ -5423,7 +5451,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 = [ @@ -5519,7 +5554,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 = [ @@ -5564,7 +5606,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 = [ @@ -5718,7 +5767,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 = [ @@ -5807,7 +5863,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 = [ @@ -10137,6 +10200,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( @@ -10197,6 +10261,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( @@ -10219,6 +10284,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. @@ -10638,6 +10782,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"}, @@ -11138,7 +11283,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() @@ -11198,7 +11350,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()) @@ -11293,7 +11452,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() @@ -11396,7 +11562,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() @@ -11481,6 +11655,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"}, @@ -11490,6 +11670,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"}, ] @@ -11497,7 +11687,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) @@ -18819,7 +19017,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( @@ -18930,10 +19135,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(), @@ -18972,6 +19183,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) @@ -18993,7 +19205,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 = [ @@ -19037,7 +19256,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. @@ -19095,7 +19321,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): @@ -19280,6 +19513,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 @@ -19311,6 +19545,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, }, } @@ -19349,6 +19584,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"}, @@ -19372,6 +19608,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, }, } @@ -19388,6 +19625,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) @@ -19421,7 +19665,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): @@ -19461,8 +19712,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 @@ -19491,6 +19743,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) @@ -19506,17 +19768,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. @@ -19563,6 +19838,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, ): diff --git a/tests/tui_gateway/test_composite_carrier_rewind.py b/tests/tui_gateway/test_composite_carrier_rewind.py new file mode 100644 index 0000000000..bef256746d --- /dev/null +++ b/tests/tui_gateway/test_composite_carrier_rewind.py @@ -0,0 +1,482 @@ +"""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\nprivate\n ", + ) + 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_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" diff --git a/tui_gateway/methods_prompt.py b/tui_gateway/methods_prompt.py index fa0dcd949f..1adc70b556 100644 --- a/tui_gateway/methods_prompt.py +++ b/tui_gateway/methods_prompt.py @@ -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", @@ -826,6 +919,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 {} ), }, diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index 564271ea13..3da3cf5cd4 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -2675,25 +2675,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 + 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}) diff --git a/tui_gateway/methods_tools.py b/tui_gateway/methods_tools.py index cece5fbbe4..d0acf8d073 100644 --- a/tui_gateway/methods_tools.py +++ b/tui_gateway/methods_tools.py @@ -706,41 +706,49 @@ 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) + ) + 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 +895,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 +967,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 +1349,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 diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 00fe114365..abdea00e6c 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -3141,6 +3141,155 @@ def _session_db(session: dict): db.close() +def _rewind_active_session_history( + session: dict, user_ordinal: int +) -> 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, + 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") + 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") + expected_active_ids = [ + int(message["_row_id"]) + for message in durable + if isinstance(message.get("_row_id"), int) + ] + 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 + + 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. @@ -7676,6 +7825,20 @@ def _history_to_messages(history: list[dict]) -> list[dict]: role = m.get("role") if role not in {"user", "assistant", "tool", "system"}: continue + if role == "user": + from agent.context_compressor import ( + is_compaction_summary_message, + user_originated_turn_view, + ) + + if is_compaction_summary_message(m): + carrier_row_id = m.get("_row_id") + live_view = user_originated_turn_view(m) + if live_view is None: + continue + if carrier_row_id is not None: + live_view["_row_id"] = carrier_row_id + m = live_view # An explicit display_kind="hidden" row is model-facing scaffolding # (compaction references, interrupted-turn checkpoints). The string # sniff below only catches the "[System:" convention; honor the