fix(webui): persist user messages for refresh

This commit is contained in:
chengyongru
2026-06-05 16:13:51 +08:00
committed by Xubin Ren
parent 3da68ac7fe
commit 710d00a179
5 changed files with 427 additions and 2 deletions
+170
View File
@@ -40,6 +40,15 @@ _FILE_EDIT_TOOL_NAMES: frozenset[str] = frozenset({
"edit_file",
"apply_patch",
})
_TURN_DISPLAY_EVENTS: frozenset[str] = frozenset({
"reasoning_delta",
"reasoning_end",
"delta",
"stream_end",
"message",
"file_edit",
"turn_end",
})
def rewrite_local_markdown_images(
@@ -155,6 +164,165 @@ def delete_webui_transcript(session_key: str) -> bool:
return False
def build_user_transcript_event(
chat_id: str,
text: str,
*,
media_paths: list[Any] | None = None,
cli_apps: list[Any] | None = None,
mcp_presets: list[Any] | None = None,
) -> dict[str, Any] | None:
paths = [str(path) for path in (media_paths or []) if path]
if not text and not paths:
return None
event: dict[str, Any] = {
"event": "user",
"chat_id": chat_id,
"text": text,
}
if paths:
event["media_paths"] = paths
apps = [dict(app) for app in (cli_apps or []) if isinstance(app, Mapping)]
if apps:
event["cli_apps"] = apps
presets = [dict(preset) for preset in (mcp_presets or []) if isinstance(preset, Mapping)]
if presets:
event["mcp_presets"] = presets
return event
def _session_user_event(
session_key: str,
message: dict[str, Any],
) -> dict[str, Any] | None:
if message.get("role") != "user":
return None
content = message.get("content")
text = content if isinstance(content, str) else ""
media = message.get("media")
cli_apps = message.get("cli_apps")
mcp_presets = message.get("mcp_presets")
chat_id = session_key.split(":", 1)[1] if ":" in session_key else session_key
return build_user_transcript_event(
chat_id,
text,
media_paths=media if isinstance(media, list) else None,
cli_apps=cli_apps if isinstance(cli_apps, list) else None,
mcp_presets=mcp_presets if isinstance(mcp_presets, list) else None,
)
def _assistant_text_signature(value: Any) -> str:
return value.strip() if isinstance(value, str) else ""
def _session_backfill_turns(
session_key: str,
session_messages: list[dict[str, Any]],
) -> list[tuple[dict[str, Any], tuple[str, ...]]]:
turns: list[tuple[dict[str, Any], tuple[str, ...]]] = []
current_user: dict[str, Any] | None = None
assistant_texts: list[str] = []
def flush() -> None:
if current_user is None:
return
signature = tuple(text for text in assistant_texts if text)
if signature:
turns.append((current_user, signature))
for message in session_messages:
role = message.get("role")
if role == "user":
flush()
current_user = _session_user_event(session_key, message)
assistant_texts = []
continue
if role == "assistant" and current_user is not None:
text = _assistant_text_signature(message.get("content"))
if text:
assistant_texts.append(text)
flush()
return turns
def _split_transcript_turns(lines: list[dict[str, Any]]) -> list[list[dict[str, Any]]]:
turns: list[list[dict[str, Any]]] = []
current: list[dict[str, Any]] = []
for rec in lines:
current.append(rec)
if rec.get("event") == "turn_end":
turns.append(current)
current = []
if current:
turns.append(current)
return turns
def _transcript_turn_signature(records: list[dict[str, Any]]) -> tuple[str, ...]:
texts: list[str] = []
for message in replay_transcript_to_ui_messages(records):
if message.get("role") != "assistant" or message.get("kind") == "trace":
continue
text = _assistant_text_signature(message.get("content"))
if text:
texts.append(text)
return tuple(texts)
def _find_unique_session_turn(
session_turns: list[tuple[dict[str, Any], tuple[str, ...]]],
signature: tuple[str, ...],
start: int,
) -> int | None:
if not signature:
return None
found: int | None = None
for index in range(start, len(session_turns)):
if session_turns[index][1] != signature:
continue
if found is not None:
return None
found = index
return found
def _with_backfilled_user(
records: list[dict[str, Any]],
user_event: dict[str, Any],
) -> list[dict[str, Any]]:
for index, rec in enumerate(records):
if rec.get("event") in _TURN_DISPLAY_EVENTS:
return [*records[:index], dict(user_event), *records[index:]]
return records
def inject_missing_user_events_from_session(
session_key: str,
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Backfill user rows for legacy WebUI transcripts that only stored assistant streams."""
if not lines or not session_messages:
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
out: list[dict[str, Any]] = []
session_cursor = 0
for turn in _split_transcript_turns(lines):
has_user = any(rec.get("event") == "user" for rec in turn)
signature = _transcript_turn_signature(turn)
match_index = _find_unique_session_turn(session_turns, signature, session_cursor)
if match_index is None:
out.extend(turn)
continue
out.extend(turn if has_user else _with_backfilled_user(turn, session_turns[match_index][0]))
session_cursor = match_index + 1
return out
def _format_tool_call_trace(call: Any) -> str | None:
if not call or not isinstance(call, dict):
return None
@@ -904,11 +1072,13 @@ def build_webui_thread_response(
augment_user_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
augment_assistant_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
augment_assistant_text: Callable[[str], str] | None = None,
session_messages: list[dict[str, Any]] | None = None,
) -> dict[str, Any] | None:
"""Return a payload compatible with ``WebuiThreadPersistedPayload``."""
lines = read_transcript_lines(session_key)
if not lines:
return None
lines = inject_missing_user_events_from_session(session_key, lines, session_messages)
msgs = replay_transcript_to_ui_messages(
lines,
augment_user_media=augment_user_media,