diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index c7f0f5dacb81..e92f93d1b31f 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -8,7 +8,7 @@ Methods covered: * ``convert_to_trajectory_format`` — internal -> trajectory-file format * ``sanitize_tool_call_arguments`` — repair corrupted JSON in tool_calls -* ``repair_message_sequence`` — enforce alternation invariants +* ``repair_message_sequence`` — repair canonical assistant/tool structure * ``strip_think_blocks`` — remove inline reasoning from stored content * ``recover_with_credential_pool`` — rotate pool entries on 429 * ``try_recover_primary_transport`` — re-create OpenAI client after rate-limit @@ -560,19 +560,12 @@ def note_turn_persisted(agent): def repair_message_sequence(agent, messages: List[Dict]) -> int: - """Collapse malformed role-alternation left in the live history. + """Repair malformed assistant/tool structure in canonical history. - Providers (OpenAI, OpenRouter, Anthropic) expect strict alternation: - after the system message, user/tool alternates with assistant, with - no two consecutive user messages and no tool-result that doesn't - follow an assistant-with-tool_calls. Violations cause silent empty - responses on most providers, which triggers the empty-retry loop. - - This runs right before the API call as a defensive belt — by the - time it fires, the scaffolding strip should already have prevented - most shapes, but external callers (gateway multi-queue replay, - session resume, cron, explicit conversation_history passed in by - host code) can feed in already-broken histories. + This canonical-history repair deliberately preserves adjacent ``user`` + messages as distinct source turns. Provider role alternation is repaired + later on the per-request ``api_messages`` copy by + :func:`drop_thinking_only_and_merge_users`. Repairs applied: 0. Consecutive ``assistant`` messages with no intervening @@ -589,8 +582,6 @@ def repair_message_sequence(agent, messages: List[Dict]) -> int: resumed histories. Refs #29148, #49147. 1. Stray ``tool`` messages whose ``tool_call_id`` doesn't match any preceding assistant tool_call — dropped. - 2. Consecutive ``user`` messages — merged with newline separator - so no user input is lost. Deliberately does NOT rewind orphan ``assistant(tool_calls)+tool`` pairs that precede a user message — that pattern IS valid when the @@ -759,54 +750,14 @@ def _is_verification_candidate(m: Dict) -> bool: matched_tool_groups = set() filtered.append(msg) - # Pass 2: merge consecutive user messages. Preserves all user input - # so nothing the user typed is lost. - merged: List[Dict] = [] - for msg in filtered: - if ( - merged - and isinstance(msg, dict) - and msg.get("role") == "user" - and isinstance(merged[-1], dict) - 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) - # content alone — collapsing image/audio blocks risks - # mangling the attachment structure. - if isinstance(prev_content, str) and isinstance(new_content, str): - prev["content"] = ( - (prev_content + "\n\n" + new_content) - if prev_content and new_content - else (prev_content or new_content) - ) - # Merged content invalidates the api_content sidecar (exact - # bytes previously sent for the pre-merge message) — drop it - # so replay can't substitute stale bytes. - drop_stale_api_content(prev) - repairs += 1 - continue - merged.append(msg) + # Adjacent user messages are canonical source boundaries, not malformed + # history. Keep them distinct here; the per-request wire copy is merged + # later by ``drop_thinking_only_and_merge_users`` for strict providers. if repairs > 0: # Rewrite in place so downstream paths (persistence, return # value, session DB flush) see the repaired sequence. - messages[:] = merged + messages[:] = filtered return repairs @@ -815,11 +766,11 @@ def repair_message_sequence_with_cursor(agent, messages: List[Dict]) -> int: """Run :func:`repair_message_sequence` and keep the SessionDB flush cursor consistent with the compacted list (#44837). - ``repair_message_sequence`` merges/drops messages in place, shrinking - the list. ``_last_flushed_db_idx`` (the DB-write cursor) indexes into - that list, so after compaction it can point past the new end — the - turn-end flush would then skip the assistant/tool chain entirely — or - past unflushed messages shifted to lower indexes. + ``repair_message_sequence`` merges assistant messages and drops orphaned + tool messages in place, shrinking the list. ``_last_flushed_db_idx`` (the + DB-write cursor) indexes into that list, so after compaction it can point + past the new end — the turn-end flush would then skip the assistant/tool + chain entirely — or past unflushed messages shifted to lower indexes. Repair preserves object identity for surviving messages, so counting the survivors from the previously-flushed prefix gives the exact new @@ -1413,13 +1364,12 @@ def drop_thinking_only_and_merge_users( *, drop_codex_reasoning_items: bool = True, ) -> List[Dict[str, Any]]: - """Drop thinking-only assistant turns; merge any adjacent user messages left behind. + """Drop thinking-only turns and merge adjacent users on the wire copy. Runs on the per-call ``api_messages`` copy only. The stored - conversation history (``agent.messages``) is never mutated, so the - user still sees the thinking block in the CLI/gateway transcript and - session persistence keeps the full trace. Only the wire copy sent to - the provider is cleaned. + conversation history (``agent.messages``) is never mutated, so canonical + user source boundaries and thinking blocks remain available to the UI and + session persistence. Only the wire copy sent to the provider is cleaned. Why drop-and-merge rather than inject stub text: - Fabricating ``"."`` / ``"(continued)"`` text lies in the history @@ -1441,8 +1391,18 @@ def drop_thinking_only_and_merge_users( ) ] dropped = len(messages) - len(kept) + has_adjacent_users = any( + previous.get("role") == "user" and current.get("role") == "user" + for previous, current in zip(kept, kept[1:]) + ) + if dropped == 0 and not has_adjacent_users: + return messages - # Pass 2: merge any newly-adjacent user messages. + # Pass 2: merge adjacent source turns for provider compatibility while + # retaining an explicit semantic boundary in the transient wire content. + # This marker never reaches canonical history or SessionDB. + boundary_text = "[Next user message]" + boundary_block = {"type": "text", "text": boundary_text} merged: List[Dict[str, Any]] = [] merges = 0 for m in kept: @@ -1463,21 +1423,25 @@ def drop_thinking_only_and_merge_users( # purposes. If either side is a list (multimodal), append as a # separate block rather than collapsing. if isinstance(prev_content, str) and isinstance(cur_content, str): - sep = "\n\n" if prev_content and cur_content else "" - prev_copy["content"] = prev_content + sep + cur_content + prev_copy["content"] = "\n\n".join( + part + for part in (prev_content, boundary_text, cur_content) + if part + ) elif isinstance(prev_content, list) and isinstance(cur_content, list): - prev_copy["content"] = list(prev_content) + list(cur_content) + prev_copy["content"] = ( + list(prev_content) + [dict(boundary_block)] + list(cur_content) + ) elif isinstance(prev_content, list) and isinstance(cur_content, str): + new_blocks = list(prev_content) + [dict(boundary_block)] if cur_content: - prev_copy["content"] = list(prev_content) + [ - {"type": "text", "text": cur_content} - ] - else: - prev_copy["content"] = list(prev_content) + new_blocks.append({"type": "text", "text": cur_content}) + prev_copy["content"] = new_blocks elif isinstance(prev_content, str) and isinstance(cur_content, list): new_blocks: List[Dict[str, Any]] = [] if prev_content: new_blocks.append({"type": "text", "text": prev_content}) + new_blocks.append(dict(boundary_block)) new_blocks.extend(cur_content) prev_copy["content"] = new_blocks else: @@ -3713,6 +3677,21 @@ def repair_empty_non_final_messages( return repaired return messages +_API_SOURCE_METADATA_KEYS = ( + "timestamp", + "message_id", + "platform_message_id", + "_source_message_id", +) + + +def copy_message_for_api(message: Dict[str, Any]) -> Dict[str, Any]: + """Copy one canonical message without transcript-only source metadata.""" + api_message = message.copy() + for key in _API_SOURCE_METADATA_KEYS: + api_message.pop(key, None) + return api_message + def sanitize_api_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: """Fix orphaned tool_call / tool_result pairs before every LLM call. @@ -4644,6 +4623,7 @@ def force_close_tcp_sockets(client: Any) -> int: "invoke_tool", "repair_tool_call", "sanitize_api_messages", + "copy_message_for_api", "looks_like_codex_intermediate_ack", "copy_reasoning_content_for_api", "cleanup_dead_connections", diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index 936d1d7463bc..b664f3ca9574 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -2924,12 +2924,14 @@ def _managed_summary_call(request, callback, *, retry_count: int): append_message(messages, {"role": "user", "content": summary_request}) try: + from agent.agent_runtime_helpers import copy_message_for_api + # Build API messages, stripping internal-only fields - # (finish_reason, reasoning) that strict APIs like Mistral reject with 422 + # (finish_reason, reasoning) that strict APIs like Mistral reject with 422. _needs_sanitize = agent._should_sanitize_tool_calls() api_messages = [] for msg in messages: - api_msg = msg.copy() + api_msg = copy_message_for_api(msg) agent._copy_reasoning_content_for_api(msg, api_msg) for internal_field in ("reasoning", "finish_reason"): api_msg.pop(internal_field, None) @@ -2937,13 +2939,10 @@ def _managed_summary_call(request, callback, *, retry_count: int): # Mistral, Moonshot/Kimi) reject any message key outside the Chat # Completions schema. The main loop drops these via # ChatCompletionsTransport.convert_messages(), but the summary path - # hand-builds messages and calls chat.completions.create() directly, - # bypassing the transport — so mirror that sanitization here: - # tool_name (SQLite FTS bookkeeping), the codex_* reasoning carriers, - # timestamp (preserved on gateway user replay entries for the - # stale-confirmation expiry check — #47868 rejection class), - # and every Hermes-internal underscore-prefixed scaffolding key. - for schema_foreign in ("tool_name", "codex_reasoning_items", "codex_message_items", "timestamp"): + # calls chat.completions.create() directly. Reuse the shared source + # metadata copy policy above, then strip remaining schema-foreign + # bookkeeping and every Hermes-internal underscore-prefixed key. + for schema_foreign in ("tool_name", "codex_reasoning_items", "codex_message_items"): api_msg.pop(schema_foreign, None) # api_content (the persist-what-you-send sidecar) carries the # exact bytes every main-loop call sent for this message — diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index f6f36d1903e9..302d5a5b65f1 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -1832,6 +1832,7 @@ def run_conversation( persist_user_display_kind: Optional[str] = None, persist_user_display_metadata: Optional[Dict[str, Any]] = None, moa_config: Optional[dict[str, Any]] = None, + persist_user_message_id: Optional[str] = None, ) -> Dict[str, Any]: """ Run a complete conversation with tool calling until completion. @@ -1857,6 +1858,9 @@ def run_conversation( persist_user_display_metadata: Optional payload for that event (e.g. a delegation's task count). or queuing follow-up prefetch work. + moa_config: Optional mixture-of-agents configuration for this turn. + persist_user_message_id: Optional stable source identity to retain on + the canonical user message and persisted row. Returns: Dict: Complete conversation result with final response and message history @@ -1908,6 +1912,7 @@ def run_conversation( persist_user_timestamp, persist_user_display_kind=persist_user_display_kind, persist_user_display_metadata=persist_user_display_metadata, + persist_user_message_id=persist_user_message_id, restore_or_build_system_prompt=_restore_or_build_system_prompt, install_safe_stdio=_install_safe_stdio, sanitize_surrogates=_sanitize_surrogates, @@ -2216,13 +2221,15 @@ def run_conversation( ) ] - # Defensive: repair malformed role-alternation before API call. - # Catches cases where the history got wedged into a - # ``tool → user`` or ``user → user`` tail (e.g. after empty- - # response scaffolding was stripped and a new user message - # landed after an orphan tool result). Most providers return - # empty content on malformed sequences, which would otherwise - # retrigger the empty-retry loop indefinitely. + # Defensive: repair malformed assistant/tool structure in canonical + # history before the API call. This collapses split assistant turns and + # drops orphaned tool results without collapsing adjacent user source + # messages. Strict-provider user-role alternation is repaired later by + # ``_drop_thinking_only_and_merge_users`` on the per-request copy. + # ``repair_message_sequence_with_cursor`` also recomputes the SessionDB + # flush cursor (_last_flushed_db_idx) when canonical repair compacts the + # list, so turn-end flushing cannot skip shifted assistant/tool rows + # (#44837). # repair_message_sequence_with_cursor also recomputes the SessionDB # flush cursor (_last_flushed_db_idx) when repair compacts the list, # so the turn-end flush doesn't skip the assistant/tool chain (#44837). @@ -2284,6 +2291,17 @@ def run_conversation( # Bookkeeping, never a provider field — only the chat-completions # transport strips underscore keys, so drop it centrally here. api_msg.pop("_row_id", None) + # Source ordering/deduplication metadata belongs to the canonical + # transcript and SessionDB, never to a provider request. Strip it + # at the common API-copy boundary so native Anthropic/Codex paths + # do not depend on transport-specific unknown-field filtering. + for metadata_key in ( + "timestamp", + "message_id", + "platform_message_id", + "_source_message_id", + ): + api_msg.pop(metadata_key, None) # Inject ephemeral context into the current turn's user message. # Sources: memory manager prefetch + plugin pre_llm_call hooks diff --git a/agent/transports/chat_completions.py b/agent/transports/chat_completions.py index bec0f9a82b0a..4b2439be1f16 100644 --- a/agent/transports/chat_completions.py +++ b/agent/transports/chat_completions.py @@ -285,6 +285,9 @@ def convert_messages( ``Extra inputs are not permitted, field: 'messages[N].tool_name'``. Permissive providers (OpenRouter, MiniMax) silently ignore the field, which masked the bug for months. + - Transcript metadata: ``timestamp``, ``message_id``, and + ``platform_message_id`` are persisted for ordering/deduplication but + are not Chat Completions message fields. - Hermes-internal scaffolding markers — any top-level message key starting with ``_`` (e.g. ``_empty_recovery_synthetic``, ``_empty_terminal_sentinel``, ``_thinking_prefill``). These are @@ -310,6 +313,8 @@ def convert_messages( or "effect_disposition" in msg or "timestamp" in msg # #47868 — strict providers reject this or "api_content" in msg # persist-what-you-send sidecar + or "message_id" in msg + or "platform_message_id" in msg ): needs_sanitize = True break @@ -381,6 +386,8 @@ def mutable_msg() -> dict[str, Any]: or "effect_disposition" in msg or "timestamp" in msg # #47868 — leak into strict providers or "api_content" in msg # persist-what-you-send sidecar + or "message_id" in msg + or "platform_message_id" in msg ): out_msg = mutable_msg() out_msg.pop("codex_reasoning_items", None) @@ -389,6 +396,8 @@ def mutable_msg() -> dict[str, Any]: out_msg.pop("effect_disposition", None) out_msg.pop("timestamp", None) # #47868 — leak into strict providers out_msg.pop("api_content", None) # persist-what-you-send sidecar + out_msg.pop("message_id", None) + out_msg.pop("platform_message_id", None) # Drop all Hermes-internal scaffolding markers (``_``-prefixed). diff --git a/agent/turn_context.py b/agent/turn_context.py index 1a7a72c0f0b9..327445dda35b 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -502,6 +502,7 @@ def build_turn_context( stream_callback, persist_user_message: Optional[Any], persist_user_timestamp: Optional[float] = None, + persist_user_message_id: Optional[str] = None, *, persist_user_display_kind: Optional[str] = None, persist_user_display_metadata: Optional[Dict[str, Any]] = None, @@ -721,6 +722,8 @@ def build_turn_context( # CLI input is stamped when staged. Gateway input may carry the platform # event time. Preserve either value and cover any legacy unstamped handoff. stamp_message_timestamp(user_msg, timestamp=persist_user_timestamp) + if persist_user_message_id is not None: + user_msg["_source_message_id"] = persist_user_message_id # Hydrate todo store from conversation history. if conversation_history and not agent._todo_store.has_items(): @@ -784,6 +787,7 @@ def build_turn_context( should_review_memory = True agent._turns_since_memory = 0 + # Cosmetic side-signal: detect an affection "reaction" (ily / <3 / good bot) # and notify the host so it can play hearts. Token-free, never touches the # conversation, and never fatal — a purely optional UI beat. diff --git a/apps/desktop/src/app/chat/composer/hooks/use-composer-queue.ts b/apps/desktop/src/app/chat/composer/hooks/use-composer-queue.ts index d1731aaf6ee8..3984249b7419 100644 --- a/apps/desktop/src/app/chat/composer/hooks/use-composer-queue.ts +++ b/apps/desktop/src/app/chat/composer/hooks/use-composer-queue.ts @@ -220,8 +220,10 @@ export function useComposerQueue({ attachments: entry.attachments, ...(entry.displayText ? { displayText: entry.displayText } : {}), fromQueue: true, + messageId: entry.id, sessionId: drainRuntimeSessionId, - storedSessionId: drainQueueSessionKey + storedSessionId: drainQueueSessionKey, + submittedAt: entry.queuedAt }) ) diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx index 0bc1f47c65ad..77ec28af9602 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx @@ -4,11 +4,13 @@ import type { MutableRefObject } from 'react' import { useEffect, useRef } from 'react' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import type { QueueEditState } from '@/app/chat/composer/composer-utils' +import { useComposerQueue } from '@/app/chat/composer/hooks/use-composer-queue' import { getSession } from '@/hermes' import { textPart } from '@/lib/chat-messages' import { createClientSessionState } from '@/lib/chat-runtime' import { $composerAttachments, $composerDraft, type ComposerAttachment, setComposerDraft } from '@/store/composer' -import { $queuedPromptsBySession, getQueuedPrompts } from '@/store/composer-queue' +import { $queuedPromptsBySession, enqueueQueuedPrompt, getQueuedPrompts } from '@/store/composer-queue' import { requestGatewayForAgent } from '@/store/gateway' import { $goalsBySession, setSessionGoal } from '@/store/goals' import { $hudMode } from '@/store/hud' @@ -248,6 +250,107 @@ function Harness({ return null } +interface QueueHarnessHandle { + drainNextQueued: () => Promise +} + +function QueueHarness({ + activeQueueSessionKey, + activeSessionId, + activeSessionIdRef, + onReady, + requestGateway, + routeTokenRef, + selectedStoredSessionIdRef +}: { + activeQueueSessionKey: string + activeSessionId: null | string + activeSessionIdRef: MutableRefObject + onReady: (handle: QueueHarnessHandle) => void + requestGateway: (method: string, params?: Record) => Promise + routeTokenRef: MutableRefObject + selectedStoredSessionIdRef: MutableRefObject +}) { + const stateRef = useRef({ + messages: [], + busy: false, + awaitingResponse: false, + interrupted: false + } as never) + + const draftRef = useRef('') + const queueEditRef = useRef(null) + + const actions = usePromptActions({ + activeSessionId, + activeSessionIdRef, + branchCurrentSession: async () => true, + busyRef: { current: false }, + createBackendSessionForSend: async () => null, + getRoutedStoredSessionId: () => selectedStoredSessionIdRef.current, + getRuntimeIdForStoredSession: storedSessionId => + storedSessionId === selectedStoredSessionIdRef.current + ? activeSessionIdRef.current + : null, + resumeStoredSession: async storedSessionId => { + const routeToken = routeTokenRef.current + const selectedStoredSessionId = selectedStoredSessionIdRef.current + + const resumed = await requestGateway<{ session_id: string }>('session.resume', { + session_id: storedSessionId, + source: 'desktop' + }) + + if ( + routeTokenRef.current !== routeToken || + selectedStoredSessionIdRef.current !== selectedStoredSessionId + ) { + return + } + + selectedStoredSessionIdRef.current = storedSessionId + activeSessionIdRef.current = resumed.session_id + }, + handleSkinCommand: () => '', + openMemoryGraph: () => undefined, + refreshSessions: async () => undefined, + requestGateway, + getRouteToken: () => routeTokenRef.current, + runtimeIdByStoredSessionIdRef: { current: new Map() }, + selectedStoredSessionIdRef, + startFreshSessionDraft: () => undefined, + sttEnabled: false, + updateSessionState: (_sessionId, updater) => { + const next = updater(stateRef.current) as never + stateRef.current = next + + return next + } + }) + + const queue = useComposerQueue({ + activeQueueSessionKey, + attachments: [], + busy: true, + clearDraft: () => undefined, + draftRef, + focusInput: () => undefined, + loadIntoComposer: () => undefined, + onCancel: actions.cancelRun, + onSteer: actions.steerPrompt, + onSubmit: actions.submitText, + queueEditRef, + queueSessionKey: activeQueueSessionKey, + sessionId: activeSessionId + }) + + useEffect(() => { + onReady({ drainNextQueued: queue.drainNextQueued }) + }, [onReady, queue.drainNextQueued]) + + return null +} + describe('usePromptActions /title', () => { beforeEach(() => { setSessions(() => [sessionInfo()]) @@ -2192,6 +2295,40 @@ describe('usePromptActions submit / queue drain semantics', () => { ).toBe(true) }) + it('forwards queued source identity and submission time to prompt.submit', async () => { + const busyRef = { current: true } + const requestGateway = vi.fn(async () => ({}) as never) + + let handle: HarnessHandle | null = null + render( + (handle = h)} + refreshSessions={async () => undefined} + requestGateway={requestGateway} + /> + ) + + const accepted = await handle!.submitText('queued message', { + fromQueue: true, + messageId: 'queued-1700000000000-source', + submittedAt: 1_700_000_000_000 + }) + + expect(accepted).toBe(true) + expect(requestGateway).toHaveBeenCalledWith( + 'prompt.submit', + { + message_id: 'queued-1700000000000-source', + queued: true, + session_id: RUNTIME_SESSION_ID, + submitted_at: 1_700_000_000, + text: 'queued message' + }, + 1_800_000 + ) + }) + it('a rejected fromQueue drain returns false (entry stays queued) and a later retry sends it', async () => { // A stale-session 404 must not strand the queued entry: submitPrompt returns // false on failure so the composer keeps it, and the edge-independent @@ -2300,6 +2437,188 @@ describe('usePromptActions submit / queue drain semantics', () => { }) }) +describe('useComposerQueue source-session retention', () => { + const STORED_SESSION_A = 'stored-queue-a' + const STORED_SESSION_B = 'stored-queue-b' + + beforeEach(() => { + window.localStorage.removeItem('hermes.desktop.composerQueue.v1') + $queuedPromptsBySession.set({}) + }) + + afterEach(() => { + cleanup() + vi.restoreAllMocks() + }) + + it('keeps A queued across a pending switch and retries the same source on A recovery', async () => { + vi.spyOn(Date, 'now').mockReturnValue(1_700_000_000_000) + + const entry = enqueueQueuedPrompt(STORED_SESSION_A, { + attachments: [], + text: 'keep this in session A' + })! + + const selectedStoredSessionIdRef: MutableRefObject = { + current: STORED_SESSION_A + } + + const activeSessionIdRef: MutableRefObject = { current: null } + const routeTokenRef: MutableRefObject = { current: 'route-a' } + const calls: { method: string; params?: Record }[] = [] + const promptCalls: Record[] = [] + const ownedSourceIds = new Set() + let canonicalRows = 0 + let releaseInitialResume: () => void = () => undefined + let markInitialResumeStarted: () => void = () => undefined + + const initialResumeStarted = new Promise(resolve => { + markInitialResumeStarted = resolve + }) + + let firstResume = true + + const requestGateway = vi.fn(async (method: string, params?: Record) => { + calls.push({ method, params }) + + if (method === 'session.resume') { + if (firstResume) { + firstResume = false + markInitialResumeStarted() + await new Promise(resolve => { + releaseInitialResume = resolve + }) + + return { session_id: 'rt-a-abandoned' } as never + } + + return { session_id: 'rt-a-recovered' } as never + } + + if (method === 'prompt.submit') { + const payload = params ?? {} + const sourceId = String(payload.message_id) + promptCalls.push(payload) + + // The runtime minted by the abandoned mid-switch resume is dead on + // the backend: the retry recovers through the #91276 cache, probes + // it once, and falls through to a fresh resume on the 404. + if (payload.session_id === 'rt-a-abandoned') { + throw new Error('session not found') + } + + if (!ownedSourceIds.has(sourceId)) { + ownedSourceIds.add(sourceId) + canonicalRows += 1 + throw new Error('request timed out: prompt.submit') + } + + return { status: 'duplicate' } as never + } + + return {} as never + }) + + let handle: QueueHarnessHandle | null = null + + const view = render( + (handle = next)} + requestGateway={requestGateway} + routeTokenRef={routeTokenRef} + selectedStoredSessionIdRef={selectedStoredSessionIdRef} + /> + ) + + await waitFor(() => expect(handle).not.toBeNull()) + + let firstAttempt!: Promise + act(() => { + firstAttempt = handle!.drainNextQueued() + }) + await initialResumeStarted + + selectedStoredSessionIdRef.current = STORED_SESSION_B + activeSessionIdRef.current = 'rt-b' + routeTokenRef.current = 'route-b' + view.rerender( + (handle = next)} + requestGateway={requestGateway} + routeTokenRef={routeTokenRef} + selectedStoredSessionIdRef={selectedStoredSessionIdRef} + /> + ) + + let firstAccepted = true + await act(async () => { + releaseInitialResume() + firstAccepted = await firstAttempt + }) + + expect(firstAccepted).toBe(false) + expect(getQueuedPrompts(STORED_SESSION_A)).toEqual([entry]) + expect(getQueuedPrompts(STORED_SESSION_B)).toEqual([]) + expect(promptCalls).toEqual([]) + + selectedStoredSessionIdRef.current = STORED_SESSION_A + activeSessionIdRef.current = 'rt-a-stale' + routeTokenRef.current = 'route-a' + view.rerender( + (handle = next)} + requestGateway={requestGateway} + routeTokenRef={routeTokenRef} + selectedStoredSessionIdRef={selectedStoredSessionIdRef} + /> + ) + + let retryAccepted = false + await act(async () => { + retryAccepted = await handle!.drainNextQueued() + }) + + expect(retryAccepted).toBe(true) + expect(promptCalls).toEqual([ + { + message_id: entry.id, + queued: true, + session_id: 'rt-a-stale', + submitted_at: entry.queuedAt / 1000, + text: entry.text + }, + { + message_id: entry.id, + queued: true, + session_id: 'rt-a-abandoned', + submitted_at: entry.queuedAt / 1000, + text: entry.text + }, + { + message_id: entry.id, + queued: true, + session_id: 'rt-a-recovered', + submitted_at: entry.queuedAt / 1000, + text: entry.text + } + ]) + expect(canonicalRows).toBe(1) + expect(ownedSourceIds).toEqual(new Set([entry.id])) + expect(getQueuedPrompts(STORED_SESSION_A)).toEqual([]) + expect(getQueuedPrompts(STORED_SESSION_B)).toEqual([]) + expect(calls.filter(call => call.method === 'session.resume')).toHaveLength(2) + }) +}) + describe('usePromptActions redirectPrompt', () => { afterEach(() => { cleanup() diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts index fc6115650b8e..b3bb32c5df40 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts @@ -753,7 +753,12 @@ export function useSubmitPrompt(deps: SubmitPromptDeps) { rewriteOptimistic(liveSessionId) const text = buildContextText(syncedAttachments) + const sourceMetadata = { + ...(options?.messageId ? { message_id: options.messageId } : {}), + ...(options?.submittedAt !== undefined ? { submitted_at: options.submittedAt / 1000 } : {}) + } const submitParams = (targetId: string) => ({ + ...sourceMetadata, session_id: targetId, text, ...(interrupted && { interrupted }), diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/utils.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/utils.ts index c5bfb83da70f..b043dcd78a47 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/utils.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/utils.ts @@ -708,6 +708,7 @@ export interface SubmitTextOptions { * still receives the text as a normal user turn. */ displayKind?: 'hidden' fromQueue?: boolean + messageId?: string /** Runtime session id to submit into. Queue drains pass this so a * backgrounded/source session cannot be replaced by the current foreground * session between enqueue and drain. */ @@ -715,4 +716,5 @@ export interface SubmitTextOptions { /** Stable stored session id for optimistic/cache updates and stale-runtime * recovery. Distinct from the runtime session id minted by the gateway. */ storedSessionId?: string | null + submittedAt?: number } diff --git a/contributors/emails/yingliangzhang@users.noreply.github.com b/contributors/emails/yingliangzhang@users.noreply.github.com new file mode 100644 index 000000000000..7dedf20d9438 --- /dev/null +++ b/contributors/emails/yingliangzhang@users.noreply.github.com @@ -0,0 +1 @@ +yingliang-zhang diff --git a/cron/scheduler.py b/cron/scheduler.py index 220e92f4393a..b242cb0443ff 100644 --- a/cron/scheduler.py +++ b/cron/scheduler.py @@ -1901,10 +1901,11 @@ def _maybe_mirror_cron_delivery( # The brief is not the agent speaking; an assistant-role mirror lands as # assistant→assistant after the agent's last turn and breaks strict # alternation (issue #2221, the exact failure #2313 removed). A - # user-role turn collapses safely via repair_message_sequence's - # consecutive-user merge on every provider, and the prefix preserves the - # "this came from cron" context that the dropped SQLite mirror metadata - # would otherwise lose on replay. + # user-role turn merges safely on the per-request API copy through + # ``_drop_thinking_only_and_merge_users`` for strict providers, while + # remaining a distinct canonical source message. The prefix preserves + # the "this came from cron" context that the dropped SQLite mirror + # metadata would otherwise lose on replay. ok = mirror_to_session( platform_name, str(chat_id), diff --git a/gateway/mirror.py b/gateway/mirror.py index 2b086ce4f04d..656535229792 100644 --- a/gateway/mirror.py +++ b/gateway/mirror.py @@ -57,8 +57,8 @@ def mirror_to_session( at the SQLite boundary (only role+content persist), so on replay an assistant-role mirror is indistinguishable from a real assistant turn and produces ``assistant → assistant`` pairs that break strict-alternation - providers (issue #2221). A user-role mirror collapses safely via - ``repair_message_sequence``'s consecutive-user merge on every provider. + providers (issue #2221). A user-role mirror remains distinct in canonical + history and merges only on the per-request API copy for strict providers. Returns True if mirrored successfully, False if no matching session or error. All errors are caught -- this is never fatal. diff --git a/gateway/session.py b/gateway/session.py index d7770188613a..8fe846886351 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -4095,10 +4095,10 @@ def load_transcript(self, session_id: str) -> List[Dict[str, Any]]: except Exception: pass try: - # repair_alternation: this load feeds LIVE REPLAY. A durable - # user;user wedge (e.g. a turn that persisted no assistant row) - # would otherwise re-trigger the pre-request repair on every - # request forever — heal it once at the restore boundary. + # repair_alternation: this load feeds LIVE REPLAY. Repair malformed + # assistant/tool structure in the restored copy while preserving + # adjacent user rows as canonical source boundaries; provider-wire + # normalization merges them later on a per-request copy. return self._db.get_messages_as_conversation( session_id, repair_alternation=True ) diff --git a/hermes_state.py b/hermes_state.py index 8fed65b943a3..9b5ed22ccd9e 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -11572,10 +11572,12 @@ def _insert_message_rows(self, conn, session_id: str, messages: List[Dict[str, A except (json.JSONDecodeError, TypeError): tool_calls = [] tool_calls_json = json.dumps(tool_calls) if tool_calls else None - # Accept either `platform_message_id` (new explicit name) or - # `message_id` (yuanbao's existing convention on message dicts). + # Match append persistence precedence for external and internal + # canonical source identity forms. platform_msg_id = ( - msg.get("platform_message_id") or msg.get("message_id") + msg.get("platform_message_id") + or msg.get("message_id") + or msg.get("_source_message_id") ) api_content = msg.get("api_content") @@ -12312,12 +12314,12 @@ def get_messages_as_conversation( as well. See :meth:`rewind_to_message`. ``repair_alternation=True`` runs ``repair_message_sequence`` over the - loaded list before returning it. Callers that restore a session for - LIVE REPLAY should pass it: a durable alternation violation (e.g. a - ``user;user`` pair left by a turn that persisted no assistant row) - otherwise re-triggers the pre-request defensive repair on every - single request for the rest of the session's life — the repair - mutates only the per-request list, never the stored transcript. + loaded live-replay copy before returning it. This repairs malformed + assistant/tool structure such as split assistant turns and orphaned + tool results without rewriting the durable transcript. Adjacent + ``user`` messages remain distinct canonical source turns; the + per-request provider copy later merges them via + ``drop_thinking_only_and_merge_users`` for strict role alternation. Inspection/export consumers keep the default and see the transcript verbatim. """ @@ -12545,9 +12547,9 @@ def _rows_to_conversation( repaired = repair_message_sequence(None, messages) if repaired: logger.info( - "Repaired %d message-alternation violation(s) while " - "restoring session %s — durable transcript kept them, " - "see repair_message_sequence", + "Repaired %d malformed assistant/tool sequence violation(s) " + "while restoring session %s — durable transcript retained " + "its original rows; see repair_message_sequence", repaired, session_id, ) diff --git a/run_agent.py b/run_agent.py index b710a2263bcc..c5227a04412c 100644 --- a/run_agent.py +++ b/run_agent.py @@ -1972,8 +1972,9 @@ def _apply_persist_user_message_override(self, messages: List[Dict]) -> None: that synthetic text leak into persisted transcripts or resumed session history. When an override is configured for the active turn, mutate the in-memory messages list in place so both persistence and returned - history stay clean. A paired timestamp override preserves the platform - event time as message metadata, rather than embedding it in content. + history stay clean. Paired source metadata preserves platform event + identity/time on the canonical message rather than embedding either in + content; the common API-copy boundary strips it before provider calls. """ idx = getattr(self, "_persist_user_message_idx", None) override = getattr(self, "_persist_user_message_override", None) @@ -2161,8 +2162,9 @@ def _flush_messages_to_session_db_unlocked( # silently dropped it. Instead, resolve the override here and apply it # ONLY to the value written to the DB (see the write loop below); the # live dict is never mutated, so every caller (early persist, mid-loop - # flush, /resume, /branch) is protected uniformly. Timestamp override is - # metadata and is likewise applied only to the written row. + # flush, /resume, /branch) is protected uniformly. The timestamp may + # already exist on the canonical message; this override also guarantees + # the written row retains it when a call path supplies it separately. _ov_idx = getattr(self, "_persist_user_message_idx", None) _ov_content = getattr(self, "_persist_user_message_override", None) _ov_timestamp = getattr(self, "_persist_user_message_timestamp", None) @@ -2173,9 +2175,10 @@ def _flush_messages_to_session_db_unlocked( # Positional flushing used to slice at # max(len(conversation_history), _last_flushed_db_idx). That # assumes the live `messages` list is the original history plus a - # new tail. repair_message_sequence can shrink/merge the history - # copy before the final flush, making len(conversation_history) - # larger than len(messages); the slice is then empty and delivered + # new tail. repair_message_sequence can shrink the history copy by + # merging assistant turns or dropping orphaned tools before the + # final flush, making len(conversation_history) larger than + # len(messages); the slice is then empty and delivered # assistant responses never reach state.db (#46053). # # Track persistence with an intrinsic per-message marker rather than @@ -2363,10 +2366,15 @@ def _flush_messages_to_session_db_unlocked( "codex_reasoning_items": msg.get("codex_reasoning_items"), "codex_message_items": msg.get("codex_message_items"), "_compressed_summary": bool(msg.get(COMPRESSED_SUMMARY_METADATA_KEY)), + "platform_message_id": ( + msg.get("platform_message_id") + or msg.get("message_id") + or msg.get("_source_message_id") + ), "timestamp": _row_timestamp, "api_content": _row_api_content, # Standalone reference handoffs are always hidden, even - # when the summarized transcript contained a user turn — + # when the summarized transcript contained a user turn— # otherwise they occupy the active user slot in # retry/undo/session dispatch (#80622). Merge-into-tail # carriers keep prior visibility rules so preserved tail @@ -8595,6 +8603,7 @@ def run_conversation( persist_user_display_kind: Optional[str] = None, persist_user_display_metadata: Optional[Dict[str, Any]] = None, moa_config: Optional[dict[str, Any]] = None, + persist_user_message_id: Optional[str] = None, ) -> Dict[str, Any]: """Forwarder — see ``agent.conversation_loop.run_conversation``.""" # A review deliberately shares this agent's session_id for prompt-cache @@ -8968,6 +8977,7 @@ def _interrupt_turn(message: str) -> None: persist_user_timestamp=persist_user_timestamp, persist_user_display_kind=persist_user_display_kind, persist_user_display_metadata=persist_user_display_metadata, + persist_user_message_id=persist_user_message_id, moa_config=moa_config, ) finally: diff --git a/tests/agent/transports/test_chat_completions.py b/tests/agent/transports/test_chat_completions.py index 059322ff411e..047e5812e80b 100644 --- a/tests/agent/transports/test_chat_completions.py +++ b/tests/agent/transports/test_chat_completions.py @@ -140,6 +140,25 @@ def test_convert_messages_strips_timestamp(self, transport): # Original list untouched (deepcopy-on-demand) assert msgs[0]["timestamp"] == 1781976577.0 + def test_convert_messages_strips_source_message_ids(self, transport): + """Persisted platform identities are transcript metadata, not API fields.""" + msgs = [ + { + "role": "user", + "content": "hi", + "message_id": "desktop-queued-1", + "platform_message_id": "legacy-platform-1", + }, + ] + + result = transport.convert_messages(msgs) + + assert "message_id" not in result[0] + assert "platform_message_id" not in result[0] + assert result[0]["content"] == "hi" + assert msgs[0]["message_id"] == "desktop-queued-1" + assert msgs[0]["platform_message_id"] == "legacy-platform-1" + def test_convert_messages_no_copy_without_timestamp(self, transport): """A timestamp-free message list needs no sanitize pass and is returned by identity (preserves the deepcopy-on-demand contract).""" diff --git a/tests/gateway/test_35809_auto_reset_clean_context.py b/tests/gateway/test_35809_auto_reset_clean_context.py index 4f6d4149f926..6958a2956375 100644 --- a/tests/gateway/test_35809_auto_reset_clean_context.py +++ b/tests/gateway/test_35809_auto_reset_clean_context.py @@ -153,8 +153,8 @@ def _bloat(n): # Stand-in for the oversized, post-compression "child" transcript that # could not be compressed any further (#35809). Alternates roles so the # fixture is a valid conversation: load_transcript is a live-replay - # restore site and heals alternation violations on load (#64934), so a - # degenerate all-user transcript would be merged into one message. + # restore site and repairs malformed assistant/tool structure on load + # (#64934); alternate roles here to model a normal replayable conversation. return [ { "role": "user" if i % 2 == 0 else "assistant", diff --git a/tests/hermes_state/test_restore_alternation_repair.py b/tests/hermes_state/test_restore_alternation_repair.py index f665aae2c84c..9476efecb901 100644 --- a/tests/hermes_state/test_restore_alternation_repair.py +++ b/tests/hermes_state/test_restore_alternation_repair.py @@ -1,14 +1,14 @@ -"""get_messages_as_conversation(repair_alternation=True) — heal durable -alternation violations at the restore boundary. +"""Live restore repairs malformed assistant/tool structure without erasing +canonical user source boundaries. -A turn that persists a user row but no assistant row (e.g. its reply was -suppressed, or two concurrent turns interleaved their flushes) leaves a -``user;user`` pair in state.db. Without repair at restore, the defensive -pre-request ``repair_message_sequence`` re-fires on EVERY request for the -rest of the session's life, because it mutates only the per-request list. +Adjacent persisted ``user;user`` rows are distinct source turns. Live restore +keeps them separate, while the transient provider copy merges them with an +explicit boundary marker to satisfy strict role alternation. The +``repair_alternation=True`` path remains responsible for malformed assistant +and tool structure. -Default (``repair_alternation=False``) must stay verbatim: inspection and -export consumers (trace upload, context guard) read the transcript as-is. +Default (``repair_alternation=False``) stays verbatim for inspection and +export consumers such as trace upload and the context guard. """ import pytest @@ -24,8 +24,8 @@ def db(tmp_path): session_db.close() -def _seed_wedged_session(db, session_id="s1"): - """assistant → user → user (no assistant row between): the durable wedge.""" +def _seed_adjacent_user_session(db, session_id="s1"): + """Persist two adjacent user source turns in otherwise clean history.""" db.create_session(session_id, "system prompt") db.append_message(session_id=session_id, role="user", content="first ask") db.append_message(session_id=session_id, role="assistant", content="first reply") @@ -34,27 +34,70 @@ def _seed_wedged_session(db, session_id="s1"): db.append_message(session_id=session_id, role="assistant", content="next reply") - - -def test_repair_alternation_merges_user_pair(db): - _seed_wedged_session(db) - messages = db.get_messages_as_conversation("s1", repair_alternation=True) +def test_default_load_is_verbatim(db): + _seed_adjacent_user_session(db) + messages = db.get_messages_as_conversation("s1") roles = [m["role"] for m in messages] - assert roles == ["user", "assistant", "user", "assistant"] - # Both user texts survive, merged in order — no user input is lost. - merged = messages[2]["content"] - assert "unanswered turn" in merged and "next turn" in merged - assert merged.index("unanswered turn") < merged.index("next turn") + assert roles == ["user", "assistant", "user", "user", "assistant"] -def test_repaired_load_is_stable_under_prerequest_repair(db): - """The restored list must yield ZERO further repairs — this is the whole - point: the pre-request defensive repair stops firing every turn.""" +def test_repair_alternation_preserves_user_pair_until_provider_wire(db): + from agent.agent_runtime_helpers import drop_thinking_only_and_merge_users + + _seed_adjacent_user_session(db) + messages = db.get_messages_as_conversation("s1", repair_alternation=True) + canonical = [dict(message) for message in messages] + + assert [message["role"] for message in messages] == [ + "user", "assistant", "user", "user", "assistant" + ] + assert [messages[2]["content"], messages[3]["content"]] == [ + "unanswered turn", "next turn" + ] + assert messages[2] is not messages[3] + + provider_messages = drop_thinking_only_and_merge_users( + [dict(message) for message in messages] + ) + + assert [message["role"] for message in provider_messages] == [ + "user", "assistant", "user", "assistant" + ] + assert provider_messages[2]["content"] == ( + "unanswered turn\n\n[Next user message]\n\nnext turn" + ) + assert messages == canonical + + +def test_adjacent_user_load_is_stable_under_canonical_repair(db): + """Repeated canonical repair leaves adjacent source turns untouched.""" from agent.agent_runtime_helpers import repair_message_sequence - _seed_wedged_session(db) + _seed_adjacent_user_session(db) messages = db.get_messages_as_conversation("s1", repair_alternation=True) + canonical = [dict(message) for message in messages] + assert repair_message_sequence(None, messages) == 0 + assert messages == canonical + + +def test_repair_alternation_repairs_malformed_assistant_pair(db): + from agent.agent_runtime_helpers import repair_message_sequence + + db.create_session("s3", "system prompt") + db.append_message(session_id="s3", role="user", content="ask") + db.append_message(session_id="s3", role="assistant", content="first fragment") + db.append_message(session_id="s3", role="assistant", content="second fragment") + + verbatim = db.get_messages_as_conversation("s3") + repaired = db.get_messages_as_conversation("s3", repair_alternation=True) + + assert [message["role"] for message in verbatim] == [ + "user", "assistant", "assistant" + ] + assert [message["role"] for message in repaired] == ["user", "assistant"] + assert repaired[1]["content"] == "first fragment\nsecond fragment" + assert repair_message_sequence(None, repaired) == 0 @@ -92,10 +135,13 @@ class _StubAgent: assert state is not None roles = [m["role"] for m in state.history] - # No consecutive user turns — the durable user;user wedge was healed. - assert roles == ["user", "assistant", "user", "assistant"], roles + # repair_message_sequence preserves adjacent user messages as distinct + # source turns — provider role alternation is repaired later on the + # per-request api_messages copy by drop_thinking_only_and_merge_users. + assert roles == ["user", "assistant", "user", "user", "assistant"], roles + # No consecutive assistant turns — that IS repaired. for a, b in zip(roles, roles[1:]): - assert not (a == "user" and b == "user"), "unhealed user;user in ACP live replay" - # No user input lost — both user texts survive, merged in order. - merged = state.history[2]["content"] - assert "unanswered turn" in merged and "next turn" in merged + assert not (a == "assistant" and b == "assistant"), "unhealed assistant;assistant in ACP live replay" + # No user input lost — both user texts survive as separate turns. + assert state.history[2]["content"] == "unanswered turn" + assert state.history[3]["content"] == "next turn" diff --git a/tests/run_agent/test_identity_flush.py b/tests/run_agent/test_identity_flush.py index 8cb1e54d6d56..2c6e6e90dbb5 100644 --- a/tests/run_agent/test_identity_flush.py +++ b/tests/run_agent/test_identity_flush.py @@ -90,7 +90,10 @@ def test_repair_shrunk_messages_below_history_length_still_persists_assistant(se agent = _make_agent(db) # Simulate history already loaded from state.db. - history = [{"role": "user", "content": f"u{i}"} for i in range(6)] + history = [ + {"role": "assistant", "content": f"a{i}"} + for i in range(6) + ] for msg in history: db.append_message( session_id=SESSION_ID, @@ -98,10 +101,11 @@ def test_repair_shrunk_messages_below_history_length_still_persists_assistant(se content=msg["content"], ) - # repair_message_sequence merged the six history rows into one - # dict before this turn appended the new user/assistant pair. + # Simulate canonical repair compacting the six adjacent + # assistant rows before this turn appends a new user/assistant + # pair. messages = [ - {"role": "user", "content": "\n\n".join(f"u{i}" for i in range(6))}, + {"role": "assistant", "content": "\n".join(f"a{i}" for i in range(6))}, {"role": "user", "content": "new question"}, {"role": "assistant", "content": "new answer"}, ] diff --git a/tests/run_agent/test_message_sequence_repair.py b/tests/run_agent/test_message_sequence_repair.py index 054c75656c09..724be0caa5af 100644 --- a/tests/run_agent/test_message_sequence_repair.py +++ b/tests/run_agent/test_message_sequence_repair.py @@ -42,31 +42,79 @@ def test_drop_scaffolding_rewinds_orphan_tool_tail(): # ── _repair_message_sequence ─────────────────────────────────────────────── -def test_repair_merges_consecutive_user_messages(): +def test_canonical_repair_preserves_adjacent_user_source_boundaries(): + """Canonical repair must not collapse separately sourced user turns.""" + agent = _bare_agent() + first = { + "role": "user", + "content": "interrupted turn tail", + "timestamp": 101.25, + "_source_message_id": "desktop-1", + } + second = { + "role": "user", + "content": "first queued prompt", + "timestamp": 102.5, + "_source_message_id": "desktop-2", + } + messages = [first, second] + original = [dict(message) for message in messages] + + repairs = AIAgent._repair_message_sequence(agent, messages) + + assert repairs == 0 + assert messages == original + assert messages[0] is first + assert messages[1] is second + wire_messages = [ + {"role": message["role"], "content": message["content"]} + for message in messages + ] + wire_messages = AIAgent._drop_thinking_only_and_merge_users(wire_messages) + + assert wire_messages == [ + { + "role": "user", + "content": ( + "interrupted turn tail\n\n" + "[Next user message]\n\n" + "first queued prompt" + ), + } + ] + assert messages == original + assert messages[0] is first + assert messages[1] is second + + +def test_repair_preserves_consecutive_plain_text_users(): agent = _bare_agent() messages = [ {"role": "user", "content": "first"}, {"role": "user", "content": "second"}, ] + original = [dict(message) for message in messages] repairs = AIAgent._repair_message_sequence(agent, messages) - assert repairs == 1 - assert len(messages) == 1 - assert messages[0]["role"] == "user" - assert messages[0]["content"] == "first\n\nsecond" + assert repairs == 0 + assert messages == original -def test_repair_preserves_user_content_when_one_side_empty(): +def test_repair_preserves_empty_user_source_boundary(): agent = _bare_agent() messages = [ {"role": "user", "content": ""}, {"role": "user", "content": "real message"}, ] - AIAgent._repair_message_sequence(agent, messages) + repairs = AIAgent._repair_message_sequence(agent, messages) - assert messages == [{"role": "user", "content": "real message"}] + assert repairs == 0 + assert messages == [ + {"role": "user", "content": ""}, + {"role": "user", "content": "real message"}, + ] def test_repair_does_not_rewind_ongoing_dialog_tool_pair(): @@ -164,6 +212,22 @@ def test_repair_keeps_tool_matching_only_call_id(): +def test_repair_preserves_multimodal_user_content(): + """Canonical repair preserves multimodal user boundaries unchanged.""" + agent = _bare_agent() + messages = [ + {"role": "user", "content": [{"type": "text", "text": "hi"}, + {"type": "image_url", "image_url": {"url": "..."}}]}, + {"role": "user", "content": "follow-up"}, + ] + + AIAgent._repair_message_sequence(agent, messages) + + # The multimodal user message stays distinct and retains its attachment. + assert len(messages) == 2 + assert isinstance(messages[0]["content"], list) + + @@ -339,8 +403,8 @@ def test_cursor_clamped_when_compaction_shrinks_below_cursor(): turn-end flush doesn't skip the assistant/tool chain (#44837).""" agent = _bare_agent() messages = [ - {"role": "user", "content": "first"}, - {"role": "user", "content": "second"}, + {"role": "assistant", "content": "first"}, + {"role": "assistant", "content": "second"}, ] agent._last_flushed_db_idx = 2 # both rows already flushed @@ -356,24 +420,51 @@ def test_cursor_rewinds_when_compaction_happens_before_cursor(): rewind it by the number removed, or unflushed rows get skipped. A plain min() clamp does NOT catch this case.""" agent = _bare_agent() - flushed_a = {"role": "user", "content": "first"} - flushed_b = {"role": "user", "content": "second"} # merged into flushed_a - unflushed_assistant = {"role": "assistant", "content": "answer"} - messages = [flushed_a, flushed_b, unflushed_assistant] - agent._last_flushed_db_idx = 2 # the two user rows are flushed + flushed_a = {"role": "assistant", "content": "first"} + flushed_b = {"role": "assistant", "content": "second"} + unflushed_user = {"role": "user", "content": "question"} + messages = [flushed_a, flushed_b, unflushed_user] + agent._last_flushed_db_idx = 2 # the two assistant rows are flushed repairs = repair_message_sequence_with_cursor(agent, messages) assert repairs == 1 assert len(messages) == 2 - # Cursor must now point at the assistant (index 1), not stay at 2 — + # Cursor must now point at the user (index 1), not stay at 2 — # min(2, len=2) would leave it at 2 and the flush would skip it. assert agent._last_flushed_db_idx == 1 - assert messages[agent._last_flushed_db_idx] is unflushed_assistant + assert messages[agent._last_flushed_db_idx] is unflushed_user + + + + +def test_cursor_untouched_when_no_repairs(): + agent = _bare_agent() + messages = [ + {"role": "user", "content": "first"}, + {"role": "user", "content": "second"}, + ] + agent._last_flushed_db_idx = 1 + + repairs = repair_message_sequence_with_cursor(agent, messages) + + assert repairs == 0 + assert agent._last_flushed_db_idx == 1 + assert len(messages) == 2 +def test_cursor_helper_safe_without_cursor_attribute(): + """Bare agents (no _last_flushed_db_idx) must not crash.""" + agent = _bare_agent() + messages = [ + {"role": "assistant", "content": "a"}, + {"role": "assistant", "content": "b"}, + ] + repairs = repair_message_sequence_with_cursor(agent, messages) + assert repairs == 1 + assert not hasattr(agent, "_last_flushed_db_idx") def test_flush_guard_clamps_overshooting_cursor(): diff --git a/tests/run_agent/test_run_agent.py b/tests/run_agent/test_run_agent.py index 1e7b7d0f5ab9..4e71c1bb35f5 100644 --- a/tests/run_agent/test_run_agent.py +++ b/tests/run_agent/test_run_agent.py @@ -2784,7 +2784,14 @@ def test_summary_strips_strict_schema_foreign_fields(self, agent): agent.client.chat.completions.create.return_value = _mock_response(content="Summary") agent._cached_system_prompt = "You are helpful." messages = [ - {"role": "user", "content": "do stuff"}, + { + "role": "user", + "content": "do stuff", + "timestamp": 101.25, + "message_id": "message-1", + "platform_message_id": "platform-1", + "_source_message_id": "source-1", + }, { "role": "assistant", "tool_calls": [{"id": "call_1", "function": {"name": "execute_code", "arguments": "{}"}}], @@ -2798,12 +2805,25 @@ def test_summary_strips_strict_schema_foreign_fields(self, agent): assert result == "Summary" sent_msgs = agent.client.chat.completions.create.call_args.kwargs.get("messages", []) + source_metadata = ( + "timestamp", + "message_id", + "platform_message_id", + "_source_message_id", + ) for m in sent_msgs: assert "tool_name" not in m, m assert "codex_reasoning_items" not in m, m assert "codex_message_items" not in m, m + assert all(key not in m for key in source_metadata), m assert not any(isinstance(k, str) and k.startswith("_") for k in m), m - # Internal history is untouched — the path copies each message. + # Canonical history is untouched — the path copies each message. + assert {key: messages[0][key] for key in source_metadata} == { + "timestamp": 101.25, + "message_id": "message-1", + "platform_message_id": "platform-1", + "_source_message_id": "source-1", + } assert messages[2]["tool_name"] == "execute_code" assert messages[1]["codex_reasoning_items"] == [{"id": "rs_1"}] diff --git a/tests/run_agent/test_thinking_only_sanitizer.py b/tests/run_agent/test_thinking_only_sanitizer.py index f6190fab722a..f4cf0e40b33e 100644 --- a/tests/run_agent/test_thinking_only_sanitizer.py +++ b/tests/run_agent/test_thinking_only_sanitizer.py @@ -11,6 +11,8 @@ backstory on why the alternative — fabricating "." stub text — was rejected. """ +from copy import deepcopy + from run_agent import AIAgent @@ -100,10 +102,48 @@ def test_adjacent_users_merge_even_when_no_thinking_row_was_dropped(self): out = AIAgent._drop_thinking_only_and_merge_users(msgs) - assert out == [{"role": "user", "content": "SUMMARY SCAFFOLD\n\nREAL ASK"}] + assert out == [{"role": "user", "content": "SUMMARY SCAFFOLD\n\n[Next user message]\n\nREAL ASK"}] assert scaffold["content"] == "SUMMARY SCAFFOLD" assert live_ask["content"] == "REAL ASK" + def test_adjacent_users_always_keep_explicit_wire_boundary(self): + boundary = {"type": "text", "text": "[Next user message]"} + text = lambda value: {"type": "text", "text": value} + cases = [ + ("", "real", "[Next user message]\n\nreal"), + ("real", "", "real\n\n[Next user message]"), + ("", "", "[Next user message]"), + ("", [text("right")], [boundary, text("right")]), + ("left", [], [text("left"), boundary]), + ([], "right", [boundary, text("right")]), + ([text("left")], "", [text("left"), boundary]), + ([text("left")], "right", [text("left"), boundary, text("right")]), + ] + + for previous, current, expected in cases: + messages = [ + {"role": "user", "content": previous}, + {"role": "user", "content": current}, + ] + canonical = deepcopy(messages) + + out = AIAgent._drop_thinking_only_and_merge_users(messages) + + assert out == [{"role": "user", "content": expected}] + assert messages == canonical + + def test_drops_thinking_only_between_user_messages_and_merges(self): + msgs = [ + {"role": "user", "content": "help me with X"}, + {"role": "assistant", "content": "", "reasoning": "let me think"}, + {"role": "user", "content": "ok continue"}, + ] + out = AIAgent._drop_thinking_only_and_merge_users(msgs) + assert len(out) == 1 + assert out[0]["role"] == "user" + assert out[0]["content"] == ( + "help me with X\n\n[Next user message]\n\nok continue" + ) def test_preserves_alternation_after_drop(self): msgs = [ @@ -115,10 +155,20 @@ def test_preserves_alternation_after_drop(self): out = AIAgent._drop_thinking_only_and_merge_users(msgs) roles = [m["role"] for m in out] assert roles == ["user", "assistant"] - assert out[0]["content"] == "u1\n\nu2" + assert out[0]["content"] == "u1\n\n[Next user message]\n\nu2" assert out[1]["content"] == "real reply" + def test_multiple_thinking_only_in_sequence_collapses(self): + msgs = [ + {"role": "user", "content": "u1"}, + {"role": "assistant", "content": "", "reasoning": "r1"}, + {"role": "assistant", "content": "", "reasoning": "r2"}, + {"role": "user", "content": "u2"}, + ] + out = AIAgent._drop_thinking_only_and_merge_users(msgs) + assert len(out) == 1 + assert out[0]["content"] == "u1\n\n[Next user message]\n\nu2" def test_does_not_touch_stored_messages_original_list_unmutated(self): original_first_user = {"role": "user", "content": "u1"} @@ -157,6 +207,7 @@ def test_merge_concatenates_list_content_user_messages(self): assert len(out) == 1 assert out[0]["content"] == [ {"type": "text", "text": "first"}, + {"type": "text", "text": "[Next user message]"}, {"type": "text", "text": "second"}, ] @@ -188,7 +239,7 @@ def test_system_messages_ignored_by_pass(self): assert len(out) == 2 assert out[0]["role"] == "system" assert out[1]["role"] == "user" - assert out[1]["content"] == "u1\n\nu2" + assert out[1]["content"] == "u1\n\n[Next user message]\n\nu2" # --------------------------------------------------------------------------- diff --git a/tests/test_hermes_state.py b/tests/test_hermes_state.py index 28480c1d1548..bd01a232f384 100644 --- a/tests/test_hermes_state.py +++ b/tests/test_hermes_state.py @@ -608,6 +608,370 @@ def test_startup_heals_null_active_rows(self, tmp_path): + def test_dict_content_round_trip(self, db): + """Dict-shaped content (e.g. provider wrappers) also round-trips.""" + db.create_session(session_id="s1", source="cli") + content = {"parts": [{"text": "hi"}]} + + db.append_message("s1", role="user", content=content) + msgs = db.get_messages("s1") + assert msgs[0]["content"] == content + + def test_string_content_unchanged_by_encoding(self, db): + """Plain strings must not be wrapped — FTS search and legacy + consumers depend on raw-string storage for text content. + """ + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="plain text") + + # Peek at the raw column to confirm no encoding was applied + with db._lock: + row = db._conn.execute( + "SELECT content FROM messages WHERE session_id = ?", ("s1",) + ).fetchone() + assert row["content"] == "plain text" + + def test_replace_messages_persists_tool_name(self, db): + """`replace_messages` (used by /retry, /undo, /compress) must write + tool_name to the DB for messages built by make_tool_result_message.""" + from agent.tool_dispatch_helpers import make_tool_result_message + db.create_session(session_id="s1", source="cli") + db.replace_messages( + "s1", + [ + {"role": "user", "content": "do something"}, + make_tool_result_message("web_search", "some results", "c1"), + ], + ) + + msgs = db.get_messages("s1") + tool_msg = next(m for m in msgs if m["role"] == "tool") + assert tool_msg["tool_name"] == "web_search" + + def test_tool_effect_disposition_round_trips_through_session_db(self, db): + from agent.tool_dispatch_helpers import make_tool_result_message + + db.create_session(session_id="s1", source="cli") + db.replace_messages( + "s1", + [make_tool_result_message( + "write_file", "worker detached", "c1", effect_disposition="unknown" + )], + ) + + assert db.get_messages_as_conversation("s1")[0]["effect_disposition"] == "unknown" + + def test_replace_messages_handles_multimodal_content(self, db): + """`replace_messages` (used by /retry, /undo, /compress) must also + handle list content without crashing.""" + db.create_session(session_id="s1", source="cli") + content = [ + {"type": "text", "text": "look at this"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAA"}}, + ] + + db.replace_messages( + "s1", + [ + {"role": "user", "content": content}, + {"role": "assistant", "content": "I see a screenshot."}, + ], + ) + + msgs = db.get_messages("s1") + assert len(msgs) == 2 + assert msgs[0]["content"] == content + assert msgs[1]["content"] == "I see a screenshot." + + def test_get_messages_as_conversation(self, db): + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="Hello") + db.append_message("s1", role="assistant", content="Hi!") + + conv = db.get_messages_as_conversation("s1") + assert len(conv) == 2 + assert conv[0]["role"] == "user" + assert conv[0]["content"] == "Hello" + assert isinstance(conv[0]["timestamp"], float) + assert conv[1]["role"] == "assistant" + assert conv[1]["content"] == "Hi!" + assert isinstance(conv[1]["timestamp"], float) + + def test_get_messages_as_conversation_orders_by_id_not_timestamp(self, db): + """Replay must follow AUTOINCREMENT id (insertion order), never the + wall-clock timestamp. + + ``append_message`` stamps each row with ``time.time()``, which is not + monotonic — on WSL2, after an NTP step, or when a VM/laptop resumes + from sleep the clock can jump backwards mid-conversation. A later + row then carries an *earlier* timestamp than the row before it. If + ``get_messages_as_conversation`` ordered by ``timestamp`` it would + sort an assistant ``tool_calls`` row after its ``tool`` response, + orphaning the tool call and triggering an HTTP 400 on the next + completion. Ordering by ``id`` keeps the real insertion order + regardless of clock skew. See c03acca50. + """ + db.create_session(session_id="s1", source="cli") + + # Simulate a clock regression across a single tool round-trip: the + # assistant tool_calls row is inserted first but stamped LATER than + # the tool response that follows it. + tool_calls = [ + {"id": "call_1", "function": {"name": "web_search", "arguments": "{}"}}, + ] + db.append_message( + "s1", role="assistant", content="", tool_calls=tool_calls, + timestamp=1000.0, + ) + db.append_message( + "s1", role="tool", content="result", tool_name="web_search", + tool_call_id="call_1", timestamp=999.0, + ) + db.append_message("s1", role="user", content="thanks", timestamp=998.0) + + conv = db.get_messages_as_conversation("s1") + + # Insertion order is preserved even though timestamps decrease. + assert [m["role"] for m in conv] == ["assistant", "tool", "user"] + # The tool response stays immediately after the assistant tool_calls + # row — the adjacency invariant the model API enforces. + assert conv[0]["tool_calls"][0]["id"] == "call_1" + assert conv[1]["role"] == "tool" + assert conv[1]["tool_call_id"] == "call_1" + + def test_platform_message_id_round_trips(self, db): + """Platform-side message ids (yuanbao msg_id, telegram update_id, …) + survive append → get_messages_as_conversation under the + ``message_id`` key so platform recall flows can match by exact id.""" + db.create_session(session_id="s_pmi", source="yuanbao") + db.append_message( + "s_pmi", + role="user", + content="hi", + platform_message_id="abc-123", + ) + db.append_message("s_pmi", role="assistant", content="hello") + + conv = db.get_messages_as_conversation("s_pmi") + user_msg = next(m for m in conv if m["role"] == "user") + assistant_msg = next(m for m in conv if m["role"] == "assistant") + assert user_msg.get("message_id") == "abc-123" + # Assistant row had no platform id — must not gain one spuriously. + assert "message_id" not in assistant_msg + + def test_replace_messages_preserves_platform_message_id(self, db): + """``rewrite_transcript`` (which goes through replace_messages) must + keep the platform_message_id round-trip working for /retry, /undo, + /compress and yuanbao's recall rewrite path.""" + db.create_session(session_id="s_rep", source="yuanbao") + db.replace_messages( + "s_rep", + [ + {"role": "user", "content": "x", "message_id": "ext-1"}, + {"role": "assistant", "content": "y"}, + ], + ) + conv = db.get_messages_as_conversation("s_rep") + assert next(m for m in conv if m["role"] == "user").get("message_id") == "ext-1" + assert "message_id" not in next(m for m in conv if m["role"] == "assistant") + + def test_replace_messages_preserves_internal_source_message_id(self, db): + db.create_session(session_id="s_source_replace", source="desktop") + db.replace_messages( + "s_source_replace", + [ + { + "role": "user", + "content": "queued source", + "_source_message_id": "desktop-rewrite-1", + } + ], + ) + + conv = db.get_messages_as_conversation("s_source_replace") + assert conv[0]["message_id"] == "desktop-rewrite-1" + assert db.has_platform_message_id( + "s_source_replace", "desktop-rewrite-1" + ) + + def test_archive_and_compact_preserves_internal_source_message_id(self, db): + db.create_session(session_id="s_source_compact", source="desktop") + db.append_message( + "s_source_compact", + role="user", + content="pre-compaction turn", + ) + + inserted = db.archive_and_compact( + "s_source_compact", + [ + { + "role": "user", + "content": "compacted live source", + "_source_message_id": "desktop-compact-1", + } + ], + ) + + conv = db.get_messages_as_conversation("s_source_compact") + assert inserted == 1 + assert [message["content"] for message in conv] == [ + "compacted live source" + ] + assert conv[0]["message_id"] == "desktop-compact-1" + assert db.has_platform_message_id( + "s_source_compact", "desktop-compact-1" + ) + + def test_get_messages_as_conversation_includes_ancestor_chain(self, db): + db.create_session("root", "tui") + db.append_message("root", role="user", content="first prompt") + db.append_message("root", role="assistant", content="first answer") + db.create_session("child", "tui", parent_session_id="root") + db.append_message("child", role="user", content="second prompt") + db.append_message("child", role="assistant", content="second answer") + + conv = db.get_messages_as_conversation("child", include_ancestors=True) + + assert [m["content"] for m in conv] == [ + "first prompt", + "first answer", + "second prompt", + "second answer", + ] + + def test_get_messages_as_conversation_avoids_repeated_resume_prompts_from_ancestors(self, db): + db.create_session("root", "tui") + db.append_message("root", role="user", content="same prompt") + db.append_message("root", role="user", content="same prompt") + db.append_message("root", role="assistant", content="answer") + db.create_session("child", "tui", parent_session_id="root") + db.append_message("child", role="user", content="next prompt") + + conv = db.get_messages_as_conversation("child", include_ancestors=True) + + assert [m["content"] for m in conv if m["role"] == "user"] == ["same prompt", "next prompt"] + + def test_get_resume_conversations_matches_separate_reads(self, db): + """The one-fetch resume projections must be byte-identical to the two + separate get_messages_as_conversation reads they replace — the whole + point of the single-SELECT optimization (desktop audit P1). Includes a + dangling tool-call tail so repair_alternation drops rows and the model / + display lengths diverge (exercises session.resume's prefix computation). + """ + db.create_session("root", "tui") + db.append_message("root", role="user", content="first prompt") + db.append_message("root", role="assistant", content="first answer") + db.create_session("child", "tui", parent_session_id="root") + db.append_message("child", role="user", content="second prompt") + db.append_message( + "child", role="assistant", content="second answer", finish_reason="stop" + ) + # Dangling assistant(tool_calls) tail with no tool response → repair + # drops it, so model_history is shorter than display_history. + db.append_message( + "child", + role="assistant", + content="", + tool_calls=[ + {"id": "t1", "type": "function", "function": {"name": "x", "arguments": "{}"}} + ], + ) + + model_expected = db.get_messages_as_conversation("child", repair_alternation=True, include_row_ids=True) + display_expected = db.get_messages_as_conversation("child", include_ancestors=True, include_row_ids=True) + + model_history, display_history = db.get_resume_conversations("child") + + assert model_history == model_expected + assert display_history == display_expected + # Sanity: the tail really did diverge the two projections. + assert len(display_history) > len(model_history) + + def test_get_resume_conversations_single_session_no_ancestors(self, db): + db.create_session("solo", "cli") + db.append_message("solo", role="user", content="hi") + db.append_message("solo", role="assistant", content="hello") + + model_expected = db.get_messages_as_conversation("solo", repair_alternation=True, include_row_ids=True) + display_expected = db.get_messages_as_conversation("solo", include_ancestors=True, include_row_ids=True) + model_history, display_history = db.get_resume_conversations("solo") + + assert model_history == model_expected + assert display_history == display_expected + + def test_get_resume_conversations_dedupes_replayed_ancestor_user(self, db): + db.create_session("root", "tui") + db.append_message("root", role="user", content="same prompt") + db.append_message("root", role="user", content="same prompt") + db.append_message("root", role="assistant", content="answer") + db.create_session("child", "tui", parent_session_id="root") + db.append_message("child", role="user", content="next prompt") + + model_expected = db.get_messages_as_conversation("child", repair_alternation=True, include_row_ids=True) + display_expected = db.get_messages_as_conversation("child", include_ancestors=True, include_row_ids=True) + model_history, display_history = db.get_resume_conversations("child") + + assert model_history == model_expected + assert display_history == display_expected + + def test_get_ancestor_display_prefix_single_session_returns_empty(self, db): + """A session with no compression ancestors has an empty prefix.""" + db.create_session("solo", "cli") + db.append_message("solo", role="user", content="hi") + db.append_message("solo", role="assistant", content="hello") + + assert db.get_ancestor_display_prefix("solo") == [] + + def test_get_ancestor_display_prefix_returns_ancestor_only_messages(self, db): + """The prefix contains ONLY ancestor messages, not tip messages. + + Previously the prefix was calculated as + display_history[:len(display) - len(raw)], which overcounts when + repair_message_sequence removes messages from the MIDDLE of the + tip history — the length difference includes both ancestor messages + AND repair-removed tip messages, but the slice captures the first N + display messages (tip messages when there are no ancestors), + causing duplication in _live_session_payload. (#65919) + """ + db.create_session("root", "tui") + db.append_message("root", role="user", content="ancestor prompt") + db.append_message("root", role="assistant", content="ancestor reply") + db.create_session("child", "tui", parent_session_id="root") + db.append_message("child", role="user", content="tip prompt") + db.append_message("child", role="assistant", content="tip reply") + # A verification candidate that repair_message_sequence collapses + # (consecutive-assistant merge replaces it with the next assistant). + db.append_message( + "child", + role="assistant", + content="verification candidate", + finish_reason="verification_required", + ) + db.append_message("child", role="assistant", content="post-verification reply") + + prefix = db.get_ancestor_display_prefix("child") + # Only the ancestor messages, not any tip messages. + assert len(prefix) == 2 + assert prefix[0]["role"] == "user" + assert prefix[0]["content"] == "ancestor prompt" + assert prefix[1]["role"] == "assistant" + assert prefix[1]["content"] == "ancestor reply" + + # The old broken calculation would produce a non-empty prefix + # (because repair collapses the verification candidate, making + # len(display) > len(raw)), even though there are 2 ancestor + # messages — it would overcount. + raw, display = db.get_resume_conversations("child") + old_prefix_len = max(0, len(display) - len(raw)) + assert len(prefix) <= old_prefix_len + + def test_finish_reason_stored(self, db): + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="assistant", content="Done", finish_reason="stop") + + messages = db.get_messages("s1") + assert messages[0]["finish_reason"] == "stop" def test_get_messages_as_conversation_strips_leaked_memory_context(self, db): db.create_session(session_id="s1", source="cli") diff --git a/tests/test_tui_gateway_queue_on_busy.py b/tests/test_tui_gateway_queue_on_busy.py index 86edf8aa598c..89426d63e65c 100644 --- a/tests/test_tui_gateway_queue_on_busy.py +++ b/tests/test_tui_gateway_queue_on_busy.py @@ -14,6 +14,10 @@ import types import tools.async_delegation as ad +from hermes_state import SessionDB +from run_agent import AIAgent + + from tui_gateway import server @@ -36,7 +40,8 @@ def _session(agent=None, **extra): def test_enqueue_pins_text_and_transport(): session = _session() server._enqueue_prompt(session, "hello", "ws-1") - assert session["queued_prompt"] == {"text": "hello", "transport": "ws-1"} + assert session["queued_prompt"]["text"] == "hello" + assert session["queued_prompt"]["transport"] == "ws-1" def test_enqueue_preserves_order_after_an_image_turn(): @@ -52,6 +57,53 @@ def test_enqueue_preserves_order_after_an_image_turn(): ] +def test_enqueue_preserves_distinct_messages_and_submission_metadata(): + session = _session() + server._enqueue_prompt( + session, + "first", + "ws-1", + submitted_at=101.25, + message_id="desktop-1", + ) + server._enqueue_prompt( + session, + "second", + "ws-2", + submitted_at=102.5, + message_id="desktop-2", + ) + + assert session["queued_prompt"] == { + "text": "first", + "transport": "ws-1", + "submitted_at": 101.25, + "message_id": "desktop-1", + } + assert session["queued_prompts"] == [ + { + "text": "second", + "transport": "ws-2", + "submitted_at": 102.5, + "message_id": "desktop-2", + } + ] + + +def test_enqueue_keeps_one_multi_paragraph_prompt_as_one_message(): + session = _session() + text = "first paragraph\n\nsecond paragraph" + + server._enqueue_prompt( + session, + text, + "ws-1", + submitted_at=101.25, + message_id="desktop-1", + ) + + assert session["queued_prompt"]["text"] == text + assert session.get("queued_prompts", []) == [] # ── _handle_busy_submit (policy) ─────────────────────────────────────────── @@ -346,17 +398,89 @@ def test_busy_interrupt_mode_ignores_completed_background_delegation(monkeypatch -def test_busy_steer_mode_injects_when_accepted(monkeypatch): +def test_busy_steer_mode_injects_when_accepted_without_enqueueing(monkeypatch): monkeypatch.setattr(server, "_load_busy_input_mode", lambda: "steer") - agent = types.SimpleNamespace(steer=lambda text: True, interrupt=lambda *a, **k: None) + calls = {"interrupt": 0, "steer": []} + agent = types.SimpleNamespace( + interrupt=lambda: calls.__setitem__("interrupt", calls["interrupt"] + 1), + steer=lambda text: (calls["steer"].append(text), True)[1], + ) session = _session(agent=agent, running=True) - resp = server._handle_busy_submit("r1", "sid", session, "nudge", "ws-1") + resp = server._handle_busy_submit( + "r1", + "sid", + session, + "nudge", + "ws-1", + submitted_at=101.25, + message_id="desktop-steer-1", + ) assert resp["result"]["status"] == "steered" + assert calls == {"interrupt": 0, "steer": ["nudge"]} assert session.get("queued_prompt") is None +def test_busy_steer_mode_rejection_queues_with_source_identity(monkeypatch): + monkeypatch.setattr(server, "_load_busy_input_mode", lambda: "steer") + calls = {"interrupt": 0, "steer": []} + agent = types.SimpleNamespace( + interrupt=lambda: calls.__setitem__("interrupt", calls["interrupt"] + 1), + steer=lambda text: (calls["steer"].append(text), False)[1], + ) + session = _session(agent=agent, running=True) + + resp = server._handle_busy_submit( + "r1", + "sid", + session, + "nudge", + "ws-1", + submitted_at=101.25, + message_id="desktop-steer-1", + ) + + assert resp["result"]["status"] == "queued" + # #86134: steer fall-through must not hard-interrupt (kills buffered steers). + assert calls == {"interrupt": 0, "steer": ["nudge"]} + assert session["queued_prompt"] == { + "text": "nudge", + "transport": "ws-1", + "submitted_at": 101.25, + "message_id": "desktop-steer-1", + } + + +def test_busy_steer_mode_unavailable_queues_with_source_identity(monkeypatch): + monkeypatch.setattr(server, "_load_busy_input_mode", lambda: "steer") + calls = {"interrupt": 0} + agent = types.SimpleNamespace( + interrupt=lambda: calls.__setitem__("interrupt", calls["interrupt"] + 1) + ) + session = _session(agent=agent, running=True) + + resp = server._handle_busy_submit( + "r1", + "sid", + session, + "nudge", + "ws-1", + submitted_at=101.25, + message_id="desktop-steer-1", + ) + + assert resp["result"]["status"] == "queued" + # #86134: unavailable steer queues without a hard interrupt. + assert calls["interrupt"] == 0 + assert session["queued_prompt"] == { + "text": "nudge", + "transport": "ws-1", + "submitted_at": 101.25, + "message_id": "desktop-steer-1", + } + + # ── steer-mode burst preservation (#86134) ───────────────────────────────── def test_busy_steer_fallthrough_queues_without_interrupting(monkeypatch): @@ -509,6 +633,118 @@ def test_busy_helper_retries_when_turn_finished(monkeypatch): assert server._handle_busy_submit("r1", "sid", session, "run now", "ws-1") is None assert session.get("queued_prompt") is None +def test_prompt_submit_dedupes_explicit_id_already_inflight(monkeypatch): + calls = {"interrupt": 0} + agent = types.SimpleNamespace( + interrupt=lambda: calls.__setitem__("interrupt", calls["interrupt"] + 1) + ) + session = _session( + agent=agent, + running=True, + inflight_turn={"message_id": "desktop-1", "user": "first"}, + ) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "current_transport", lambda: "ws-2") + + response = server.handle_request( + { + "id": "rpc-2", + "method": "prompt.submit", + "params": { + "message_id": "desktop-1", + "session_id": "sid", + "text": "first", + }, + } + ) + + assert response is not None + assert response["result"]["status"] == "duplicate" + assert session.get("queued_prompt") is None + assert calls["interrupt"] == 0 + + + +def test_prompt_submit_duplicate_rehomes_only_matching_queued_source(monkeypatch): + session = _session( + running=True, + transport="ws-current", + queued_prompt={ + "text": "first", + "transport": "ws-old", + "message_id": "desktop-1", + }, + queued_prompts=[ + { + "text": "second", + "transport": "ws-still-live", + "message_id": "desktop-2", + } + ], + ) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "current_transport", lambda: "ws-retry") + + response = server.handle_request( + { + "id": "rpc-retry", + "method": "prompt.submit", + "params": { + "message_id": "desktop-1", + "session_id": "sid", + "text": "first", + }, + } + ) + + assert response is not None + assert response["result"]["status"] == "duplicate" + assert session["transport"] == "ws-retry" + assert session["queued_prompt"]["transport"] == "ws-retry" + assert session["queued_prompts"][0]["transport"] == "ws-still-live" + + +def test_prompt_submit_does_not_dedupe_reused_rpc_id_without_explicit_id( + monkeypatch, +): + monkeypatch.setattr(server, "_load_busy_input_mode", lambda: "queue") + session = _session( + running=True, + inflight_turn={"message_id": "rpc-1", "user": "prior connection"}, + ) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "current_transport", lambda: "ws-new") + + response = server.handle_request( + { + "id": "rpc-1", + "method": "prompt.submit", + "params": {"session_id": "sid", "text": "new connection prompt"}, + } + ) + + assert response is not None + assert response["result"]["status"] == "queued" + assert session["queued_prompt"]["text"] == "new connection prompt" + assert "message_id" not in session["queued_prompt"] + + +def test_prompt_id_dedupe_uses_persisted_source_id(tmp_path): + db = SessionDB(tmp_path / "dedupe.db") + try: + db.create_session("session-key", source="desktop", model="test/model") + db.append_message( + session_id="session-key", + role="user", + content="already accepted", + platform_message_id="desktop-persisted", + ) + session = _session(agent=types.SimpleNamespace(_session_db=db)) + + assert server._has_prompt_message_id(session, "desktop-persisted") is True + assert server._has_prompt_message_id(session, "desktop-new") is False + finally: + db.close() @@ -678,15 +914,84 @@ def test_drain_preserves_queued_prompt_when_session_is_closing(monkeypatch): assert session["running"] is False -def test_drain_releases_running_on_dispatch_failure(monkeypatch): - def _boom(*a, **k): +def test_drain_failure_restores_exact_item_before_later_arrivals(monkeypatch): + first = { + "text": "first", + "transport": "ws-1", + "submitted_at": 101.25, + "message_id": "desktop-1", + } + second = { + "text": "second", + "transport": "ws-2", + "submitted_at": 102.5, + "message_id": "desktop-2", + } + + def _boom(_rid, _sid, session, _text, **_kwargs): + server._enqueue_prompt( + session, + "third", + "ws-3", + submitted_at=103.75, + message_id="desktop-3", + ) raise RuntimeError("dispatch failed") + monkeypatch.setattr(server, "_run_prompt_submit", _boom) - session = _session(queued_prompt={"text": "go", "transport": None}) + session = _session(queued_prompt=first, queued_prompts=[second]) assert server._drain_queued_prompt("r1", "sid", session) is True - # Failure must not leave the session wedged as running. assert session["running"] is False + assert session["inflight_turn"] is None + assert session["queued_prompt"] is first + assert session["queued_prompts"] == [ + second, + { + "text": "third", + "transport": "ws-3", + "submitted_at": 103.75, + "message_id": "desktop-3", + }, + ] + + +def test_drain_claim_dedupes_retry_before_dispatch(monkeypatch): + retry_response = None + queued = { + "text": "first", + "transport": "ws-original", + "submitted_at": 101.25, + "message_id": "stable-1", + } + session = _session(queued_prompt=queued) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "current_transport", lambda: "ws-retry") + + def _run(_rid, _sid, _session, _text, **_kwargs): + nonlocal retry_response + retry_response = server.handle_request( + { + "id": "rpc-retry", + "method": "prompt.submit", + "params": { + "message_id": "stable-1", + "session_id": "sid", + "submitted_at": 101.25, + "text": "first", + }, + } + ) + + monkeypatch.setattr(server, "_run_prompt_submit", _run) + + assert server._drain_queued_prompt("rpc-original", "sid", session) is True + assert retry_response is not None + assert retry_response["result"]["status"] == "duplicate" + assert session.get("queued_prompt") is None + assert session.get("queued_prompts", []) == [] + assert session["inflight_turn"]["message_id"] == "stable-1" + assert session["inflight_turn"]["submitted_at"] == 101.25 def test_drain_does_not_dispatch_a_prompt_cancelled_after_claim(monkeypatch): @@ -780,3 +1085,543 @@ def _run(_rid, _sid, session, text, **_kwargs): assert session["queued_prompt"] is None assert session.get("queued_prompts") is None +def test_repeated_arrivals_drain_once_in_order_to_their_own_transports(monkeypatch): + fired = [] + + def _run(rid, sid, session, text, **kwargs): + kwargs.pop("queued_prompt_generation", None) + fired.append( + { + "rid": rid, + "sid": sid, + "text": text, + "transport": session["transport"], + **kwargs, + } + ) + session["running"] = False + + monkeypatch.setattr(server, "_run_prompt_submit", _run) + session = _session() + for index in range(3): + server._enqueue_prompt( + session, + f"message-{index}", + f"ws-{index}", + submitted_at=100.0 + index, + message_id=f"desktop-{index}", + ) + + assert server._drain_queued_prompt("r1", "sid", session) is True + assert server._drain_queued_prompt("r1", "sid", session) is True + assert server._drain_queued_prompt("r1", "sid", session) is True + assert server._drain_queued_prompt("r1", "sid", session) is False + + assert fired == [ + { + "rid": "r1", + "sid": "sid", + "text": f"message-{index}", + "transport": f"ws-{index}", + "submitted_at": 100.0 + index, + "message_id": f"desktop-{index}", + } + for index in range(3) + ] + assert session["queued_prompt"] is None + assert session.get("queued_prompts", []) == [] + + +class _RecordingTransport: + def __init__(self, completed: threading.Event | None = None): + self._closed = False + self.completed = completed + self.frames = [] + + def write(self, obj): + self.frames.append(obj) + event_type = ((obj.get("params") or {}).get("type")) + if event_type == "message.complete" and self.completed is not None: + self.completed.set() + return not self._closed + + def close(self): + self._closed = True + + +def test_session_activate_rehomes_dead_queue_item_and_preserves_live_tail( + monkeypatch, +): + dead_head_transport = _RecordingTransport() + dead_head_transport.close() + current_live_transport = _RecordingTransport() + live_tail_transport = _RecordingTransport() + activated_transport = _RecordingTransport() + session = _session( + transport=current_live_transport, + queued_prompt={ + "text": "dead head", + "transport": dead_head_transport, + "message_id": "desktop-dead", + }, + queued_prompts=[ + { + "text": "live tail", + "transport": live_tail_transport, + "message_id": "desktop-live", + } + ], + ) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "_session_info", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "current_transport", lambda: activated_transport) + + response = server.handle_request( + { + "id": "rpc-activate", + "method": "session.activate", + "params": {"session_id": "sid"}, + } + ) + + assert response["result"]["session_id"] == "sid" + assert session["transport"] is activated_transport + assert session["queued_prompt"]["transport"] is activated_transport + assert session["queued_prompts"][0]["transport"] is live_tail_transport + + +def test_disconnect_snapshot_cannot_overwrite_inflight_duplicate_retry( + monkeypatch, +): + snapshot_taken = threading.Event() + finish_snapshot = threading.Event() + old_transport = _RecordingTransport() + new_transport = _RecordingTransport() + old_transport.close() + sid = "race-ui" + session = _session( + running=True, + transport=old_transport, + inflight_turn={ + "message_id": "desktop-race-1", + "user": "survive disconnect race", + }, + ) + + class _SnapshotBarrierSessions(dict): + def items(self): + snapshot = list(super().items()) + snapshot_taken.set() + finish_snapshot.wait(10) + return snapshot + + sessions = _SnapshotBarrierSessions({sid: session}) + disconnect_result = {} + disconnect_errors = [] + + def _disconnect(): + try: + disconnect_result["value"] = server._close_sessions_for_transport( + old_transport + ) + except BaseException as exc: + disconnect_errors.append(exc) + + monkeypatch.setattr(server, "_sessions", sessions) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "current_transport", lambda: new_transport) + monkeypatch.setattr(server, "_schedule_ws_orphan_reap", lambda *_a, **_k: None) + disconnect_thread = threading.Thread(target=_disconnect) + disconnect_thread.start() + + try: + assert snapshot_taken.wait(10), "disconnect did not snapshot the old owner" + duplicate = server.handle_request( + { + "id": "rpc-retry", + "method": "prompt.submit", + "params": { + "message_id": "desktop-race-1", + "session_id": sid, + "text": "survive disconnect race", + }, + } + ) + assert duplicate["result"]["status"] == "duplicate" + assert session["transport"] is new_transport + finally: + finish_snapshot.set() + disconnect_thread.join(10) + + assert not disconnect_thread.is_alive() + assert disconnect_errors == [] + assert disconnect_result["value"] == (0, 0) + assert session["transport"] is new_transport + + server._emit("message.start", sid) + server._emit("message.delta", sid, {"text": "delta"}) + server._emit("message.complete", sid, {"text": "complete"}) + assert old_transport.frames == [] + assert [ + (frame.get("params") or {}).get("type") for frame in new_transport.frames + ] == ["message.start", "message.delta", "message.complete"] + + +def _model_response(text): + message = types.SimpleNamespace( + content=text, + tool_calls=None, + reasoning_content=None, + reasoning=None, + ) + choice = types.SimpleNamespace(message=message, finish_reason="stop") + return types.SimpleNamespace(choices=[choice], model="test/model", usage=None) + + +def test_busy_steer_rejection_dedupes_and_persists_one_canonical_turn( + monkeypatch, + tmp_path, +): + db = SessionDB(tmp_path / "steer-source.db") + session_key = "steer-source" + db.create_session(session_key, source="desktop", model="test/model") + agent = AIAgent( + api_key="test-key", + base_url="https://example.invalid/v1", + provider="custom", + model="test/model", + api_mode="chat_completions", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + session_db=db, + session_id=session_key, + ) + agent._session_db_created = True + agent._cached_system_prompt = "You are a test assistant." + agent._disable_streaming = True + + steer_calls = [] + interrupt_calls = [] + wire_requests = [] + completed = threading.Event() + monkeypatch.setattr(agent, "steer", lambda text: (steer_calls.append(text), False)[1]) + monkeypatch.setattr(agent, "interrupt", lambda: interrupt_calls.append(True)) + monkeypatch.setattr( + agent, + "_interruptible_api_call", + lambda api_kwargs: ( + wire_requests.append(api_kwargs["messages"]), + _model_response("ack"), + )[1], + ) + monkeypatch.setattr(agent, "_cleanup_task_resources", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_load_busy_input_mode", lambda: "steer") + monkeypatch.setattr(server, "_sync_agent_model_with_config", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_wire_callbacks", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_register_session_cwd", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_session_info", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "_get_usage", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "_voice_tts_enabled", lambda: False) + monkeypatch.setattr("agent.title_generator.maybe_auto_title", lambda *_a, **_k: None) + + session = _session(agent=agent, session_key=session_key, running=True) + monkeypatch.setattr(server, "_sess_nowait", lambda *_a, **_k: (session, None)) + monkeypatch.setattr(server, "current_transport", lambda: "ws-steer") + monkeypatch.setattr( + server, + "_emit", + lambda event, _sid, _payload=None: completed.set() + if event == "message.complete" + else None, + ) + request = { + "id": "rpc-steer", + "method": "prompt.submit", + "params": { + "message_id": "desktop-steer-1", + "session_id": "ui-session", + "submitted_at": 101.25, + "text": "canonical nudge", + }, + } + + first = server.handle_request(request) + duplicate = server.handle_request({**request, "id": "rpc-steer-retry"}) + + assert first["result"]["status"] == "queued" + assert duplicate["result"]["status"] == "duplicate" + assert steer_calls == ["canonical nudge"] + # #86134: steer rejection queues without hard-interrupting the live turn. + assert interrupt_calls == [] + assert session["queued_prompt"]["message_id"] == "desktop-steer-1" + assert session.get("queued_prompts", []) == [] + + session["running"] = False + assert server._drain_queued_prompt("rpc-steer", "ui-session", session) is True + assert completed.wait(10), "steer-fallback queued turn did not complete" + + canonical_users = [ + message for message in session["history"] if message.get("role") == "user" + ] + user_rows = [row for row in db.get_messages(session_key) if row["role"] == "user"] + assert [message["content"] for message in canonical_users] == ["canonical nudge"] + assert [(row["content"], row["platform_message_id"]) for row in user_rows] == [ + ("canonical nudge", "desktop-steer-1") + ] + assert len(wire_requests) == 1 + + +def test_reconnect_rehomes_queued_turn_and_routes_all_events_to_live_transport( + monkeypatch, + tmp_path, +): + db = SessionDB(tmp_path / "reconnect-source.db") + session_key = "reconnect-source" + sid = "reconnect-ui" + db.create_session(session_key, source="desktop", model="test/model") + agent = AIAgent( + api_key="test-key", + base_url="https://example.invalid/v1", + provider="custom", + model="test/model", + api_mode="chat_completions", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + session_db=db, + session_id=session_key, + ) + agent._session_db_created = True + agent._cached_system_prompt = "You are a test assistant." + agent._disable_streaming = True + + wire_requests = [] + completed = threading.Event() + old_transport = _RecordingTransport() + new_transport = _RecordingTransport(completed) + active_transport = {"value": old_transport} + original_run = agent.run_conversation + + def _run_with_delta(user_message, **kwargs): + result = original_run(user_message, **kwargs) + kwargs["stream_callback"]("ack-delta") + return result + + monkeypatch.setattr(agent, "run_conversation", _run_with_delta) + monkeypatch.setattr( + agent, + "_interruptible_api_call", + lambda api_kwargs: ( + wire_requests.append(api_kwargs["messages"]), + _model_response("ack"), + )[1], + ) + monkeypatch.setattr(agent, "_cleanup_task_resources", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_load_busy_input_mode", lambda: "queue") + monkeypatch.setattr(server, "_sync_agent_model_with_config", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_wire_callbacks", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_register_session_cwd", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_session_info", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "_get_usage", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "_voice_tts_enabled", lambda: False) + monkeypatch.setattr(server, "_get_db", lambda: db) + monkeypatch.setattr(server, "current_transport", lambda: active_transport["value"]) + monkeypatch.setattr(server, "_schedule_ws_orphan_reap", lambda *_a, **_k: None) + monkeypatch.setattr("agent.title_generator.maybe_auto_title", lambda *_a, **_k: None) + + session = _session( + agent=agent, + session_key=session_key, + running=True, + transport=old_transport, + ) + missing = object() + previous_session = server._sessions.get(sid, missing) + server._sessions[sid] = session + request = { + "id": "rpc-original", + "method": "prompt.submit", + "params": { + "message_id": "desktop-reconnect-1", + "session_id": sid, + "submitted_at": 101.25, + "text": "survive reconnect", + }, + } + + try: + first = server.handle_request(request) + assert first["result"]["status"] == "queued" + assert session["queued_prompt"]["transport"] is old_transport + + old_transport.close() + server._close_sessions_for_transport(old_transport) + assert session["transport"] is server._detached_ws_transport + + active_transport["value"] = new_transport + resumed = server.handle_request( + { + "id": "rpc-resume", + "method": "session.resume", + "params": {"session_id": session_key}, + } + ) + assert resumed["result"]["session_id"] == sid + assert session["transport"] is new_transport + assert session["queued_prompt"]["transport"] is new_transport + + duplicate = server.handle_request({**request, "id": "rpc-retry"}) + assert duplicate["result"]["status"] == "duplicate" + assert session["queued_prompt"]["transport"] is new_transport + + with session["history_lock"]: + session["running"] = False + assert server._drain_queued_prompt("rpc-drain", sid, session) is True + assert completed.wait(10), "reconnected client did not receive completion" + run_thread = session.get("_run_thread") + assert run_thread is not None + run_thread.join(10) + assert not run_thread.is_alive() + + old_event_types = [ + (frame.get("params") or {}).get("type") for frame in old_transport.frames + ] + new_event_types = [ + (frame.get("params") or {}).get("type") for frame in new_transport.frames + ] + assert old_event_types == [] + assert { + "message.start", + "message.delta", + "message.complete", + }.issubset(new_event_types) + + canonical_users = [ + message for message in session["history"] if message.get("role") == "user" + ] + assert len(canonical_users) == 1 + assert canonical_users[0]["content"] == "survive reconnect" + assert canonical_users[0]["timestamp"] == 101.25 + assert canonical_users[0]["_source_message_id"] == "desktop-reconnect-1" + + user_rows = [row for row in db.get_messages(session_key) if row["role"] == "user"] + assert [(row["content"], row["platform_message_id"]) for row in user_rows] == [ + ("survive reconnect", "desktop-reconnect-1") + ] + assert len(wire_requests) == 1 + assert session["queued_prompt"] is None + assert session.get("queued_prompts", []) == [] + assert session["inflight_turn"] is None + finally: + run_thread = session.get("_run_thread") + if run_thread is not None: + run_thread.join(10) + if previous_session is missing: + server._sessions.pop(sid, None) + else: + server._sessions[sid] = previous_session + db.close() + + +def test_drain_persists_distinct_users_and_sends_valid_ordered_wire_history( + monkeypatch, + tmp_path, +): + """Exercise the real gateway drain, AIAgent loop, SessionDB, and wire copy.""" + db = SessionDB(tmp_path / "state.db") + session_key = "queued-boundaries" + db.create_session(session_key, source="desktop", model="test/model") + agent = AIAgent( + api_key="test-key", + base_url="https://example.invalid/v1", + provider="custom", + model="test/model", + api_mode="chat_completions", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + session_db=db, + session_id=session_key, + ) + agent._session_db_created = True + agent._cached_system_prompt = "You are a test assistant." + agent._disable_streaming = True + + wire_requests = [] + replies = iter(("ack-first", "ack-second")) + + def _api_call(api_kwargs): + wire_requests.append(api_kwargs["messages"]) + return _model_response(next(replies)) + + monkeypatch.setattr(agent, "_interruptible_api_call", _api_call) + monkeypatch.setattr(agent, "_cleanup_task_resources", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_sync_agent_model_with_config", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_wire_callbacks", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_register_session_cwd", lambda *_a, **_k: None) + monkeypatch.setattr(server, "_session_info", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "_get_usage", lambda *_a, **_k: {}) + monkeypatch.setattr(server, "_voice_tts_enabled", lambda: False) + monkeypatch.setattr("agent.title_generator.maybe_auto_title", lambda *_a, **_k: None) + + completed = threading.Event() + completion_count = 0 + + def _emit(event, _sid, _payload=None): + nonlocal completion_count + if event == "message.complete": + completion_count += 1 + if completion_count == 2: + completed.set() + + monkeypatch.setattr(server, "_emit", _emit) + + session = _session(agent=agent, session_key=session_key) + server._enqueue_prompt( + session, + "first queued prompt", + "ws-1", + submitted_at=101.25, + message_id="desktop-1", + ) + server._enqueue_prompt( + session, + "second queued prompt", + "ws-2", + submitted_at=102.5, + message_id="desktop-2", + ) + + assert server._drain_queued_prompt("r1", "ui-session", session) is True + assert completed.wait(10), "queued turns did not both complete" + + user_rows = [row for row in db.get_messages(session_key) if row["role"] == "user"] + assert [row["content"] for row in user_rows] == [ + "first queued prompt", + "second queued prompt", + ] + assert [row["timestamp"] for row in user_rows] == [101.25, 102.5] + assert [row["platform_message_id"] for row in user_rows] == [ + "desktop-1", + "desktop-2", + ] + + assert len(wire_requests) == 2 + assert [ + message["content"] + for message in wire_requests[1] + if message.get("role") == "user" + ] == ["first queued prompt", "second queued prompt"] + for request in wire_requests: + non_system_roles = [ + message["role"] for message in request if message.get("role") != "system" + ] + assert all( + left != right + for left, right in zip(non_system_roles, non_system_roles[1:]) + ) + assert all("timestamp" not in message for message in request) + assert all("_source_message_id" not in message for message in request) + assert all("message_id" not in message for message in request) + assert all("platform_message_id" not in message for message in request) diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index f8290f8442fb..a4d958edf74e 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -414,11 +414,19 @@ def start(self): monkeypatch.setattr(server, "_wait_agent", lambda _session, _rid: None) # The deferred inline-fallback thread now waits via the patient variant. monkeypatch.setattr(server, "_wait_agent_for_prompt", lambda _session, _rid, _sid: None) - monkeypatch.setattr( - server, - "_run_prompt_submit", - lambda rid, sid, _session, text, **_kwargs: inline_calls.append((rid, sid, text)), - ) + def _run_inline( + rid, + sid, + _session, + text, + *, + submitted_at=None, + message_id=None, + **_kwargs, + ): + inline_calls.append((rid, sid, text, submitted_at, message_id)) + + monkeypatch.setattr(server, "_run_prompt_submit", _run_inline) monkeypatch.setattr(server.threading, "Thread", _ImmediateThread) try: @@ -426,7 +434,12 @@ def start(self): { "id": "fallback-turn", "method": "prompt.submit", - "params": {"session_id": "iso-fallback", "text": "hello"}, + "params": { + "session_id": "iso-fallback", + "text": "hello", + "submitted_at": 123.5, + "message_id": "fallback-message", + }, } ) finally: @@ -437,7 +450,9 @@ def start(self): "id": "fallback-turn", "result": {"status": "streaming"}, } - assert inline_calls == [("fallback-turn", "iso-fallback", "hello")] + assert inline_calls == [ + ("fallback-turn", "iso-fallback", "hello", 123.5, "fallback-message") + ] assert session.get("_compute_host_active") is not True @@ -17760,7 +17775,11 @@ def test_close_sessions_for_transport_closes_flagged_repoints_rest(monkeypatch): transport = object() # the disconnecting transport server._sessions.clear() server._sessions["a"] = {"transport": transport, "close_on_disconnect": True} - server._sessions["b"] = {"transport": transport, "close_on_disconnect": False} + server._sessions["b"] = { + "transport": transport, + "close_on_disconnect": False, + "history_lock": threading.Lock(), + } try: server._close_sessions_for_transport(transport, end_reason="ws_disconnect") assert seen == [("a", "ws_disconnect")] # only the flagged one closed @@ -20753,9 +20772,10 @@ def test_prompt_submit_rebind_map_clears_active_row_hidden_by_sequence_repair( 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 { + # Canonical repair preserves adjacent user turns as distinct rows (the + # user;user merge happens later on the per-request provider copy), so the + # replay view keeps the second physical user's row identity. + assert physical_ids[1] in { server._message_row_id(message) for message in repaired } @@ -20784,7 +20804,14 @@ def test_prompt_submit_rebind_map_clears_active_row_hidden_by_sequence_repair( ) 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 + # Survivors rebind to their fresh row ids — the preserved second + # user turn is a survivor, not a row hidden by repair. + assert isinstance(row_id_map[str(physical_ids[1])], int) + # Rows dropped by the truncation clear to None so the client drops + # its cached stamp instead of keeping a stale one. + assert row_id_map[str(physical_ids[3])] is None + assert row_id_map[str(physical_ids[4])] is None + # A requested id that never existed stays out of the map entirely. assert "999999" not in row_id_map finally: server._sessions.pop(sid, None) diff --git a/tests/test_tui_gateway_ws.py b/tests/test_tui_gateway_ws.py index 7dbebed65c10..358b7710acb6 100644 --- a/tests/test_tui_gateway_ws.py +++ b/tests/test_tui_gateway_ws.py @@ -82,6 +82,23 @@ def close(self): server._sessions.clear() +def test_ws_disconnect_preserves_and_repoints_reconnectable_session(monkeypatch): + server._sessions.clear() + try: + _run_disconnect( + monkeypatch, + lambda t: server._sessions.update( + plain={ + "transport": t, + "close_on_disconnect": False, + "session_key": "k", + "history_lock": threading.Lock(), + } + ), + ) + assert server._sessions["plain"]["transport"] is server._detached_ws_transport + finally: + server._sessions.clear() def test_ws_connection_registers_then_disconnect_unregisters_live_transport(monkeypatch): diff --git a/tui_gateway/compute_host.py b/tui_gateway/compute_host.py index 2cb60cf20387..081e0b313c47 100644 --- a/tui_gateway/compute_host.py +++ b/tui_gateway/compute_host.py @@ -467,7 +467,11 @@ def _run_real_turn(self, frame: dict[str, Any]) -> None: session["running"] = True session["_turn_cancel_requested"] = False session["last_active"] = time.time() - server._start_inflight_turn(session, frame.get("text") if "text" in frame else frame.get("prompt")) + server._start_inflight_turn( + session, + frame.get("text") if "text" in frame else frame.get("prompt"), + message_id=frame.get("message_id"), + ) self.emit({"type": "turn.started", "sid": sid, "request_id": request_id, "started_ns": now_ns()}) try: server._ensure_session_db_row(session) @@ -490,6 +494,8 @@ def _run_real_turn(self, frame: dict[str, Any]) -> None: session, text, display_kind=frame.get("display_kind") or None, + submitted_at=frame.get("submitted_at"), + message_id=frame.get("message_id"), ) run_thread = session.get("_run_thread") if run_thread is not None and hasattr(run_thread, "join"): diff --git a/tui_gateway/methods_prompt.py b/tui_gateway/methods_prompt.py index 7e2f816133e3..9b364c50325d 100644 --- a/tui_gateway/methods_prompt.py +++ b/tui_gateway/methods_prompt.py @@ -332,6 +332,25 @@ def _(rid, params: dict) -> dict: from tools.tts_streaming import mark_speech_interrupted mark_speech_interrupted() + + raw_submitted_at = params.get("submitted_at") + try: + explicit_submitted_at = ( + float(raw_submitted_at) + if raw_submitted_at is not None + else None + ) + except (TypeError, ValueError): + explicit_submitted_at = None + raw_message_id = params.get("message_id") + explicit_message_id = ( + str(raw_message_id).strip() if raw_message_id is not None else None + ) or None + # JSON-RPC request ids are transport-local sequence numbers that may be + # reused after reconnect. Only a client-supplied stable source id belongs + # in canonical history / SessionDB platform_message_id. + message_id = explicit_message_id + session, err = _sess_nowait(params, rid) if err: return err @@ -359,12 +378,21 @@ def _(rid, params: dict) -> dict: turn_isolation = _session_uses_compute_host(session, isolation_cfg) # Re-bind to the current client transport for this request. This keeps # streaming events on the active websocket even if an earlier disconnect - # or fallback moved the session transport to stdio. - if (t := current_transport()) is not None: - session["transport"] = t + # or fallback moved the session transport to stdio. Any matching queued + # source is re-homed atomically with the bind: a reconnect retry must not + # leave its queue entry pinned to the disconnected websocket. + t = current_transport() while True: busy_transport = None with session["history_lock"]: + if t is not None: + _rebind_session_transport( + session, t, message_id=explicit_message_id + ) + if explicit_message_id is not None and _has_prompt_message_id( + session, explicit_message_id + ): + return _ok(rid, {"status": "duplicate"}) if session.get("running"): # Don't reject a mid-turn prompt — queue it (and, by default, # interrupt the live turn) so it runs as the next turn. The @@ -376,6 +404,8 @@ def _(rid, params: dict) -> dict: busy_response = _handle_busy_submit( rid, sid, session, text, busy_transport, queued=bool(params.get("queued")), + submitted_at=explicit_submitted_at, + message_id=message_id, ) if busy_response is not None: return busy_response @@ -399,6 +429,14 @@ def _(rid, params: dict) -> dict: else None ) with session["history_lock"]: + # A Desktop queue entry keeps the same explicit ID across + # timeout/resume retries. Acknowledge an already-owned ID instead of + # accepting a duplicate turn. Do not dedupe the JSON-RPC request ID + # fallback: clients may reuse it after reconnect. + if explicit_message_id is not None and _has_prompt_message_id( + session, explicit_message_id + ): + return _ok(rid, {"status": "duplicate"}) # 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 # racing the in-flight child on the same stored session (interleaved @@ -908,7 +946,21 @@ def run_after_agent_ready() -> None: }, ) return - _run_prompt_submit(rid, sid, session, text, display_kind=display_kind) + _run_prompt_submit( + rid, + sid, + session, + text, + display_kind=display_kind, + **{ + key: value + for key, value in ( + ("submitted_at", explicit_submitted_at), + ("message_id", message_id), + ) + if value is not None + }, + ) run_thread = threading.Thread(target=run_after_agent_ready, daemon=True) # Keep a handle so session.interrupt can tell a live turn from a stuck diff --git a/tui_gateway/server.py b/tui_gateway/server.py index a2fe9bd5cb0e..58825a1d0f28 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -2614,6 +2614,9 @@ def _compute_host_turn_frame( image_paths: list[str] | None = None, queued_prompt_generation: int | None = None, display_kind: str | None = None, + *, + submitted_at: float | None = None, + message_id: str | None = None, ) -> dict: with session["history_lock"]: history = list(session.get("history", [])) @@ -2641,6 +2644,8 @@ def _compute_host_turn_frame( "source": _session_source(session), "attached_images": attached_images, "queued_prompt_generation": queued_prompt_generation, + "submitted_at": submitted_at, + "message_id": message_id, } @@ -2719,6 +2724,9 @@ def _submit_prompt_to_compute_host( image_paths: list[str] | None = None, queued_prompt_generation: int | None = None, display_kind: str | None = None, + *, + submitted_at: float | None = None, + message_id: str | None = None, ) -> dict: cfg = _load_dashboard_process_isolation_config() frame = _compute_host_turn_frame( @@ -2729,6 +2737,8 @@ def _submit_prompt_to_compute_host( image_paths=image_paths, queued_prompt_generation=queued_prompt_generation, display_kind=display_kind, + submitted_at=submitted_at, + message_id=message_id, ) def _complete(done: dict) -> None: @@ -9176,6 +9186,11 @@ def _history_to_messages(history: list[dict]) -> list[dict]: # reactions — needs this instead. if m.get("_row_id") is not None: msg["row_id"] = m["_row_id"] + if m.get("timestamp") is not None: + msg["timestamp"] = m["timestamp"] + source_message_id = m.get("message_id") or m.get("_source_message_id") + if source_message_id is not None: + msg["message_id"] = source_message_id if role == "user": invocation = _skill_scaffold_projection(content_text) if invocation: @@ -9259,7 +9274,13 @@ def _inflight_text(value: Any) -> str: return _content_display_text(value).strip() -def _start_inflight_turn(session: dict, text: Any) -> None: +def _start_inflight_turn( + session: dict, + text: Any, + *, + submitted_at: float | None = None, + message_id: str | None = None, +) -> None: now = time.time() session["inflight_turn"] = { "assistant": "", @@ -9268,6 +9289,10 @@ def _start_inflight_turn(session: dict, text: Any) -> None: "updated_at": now, "user": _inflight_text(text), } + if submitted_at is not None: + session["inflight_turn"]["submitted_at"] = submitted_at + if message_id is not None: + session["inflight_turn"]["message_id"] = message_id def _append_inflight_delta(session: dict, delta: Any) -> None: @@ -9514,15 +9539,17 @@ def _enqueue_prompt( text: Any, transport: Any, image_paths: list[str] | None = None, + *, + submitted_at: float | None = None, + message_id: str | None = None, ) -> None: - """Stash a message to run as the very next turn once the live one ends. - - Used when a prompt arrives mid-turn (see ``_handle_busy_submit``). Text-only - arrivals share a slot and merge losslessly (mirroring the consecutive-user - merge in ``repair_message_sequence``). Image-bearing submissions stay as - separate envelopes, so their attachment ownership and chronology survive. - ``transport`` is pinned so the drained turn streams back to the client that - sent it even if the session transport is rebound meanwhile. + """Append one canonical source message to the busy-time FIFO. + + ``queued_prompt`` remains the inspectable head slot for compatibility. + Later arrivals live in ``queued_prompts`` rather than being concatenated, + because text concatenation irreversibly destroys source boundaries + (``submitted_at``/``message_id``/image ownership). Each item pins its own + transport and optional source metadata until its turn is drained. """ image_paths = list(image_paths or []) # #84417: scrub any live-turn self-duplicates first so the consecutive-text @@ -9542,22 +9569,14 @@ def _enqueue_prompt( queued = {"text": text, "transport": transport} if image_paths: queued["image_paths"] = image_paths - existing = session.get("queued_prompt") - if ( - existing - and isinstance(existing.get("text"), str) - and isinstance(text, str) - and not existing.get("image_paths") - and not image_paths - and not session.get("queued_prompts") - ): - prev = existing["text"] - existing["text"] = f"{prev}\n\n{text}" if prev and text else (prev or text) - return - if existing: + if submitted_at is not None: + queued["submitted_at"] = submitted_at + if message_id is not None: + queued["message_id"] = message_id + if session.get("queued_prompt"): session.setdefault("queued_prompts", []).append(queued) - return - session["queued_prompt"] = queued + else: + session["queued_prompt"] = queued def _sanitize_queued_entry_vs_inflight_user( @@ -9636,6 +9655,49 @@ def _drop_queued_duplicates_of_inflight_user(session: dict) -> None: session.pop("queued_prompts", None) +def _has_prompt_message_id(session: dict, message_id: str) -> bool: + """Return whether a stable client message ID is already owned by a turn. + + Callers hold ``history_lock``. The in-memory checks close the window before + early persistence, while the SessionDB check covers timeout/resume retries + after the original turn has completed. + """ + inflight = session.get("inflight_turn") + if isinstance(inflight, dict) and inflight.get("message_id") == message_id: + return True + + queued_items = [session.get("queued_prompt")] + pending = session.get("queued_prompts") + if isinstance(pending, list): + queued_items.extend(pending) + if any( + isinstance(item, dict) and item.get("message_id") == message_id + for item in queued_items + ): + return True + + for item in session.get("history") or []: + if not isinstance(item, dict): + continue + source_id = ( + item.get("platform_message_id") + or item.get("message_id") + or item.get("_source_message_id") + ) + if source_id is not None and str(source_id) == message_id: + return True + + agent = session.get("agent") + db = getattr(agent, "_session_db", None) + session_key = str(session.get("session_key") or "") + if db is not None and session_key and hasattr(db, "has_platform_message_id"): + try: + return bool(db.has_platform_message_id(session_key, message_id)) + except Exception: + pass + return False + + def _interrupt_busy_session(sid: str, session: dict, agent: Any) -> None: """Interrupt a busy turn without blocking the RPC reader or session lock. @@ -9669,10 +9731,52 @@ def interrupt() -> None: session["_busy_interrupt_pending"] = False threading.Thread(target=interrupt, daemon=True, name=f"busy-interrupt-{sid}").start() +def _rebind_session_transport( + session: dict, + transport: Any, + *, + message_id: str | None = None, + migrate_dead_queued: bool = False, +) -> None: + """Bind a live client and re-home only queue entries it now owns. + + Callers hold ``history_lock``. An explicit source-ID retry transfers that + one queued item to the retrying client. Resume/activate migrates each dead + queued transport independently while preserving live per-item FIFO routing. + """ + session["transport"] = transport + + queued_items = [session.get("queued_prompt")] + pending = session.get("queued_prompts") + if isinstance(pending, list): + queued_items.extend(pending) + for item in queued_items: + if not isinstance(item, dict): + continue + matches_source = ( + message_id is not None and item.get("message_id") == message_id + ) + if matches_source or ( + migrate_dead_queued and _transport_is_dead(item.get("transport")) + ): + item["transport"] = transport + + # ``inflight_turn`` has no separate transport slot: rebinding the session + # above transfers any matching in-flight turn atomically with its ID check. + # Do not manufacture one here; completed-history/DB duplicates also pass + # through this helper solely to keep future session events on the live client. def _handle_busy_submit( - rid, sid: str, session: dict, text: Any, transport: Any, queued: bool = False + rid, + sid: str, + session: dict, + text: Any, + transport: Any, + queued: bool = False, + *, + submitted_at: float | None = None, + message_id: str | None = None, ) -> dict | None: """Apply the ``display.busy_input_mode`` policy to a prompt that lands while a turn is in flight, instead of rejecting it with ``session busy``. @@ -9750,7 +9854,14 @@ def _handle_busy_submit( if image_paths: session["attached_images"] = image_paths + list(session.get("attached_images", [])) return None - _enqueue_prompt(session, text, transport, image_paths=image_paths) + _enqueue_prompt( + session, + text, + transport, + image_paths=image_paths, + submitted_at=submitted_at, + message_id=message_id, + ) session["last_active"] = time.time() # Attachments need a separate model invocation. Queue them without @@ -9769,12 +9880,15 @@ def _handle_busy_submit( return _ok(rid, {"status": "queued"}) + + def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: - """Fire a queued next-turn prompt if one is waiting and the session is idle. + """Dispatch the FIFO head once when the session becomes idle. - Returns True if a queued prompt was dispatched (the caller should then skip - lower-priority follow-ups this cycle — the user's message wins). Mirrors the - claim-under-lock pattern used by the goal-continuation re-fire. + Returns True after claiming an item, so the caller skips lower-priority + follow-ups for this cycle. The head is advanced under ``history_lock``; + synchronous dispatch failure restores the claimed item ahead of arrivals + that raced the failed attempt. """ with session["history_lock"]: if session.get("_closing"): @@ -9790,6 +9904,17 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: session["running"] = True if queued.get("transport") is not None: session["transport"] = queued["transport"] + _start_inflight_turn( + session, + queued["text"], + submitted_at=queued.get("submitted_at"), + message_id=queued.get("message_id"), + ) + run_kwargs = { + key: queued[key] + for key in ("submitted_at", "message_id") + if queued.get(key) is not None + } use_compute_host = _session_uses_compute_host(session) with session["history_lock"]: if int(session.get("_queued_prompt_generation", 0)) != queue_generation: @@ -9821,10 +9946,13 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: queued["text"], image_paths=queued["image_paths"], queued_prompt_generation=queue_generation, + **run_kwargs, ) else: resp = _submit_prompt_to_compute_host( - rid, sid, session, queued["text"], queued_prompt_generation=queue_generation + rid, sid, session, queued["text"], + queued_prompt_generation=queue_generation, + **run_kwargs, ) if resp.get("error"): message = str(((resp.get("error") or {}).get("message")) or "queued prompt failed") @@ -9842,6 +9970,7 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: queued["text"], image_paths=queued["image_paths"], queued_prompt_generation=queue_generation, + **run_kwargs, ) else: _run_prompt_submit( @@ -9850,6 +9979,7 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: session, queued["text"], queued_prompt_generation=queue_generation, + **run_kwargs, ) except Exception as exc: print( @@ -9858,18 +9988,37 @@ def _drain_queued_prompt(rid, sid: str, session: dict) -> bool: file=sys.stderr, ) with session["history_lock"]: + # Arrivals that raced the failed attempt sit in the FIFO tail — + # restore the exact pre-dispatch order (failed item back at head, + # displaced head back ahead of the arrivals) and stop here: the + # failed item must not be retried in a tight loop while later + # submissions wait behind it. + arrivals_raced = bool(session.get("queued_prompts")) + displaced_head = session.get("queued_prompt") + if arrivals_raced: + session.setdefault("queued_prompts", []).insert(0, displaced_head) + session["queued_prompt"] = queued + elif displaced_head is None: + # Nothing behind the failed item — restore it so the user's + # message is not silently dropped by the failed dispatch. + session["queued_prompt"] = queued + _clear_inflight_turn(session) session["running"] = False dispatch_failed = True if dispatch_failed: with session["history_lock"]: - drain_next = bool(session.get("queued_prompt")) and not session.get( - "_turn_cancel_requested" + drain_next = ( + not arrivals_raced + and displaced_head is not None + and not session.get("_turn_cancel_requested") ) if drain_next: _drain_queued_prompt(rid, sid, session) return True + + def _inflight_snapshot(session: dict) -> dict | None: turn = session.get("inflight_turn") if not isinstance(turn, dict): @@ -10452,7 +10601,11 @@ def _live_session_payload( if cols is not None: session["cols"] = cols if transport is not None: - session["transport"] = transport + _rebind_session_transport( + session, + transport, + migrate_dead_queued=True, + ) # Track every transport that has shown this session (multi-window: # pop-out windows each resume the same sid). The last viewer # becomes the transport on the disconnect path so closing a @@ -11128,6 +11281,751 @@ def _serialize_subscription_preview(p) -> dict: } +@method("subscription.preview") +def _(rid, params: dict) -> dict: + """POST /api/billing/subscription/preview → serialized quote or typed error. + + params: {subscription_type_id: str}. Chargeless effect quote. Requires + billing:manage (live Stripe calls + amounts), so a 403 → insufficient_scope + drives the device step-up exactly like the mutations. + """ + from agent.subscription_view import subscription_change_preview_from_payload + from hermes_cli.nous_billing import BillingError, post_subscription_preview + + tier_id = params.get("subscription_type_id") + if not tier_id: + return _ok(rid, {"ok": False, "error": "invalid_request", "message": "subscription_type_id is required"}) + try: + preview = subscription_change_preview_from_payload( + post_subscription_preview(subscription_type_id=tier_id) + ) + return _ok(rid, _serialize_subscription_preview(preview)) + except BillingError as exc: + return _ok(rid, _serialize_billing_error(exc)) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc)}) + + +@method("subscription.change") +def _(rid, params: dict) -> dict: + """PUT /api/billing/subscription/pending-change → {ok, message} or typed error. + + params: {subscription_type_id?: str, cancel?: bool}. Schedules a downgrade / + same-price change OR a cancellation at period end (chargeless). Requires + billing:manage. + """ + from hermes_cli.nous_billing import BillingError, put_subscription_pending_change + + cancel = bool(params.get("cancel")) + tier_id = params.get("subscription_type_id") + if not cancel and not tier_id: + return _ok(rid, {"ok": False, "error": "invalid_request", "message": "subscription_type_id or cancel is required"}) + try: + result = put_subscription_pending_change(subscription_type_id=tier_id, cancel=cancel) + return _ok(rid, {"ok": True, "message": result.get("message"), "payload": result}) + except BillingError as exc: + return _ok(rid, _serialize_billing_error(exc)) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc)}) + + +@method("subscription.resume") +def _(rid, params: dict) -> dict: + """DELETE /api/billing/subscription/pending-change → {ok, message} or typed error. + + Clears a scheduled downgrade or cancellation (resume / undo). Chargeless, but it + re-enables recurring spend → requires billing:manage and honors the kill-switch. + """ + from hermes_cli.nous_billing import BillingError, delete_subscription_pending_change + + try: + result = delete_subscription_pending_change() + return _ok(rid, {"ok": True, "message": result.get("message"), "payload": result}) + except BillingError as exc: + return _ok(rid, _serialize_billing_error(exc)) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc)}) + + +@method("subscription.upgrade") +def _(rid, params: dict) -> dict: + """POST /api/billing/subscription/upgrade → {ok, status, ...} or typed error. + + params: {subscription_type_id: str, idempotency_key?: str}. The single money + route: prorate + charge the card on the subscription + flip the plan. SCA / + decline come back as status requires_action / payment_failed with a recovery_url + to finish in the portal. The idempotency key is minted if absent and echoed so + the TUI reuses it on retry of the SAME upgrade. Requires billing:manage. + """ + from agent.billing_view import new_idempotency_key + from hermes_cli.nous_billing import BillingError, post_subscription_upgrade + + tier_id = params.get("subscription_type_id") + if not tier_id: + return _ok(rid, {"ok": False, "error": "invalid_request", "message": "subscription_type_id is required"}) + key = params.get("idempotency_key") or new_idempotency_key() + try: + result = post_subscription_upgrade(subscription_type_id=tier_id, idempotency_key=key) + return _ok( + rid, + { + "ok": True, + "status": result.get("status"), + "target_tier_name": result.get("targetTierName"), + "recovery_url": result.get("recoveryUrl"), + "reason": result.get("reason"), + "idempotency_key": key, + }, + ) + except BillingError as exc: + env = _serialize_billing_error(exc) + env["idempotency_key"] = key # so the TUI can reuse on retry + return _ok(rid, env) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc), "idempotency_key": key}) + + + +@method("billing.charge") +def _(rid, params: dict) -> dict: + """POST /api/billing/charge → {ok, chargeId} or a typed error envelope. + + params: {amount_usd: str|number, idempotency_key?: str}. If no key is + supplied, the server-side core mints a fresh one and returns it so the TUI can + reuse it on retry of the SAME purchase. + """ + from hermes_cli.nous_billing import BillingError, post_charge + from agent.billing_view import new_idempotency_key + + amount = params.get("amount_usd") + if amount is None: + return _ok(rid, {"ok": False, "error": "invalid_request", "message": "amount_usd is required"}) + key = params.get("idempotency_key") or new_idempotency_key() + try: + result = post_charge(amount_usd=amount, idempotency_key=key) + return _ok(rid, {"ok": True, "charge_id": result.get("chargeId"), "idempotency_key": key}) + except BillingError as exc: + env = _serialize_billing_error(exc) + env["idempotency_key"] = key # so the TUI can reuse on retry + return _ok(rid, env) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc), "idempotency_key": key}) + + +@method("billing.charge_status") +def _(rid, params: dict) -> dict: + """GET /api/billing/charge/{id} → {ok, status, ...} or typed error. + + The poll. Caller drives the 2s/5-min cadence; this is a single status read. + """ + from hermes_cli.nous_billing import BillingError, get_charge_status + + charge_id = params.get("charge_id") + if not charge_id: + return _ok(rid, {"ok": False, "error": "invalid_charge_id", "message": "charge_id is required"}) + try: + result = get_charge_status(charge_id) + return _ok( + rid, + { + "ok": True, + "status": result.get("status"), + "amount_usd": result.get("amountUsd"), + "settled_at": result.get("settledAt"), + "reason": result.get("reason"), + }, + ) + except BillingError as exc: + return _ok(rid, _serialize_billing_error(exc)) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc)}) + + +@method("billing.auto_reload") +def _(rid, params: dict) -> dict: + """PATCH /api/billing/auto-top-up → {ok:true} or typed error (Screen 2). + + params: {enabled: bool, threshold: number, top_up_amount: number}. + """ + from hermes_cli.nous_billing import BillingError, patch_auto_top_up + + try: + enabled = bool(params.get("enabled")) + threshold = params.get("threshold") + top_up_amount = params.get("top_up_amount") + if threshold is None or top_up_amount is None: + return _ok(rid, {"ok": False, "error": "invalid_request", "message": "threshold and top_up_amount are required"}) + patch_auto_top_up(enabled=enabled, threshold=threshold, top_up_amount=top_up_amount) + return _ok(rid, {"ok": True}) + except BillingError as exc: + return _ok(rid, _serialize_billing_error(exc)) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc)}) + + +@method("billing.step_up") +def _(rid, params: dict) -> dict: + """Run the lazy billing:manage step-up device flow → {ok, granted}. + + Triggered by the TUI after a billing call returns error=insufficient_scope. + Returns granted:false when the server silently downscopes (non-admin / unticked). + + Runs on the thread pool (in _LONG_HANDLERS): the device flow blocks for the + whole device-code lifetime (minutes), so it must not stall the main stdin loop. + The verification URL/code reach the TUI via an out-of-band ``billing.step_up. + verification`` event (a plain print would be dropped by the JSON-RPC stdout + pipe), and the browser is opened TUI-side via openExternalUrl — never with the + gateway's headless webbrowser.open (hence open_browser=False). + """ + sid = params.get("session_id") or "" + try: + from hermes_cli.auth import step_up_nous_billing_scope + from hermes_cli.nous_billing import BillingError + + def _on_verification(url: str, code: str) -> None: + _emit( + "billing.step_up.verification", + sid, + {"verification_url": url, "user_code": code}, + ) + + granted = step_up_nous_billing_scope( + open_browser=False, on_verification=_on_verification + ) + return _ok(rid, {"ok": True, "granted": bool(granted)}) + except BillingError as exc: + # Route typed billing errors (e.g. session_revoked when the token expires + # mid-device-flow) through the shared spine like the other write handlers, + # so the TUI maps them to the right copy instead of a generic failure. + env = _serialize_billing_error(exc) + env["granted"] = False + return _ok(rid, env) + except Exception as exc: + return _ok(rid, {"ok": False, "error": "error", "message": str(exc), "granted": False}) + + +@method("session.status") +def _(rid, params: dict) -> dict: + session, err = _sess_nowait(params, rid) + if err: + return err + + from hermes_constants import display_hermes_home + + key = session.get("session_key") or params.get("session_id") or "" + agent = session.get("agent") + meta = {} + db = _get_db() + if db and key: + try: + meta = db.get_session(key) or {} + except Exception: + meta = {} + + def _dt(value, fallback: datetime | None = None) -> datetime: + if value: + try: + return datetime.fromtimestamp(float(value)) + except Exception: + pass + return fallback or datetime.now() + + created = _dt(meta.get("started_at")) + updated = created + for field in ("updated_at", "last_updated_at", "last_activity_at"): + if meta.get(field): + updated = _dt(meta.get(field), created) + break + + mirror = _metadata_mirror(session) + usage = _session_usage_snapshot(session) + provider = getattr(agent, "provider", None) or mirror.get("provider") or "unknown" + model = getattr(agent, "model", None) or mirror.get("model") or "(unknown)" + project = _project_info_for_cwd(_display_session_cwd(session)) + lines = [ + "Hermes TUI Status", + "", + f"Session ID: {key}", + f"Path: {display_hermes_home()}", + ] + if project: + lines.append(f"Project: {project['name']}") + title = (meta.get("title") or "").strip() + if title: + lines.append(f"Title: {title}") + lines.extend( + [ + f"Model: {model} ({provider})", + f"Created: {created.strftime('%Y-%m-%d %H:%M')}", + f"Last Activity: {updated.strftime('%Y-%m-%d %H:%M')}", + f"Tokens: {int(usage.get('total') or 0):,}", + f"Agent Running: {'Yes' if session.get('running') else 'No'}", + ] + ) + return _ok(rid, {"output": "\n".join(lines)}) + + +@method("session.history") +def _(rid, params: dict) -> dict: + session, err = _sess_nowait(params, rid) + if err: + return err + history = list(session.get("history", [])) + db = _get_db() + if db is not None and session.get("session_key"): + try: + history = db.get_messages_as_conversation( + session["session_key"], include_ancestors=True + ) + except Exception: + pass + return _ok( + rid, + { + "count": len(history), + "messages": _history_to_messages(history), + }, + ) + + +@method("session.undo") +def _(rid, params: dict) -> dict: + session, err = _sess(params, rid) + if err: + return err + # Reject during an in-flight turn. If we mutated history while + # the agent thread is running, prompt.submit's post-run history + # write would either clobber the undo (version matches) or + # silently drop the agent's output (version mismatch, see below). + # Neither is what the user wants — make them /interrupt first. + if session.get("running"): + return _err( + rid, 4009, "session busy — /interrupt the current turn before /undo" + ) + removed = 0 + with session["history_lock"]: + history = session.get("history", []) + while history and history[-1].get("role") in {"assistant", "tool"}: + history.pop() + removed += 1 + if history and history[-1].get("role") == "user": + history.pop() + removed += 1 + if removed: + session["history_version"] = int(session.get("history_version", 0)) + 1 + return _ok(rid, {"removed": removed}) + + +@method("session.compress") +def _(rid, params: dict) -> dict: + session, err = _sess_nowait(params, rid) + if err: + return err + assert session is not None + if _session_uses_compute_host(session): + sid = str(params.get("session_id") or "") + focus_topic = str(params.get("focus_topic", "") or "").strip() + command = "/compress" + (f" {focus_topic}" if focus_topic else "") + try: + ack = _send_compute_host_control( + sid, + route_name="session.compress", + command=command, + wait=True, + timeout=120.0, + ) + except Exception as exc: + return _err(rid, 5019, f"compute-host compress failed: {exc}") + if ack.get("type") in {"control.error", "error"}: + return _err(rid, 4009, str(ack.get("message") or "compute-host compress failed")) + _apply_compute_host_metadata_mirror(session, ack) + host_result = ack.get("result") + if isinstance(host_result, dict): + # The host owns the isolated session's agent/history, so preserve + # its structured compression result verbatim. In particular this + # carries `status: aborted` and `summary.aborted`; flattening the + # old text-only acknowledgement made Desktop show aborted work as a + # success toast. + return _ok(rid, {**host_result, "turn_isolation": True}) + host_info = ack.get("session_info") if isinstance(ack.get("session_info"), dict) else {} + host_messages = _history_to_messages(ack.get("messages")) if isinstance(ack.get("messages"), list) else [] + # `messages` is returned at top level for the desktop transcript + # replacement. Keep the host acknowledgement metadata, but do not send + # the same (potentially large) transcript a second time inside it. + host_ack = {key: value for key, value in ack.items() if key != "messages"} + return _ok( + rid, + { + "status": "compressed", + "turn_isolation": True, + "host_ack": host_ack, + "info": host_info, + "messages": host_messages, + "usage": host_info.get("usage") if isinstance(host_info.get("usage"), dict) else {}, + }, + ) + session, err = _sess(params, rid) + if err: + return err + if session.get("running"): + return _err( + rid, 4009, "session busy — /interrupt the current turn before /compress" + ) + from agent.conversation_compression import ( + finalize_context_engine_compression_notification, + ) + + sid = params.get("session_id", "") + focus_topic = str(params.get("focus_topic", "") or "").strip() + try: + from agent.manual_compression_feedback import summarize_manual_compression + from agent.model_metadata import estimate_request_tokens_rough + + with session["history_lock"]: + before_messages = list(session.get("history", [])) + history_version = int(session.get("history_version", 0)) + before_count = len(before_messages) + _agent = session["agent"] + _sys_prompt = getattr(_agent, "_cached_system_prompt", "") or "" + _tools = getattr(_agent, "tools", None) or None + before_tokens = ( + estimate_request_tokens_rough( + before_messages, system_prompt=_sys_prompt, tools=_tools + ) + if before_count + else 0 + ) + + if before_count >= 4: + focus_suffix = f', focus: "{focus_topic}"' if focus_topic else "" + _status_update( + sid, + "compressing", + f"⠋ compressing {before_count} messages " + f"(~{before_tokens:,} tok){focus_suffix}…", + ) + + try: + removed, usage = _compress_session_history( + session, + focus_topic, + approx_tokens=before_tokens, + before_messages=before_messages, + history_version=history_version, + ) + with session["history_lock"]: + messages = list(session.get("history", [])) + after_count = len(messages) + # Re-read system prompt + tools after compression — _compress_context + # may have rebuilt the system prompt (_cached_system_prompt=None). + _sys_prompt_after = ( + getattr(_agent, "_cached_system_prompt", "") or _sys_prompt + ) + _tools_after = getattr(_agent, "tools", None) or _tools + after_tokens = ( + estimate_request_tokens_rough( + messages, + system_prompt=_sys_prompt_after, + tools=_tools_after, + ) + if after_count + else 0 + ) + agent = session["agent"] + _sync_session_key_after_compress(sid, session) + summary = summarize_manual_compression( + before_messages, + messages, + before_tokens, + after_tokens, + compression_state=getattr(agent, "context_compressor", None), + ) + info = _session_info(agent, session) + _emit("session.info", sid, info) + finalize_context_engine_compression_notification( + agent, + committed=True, + ) + return _ok( + rid, + { + "status": "aborted" if summary["aborted"] else "compressed", + "removed": removed, + "before_messages": before_count, + "after_messages": after_count, + "before_tokens": before_tokens, + "after_tokens": after_tokens, + "summary": summary, + "usage": usage, + "info": info, + # Keep this identical to session.resume / session.history: + # raw tool results can contain large or sensitive payloads + # that belong in persisted history, not the transcript + # replacement response. + "messages": _history_to_messages(messages), + }, + ) + finally: + # Always clear the pinned compressing status so the bar + # reverts to neutral whether compaction succeeded, was a + # no-op, or raised. + _status_update(sid, "ready") + except CompressionLockHeld as e: + _status_update(sid, "ready") + from agent.manual_compression_feedback import ( + describe_compression_lock_skip, + ) + return _ok(rid, { + "compressed": False, + "lock_held": True, + "message": describe_compression_lock_skip(e.holder), + }) + except Exception as e: + finalize_context_engine_compression_notification( + session["agent"], + committed=False, + ) + return _err(rid, 5005, str(e)) + + +@method("session.save") +def _(rid, params: dict) -> dict: + session, err = _sess(params, rid) + if err: + return err + + if _session_uses_compute_host(session): + sid = str(params.get("session_id") or "") + try: + ack = _send_compute_host_control( + sid, + route_name="session.save", + wait=True, + ) + except Exception as exc: + return _err(rid, 5011, f"compute-host session save failed: {exc}") + if ack.get("type") in {"control.error", "error"}: + return _err(rid, 5011, str(ack.get("message") or "compute-host session save failed")) + result = ack.get("result") + if not isinstance(result, dict): + return _err(rid, 5011, "compute-host session save returned an invalid response") + return _ok(rid, result) + + agent = session["agent"] + # Mirror the classic CLI /save: snapshot under the Hermes profile home + # (~/.hermes/sessions/saved/) rather than the project/workspace CWD, and + # include the system prompt so the export matches the dashboard save. + saved_dir = get_hermes_home() / "sessions" / "saved" + try: + saved_dir.mkdir(parents=True, exist_ok=True) + except Exception as e: + return _err(rid, 5011, f"failed to create save directory {saved_dir}: {e}") + + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + path = saved_dir / f"hermes_conversation_{timestamp}.json" + + with session["history_lock"]: + messages = list(session.get("history", [])) + + session_id = getattr(agent, "session_id", None) or session.get("session_key") or "" + # Prefer the agent's session_start datetime (matches the classic CLI export); + # fall back to the gateway session's created_at timestamp. + agent_start = getattr(agent, "session_start", None) + if isinstance(agent_start, datetime): + session_start = agent_start.isoformat() + else: + created_at = session.get("created_at") + session_start = ( + datetime.fromtimestamp(created_at).isoformat() + if isinstance(created_at, (int, float)) + else "" + ) + + try: + with open(path, "w", encoding="utf-8") as f: + json.dump( + { + "model": getattr(agent, "model", ""), + "session_id": session_id, + "session_start": session_start, + "system_prompt": getattr(agent, "_cached_system_prompt", "") or "", + "messages": messages, + }, + f, + indent=2, + ensure_ascii=False, + ) + return _ok(rid, {"file": str(path)}) + except Exception as e: + return _err(rid, 5011, str(e)) + + +@method("session.close") +def _(rid, params: dict) -> dict: + sid = params.get("session_id", "") + # Serialize only the ownership claim against session.resume / the orphan + # reaper. Finalization may run arbitrary plugin/agent cleanup and must not + # keep every unrelated session.resume waiting behind it. + with _session_resume_lock: + session = _pop_session_by_id(sid) + closed = _teardown_popped_session(session, end_reason="tui_close") + return _ok(rid, {"closed": closed}) + + +@method("session.branch") +def _(rid, params: dict) -> dict: + session, err = _sess(params, rid) + if err: + return err + db = _get_db() + if db is None: + return _db_unavailable_error(rid, code=5008) + old_key = session["session_key"] + with session["history_lock"]: + history = [dict(msg) for msg in session.get("history", [])] + if not history: + return _err(rid, 4008, "nothing to branch — send a message first") + new_key = _new_session_key() + new_sid = uuid.uuid4().hex[:8] + source = _session_source(session) + lease, limit_message = _claim_active_session_slot( + new_key, live_session_id=new_sid, surface=source + ) + if limit_message is not None: + return _err(rid, 4090, limit_message) + branch_name = params.get("name", "") + try: + if branch_name: + title = branch_name + else: + current = db.get_session_title(old_key) or "branch" + title = ( + db.get_next_title_in_lineage(current) + if hasattr(db, "get_next_title_in_lineage") + else f"{current} (branch)" + ) + db.create_session( + new_key, + source=source, + model=_resolve_model(), + # Stable _branched_from marker so list_sessions_rich() keeps the + # branch visible in /resume and /sessions. The TUI branch leaves + # the parent live (no end_reason='branched'), so the legacy + # end_reason heuristic never matches it — the marker is the only + # thing that surfaces TUI branches. See issue #20856. + model_config={"_branched_from": old_key}, + parent_session_id=old_key, + cwd=_session_cwd(session), + ) + for msg in history: + db.append_message( + session_id=new_key, + role=msg.get("role", "user"), + content=msg.get("content"), + timestamp=msg.get("timestamp"), + ) + db.set_session_title(new_key, title) + except Exception as e: + if lease is not None: + lease.release() + return _err(rid, 5008, f"branch failed: {e}") + try: + tokens = _set_session_context(new_key) + try: + agent = _make_agent( + new_sid, + new_key, + session_id=new_key, + platform_override=source, + ) + finally: + _clear_session_context(tokens) + _init_session( + new_sid, + new_key, + agent, + list(history), + cols=session.get("cols", 80), + source=source, + ) + if new_sid in _sessions: + _sessions[new_sid]["active_session_lease"] = lease + except Exception as e: + if lease is not None: + lease.release() + return _err(rid, 5000, f"agent init failed on branch: {e}") + return _ok(rid, {"session_id": new_sid, "title": title, "parent": old_key}) + + +@method("session.interrupt") +def _(rid, params: dict) -> dict: + # Keypress barge-in: stopping the turn also silences its streaming TTS + # (voice is process-global, so no per-session scoping is needed). + _tts_stream_stop() + session, err = _sess_nowait(params, rid) + if err: + return err + if _session_uses_compute_host(session): + sid = str(params.get("session_id") or "") + if session.get("running"): + try: + _get_compute_host_supervisor().interrupt(sid, request_id=f"interrupt-{rid}") + except Exception as exc: + return _err(rid, 5019, f"compute-host interrupt failed: {exc}") + with session["history_lock"]: + session["_turn_cancel_requested"] = True + session["queued_prompt"] = None + _clear_pending(sid) + try: + from tools.approval import resolve_gateway_approval + + resolve_gateway_approval(session["session_key"], "deny", resolve_all=True) + except Exception: + pass + return _ok(rid, {"status": "interrupted", "turn_isolation": True}) + session, err = _sess(params, rid) + if err: + return err + # Safety net: if the turn's run thread is already gone but `running` stayed + # stuck (a crash/desync that skipped the run loop's `finally`), force-clear it + # so the session can't be permanently bricked at 4009 "session busy" — every + # send/restore/resume would otherwise reject until a full backend restart. + # Always tell the agent to interrupt when the session claims a run is active: + # stale flags are cleared below, and fresh turns clear the interrupt flag at + # entry. This keeps a stale/missing thread handle from making Stop a no-op. + run_thread = session.get("_run_thread") + run_thread_alive = run_thread is not None and run_thread.is_alive() + should_interrupt = bool(session.get("running")) + if should_interrupt and hasattr(session["agent"], "interrupt"): + session["agent"].interrupt() + with session["history_lock"]: + session["_turn_cancel_requested"] = True + session["queued_prompt"] = None + session["queued_prompts"] = [] + if not run_thread_alive: + with session["history_lock"]: + if session.get("running"): + session["running"] = False + _clear_inflight_turn(session) + + # Stop = stop the TURN (cooperative interrupt above also kills the in-flight + # foreground subprocess). Background processes the agent started (dev servers, + # watchers) are intentionally left running — kill those individually with the + # "x" on the task row (process.kill). Don't reap them here. + # Scope the pending-prompt release to THIS session. A global + # _clear_pending() would collaterally cancel clarify/sudo/secret + # prompts on unrelated sessions sharing the same tui_gateway + # process, silently resolving them to empty strings. + _clear_pending(params.get("session_id", "")) + try: + from tools.approval import resolve_gateway_approval + + resolve_gateway_approval(session["session_key"], "deny", resolve_all=True) + except Exception: + pass + return _ok(rid, {"status": "interrupted"}) + + # ── Delegation: subagent tree observability + controls ─────────────── # Powers the TUI's /agents overlay (see ui-tui/src/components/agentsOverlay). # The registry lives in tools/delegate_tool — these handlers are thin @@ -12204,6 +13102,8 @@ def _run_prompt_submit( session: dict, text: Any, *, + submitted_at: float | None = None, + message_id: str | None = None, display_kind: str | None = None, display_metadata: dict | None = None, image_paths: list[str] | None = None, @@ -12224,11 +13124,15 @@ def _run_prompt_submit( session["attached_images"] = [] else: images = list(image_paths) + if submitted_at is not None: + session["_pending_submitted_at"] = submitted_at + if message_id is not None: + session["_pending_message_id"] = message_id inflight = session.get("inflight_turn") # A retained failed turn (see _fail_inflight_turn) is a stale leftover # by the time a new turn starts — replace it, never append onto it. if not isinstance(inflight, dict) or inflight.get("status") == "error": - _start_inflight_turn(session, text) + _start_inflight_turn(session, text, message_id=message_id) agent = session["agent"] if hasattr(agent, "clear_interrupt"): try: @@ -12534,6 +13438,16 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: # the same value is a no-op. try: _run_params = inspect.signature(agent.run_conversation).parameters + _accepts_kwargs = any( + p.kind == inspect.Parameter.VAR_KEYWORD + for p in _run_params.values() + ) + if _accepts_kwargs or "persist_user_timestamp" in _run_params: + if submitted_at is not None: + run_kwargs["persist_user_timestamp"] = submitted_at + if _accepts_kwargs or "persist_user_message_id" in _run_params: + if message_id is not None: + run_kwargs["persist_user_message_id"] = message_id except (TypeError, ValueError): _run_params = {} if "task_id" in _run_params: