Stabilize WS chat streaming and failure handling

Use inactivity-based WS send timeout, avoid end-of-turn message list repaint, add WS heartbeat during turns, improve MiniMax tool replay compatibility, and retry tool replay protocol mismatch errors.

Made-with: Cursor
This commit is contained in:
oliver 2026-04-26 23:25:15 +08:00
parent 50c8bd19c0
commit 768f175c3e
5 changed files with 245 additions and 40 deletions

View file

@ -2649,10 +2649,6 @@ async function renderChatUi() {
const payload = msg.payload || {}; const payload = msg.payload || {};
onEvent({ event: String(msg.event || ""), payload }); onEvent({ event: String(msg.event || ""), payload });
const d = msg.event === "chat" ? { phase: String(payload.state || "") } : {}; const d = msg.event === "chat" ? { phase: String(payload.state || "") } : {};
if (msg.event === "session.message") {
const role = String(((payload.message || {}).role) || "").toLowerCase();
if (role === "assistant") return doneMeta;
}
const phase = String(d.phase || ""); const phase = String(d.phase || "");
if (phase === "final" || phase === "error" || phase === "aborted") { if (phase === "final" || phase === "error" || phase === "aborted") {
return doneMeta; return doneMeta;
@ -2783,7 +2779,13 @@ async function renderChatUi() {
let last = null; let last = null;
for (let i = msgs.length - 1; i >= 0; i--) { for (let i = msgs.length - 1; i >= 0; i--) {
const m = msgs[i]; const m = msgs[i];
if (String((m && m.role) || "").toLowerCase() === "assistant") { if (String((m && m.role) || "").toLowerCase() !== "assistant") continue;
// Recovery should only accept visible assistant body, not intermediate
// reasoning/tool-call events; otherwise we may terminate on a partial line.
const et = String((m && m.event_type) || "").trim().toLowerCase();
if (et && et !== "assistant_text" && et !== "assistant") continue;
if (!String((m && m.content) || "").trim()) continue;
{
last = m; last = m;
break; break;
} }
@ -3055,6 +3057,10 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
const abortController = new AbortController(); const abortController = new AbortController();
currentStreamAbortController = abortController; currentStreamAbortController = abortController;
currentAbortMeta = { sessionId: String(activeId || ""), runId: "" }; currentAbortMeta = { sessionId: String(activeId || ""), runId: "" };
// End-of-turn reload gating: avoid clearing/repainting messagesEl after we already
// finalized the stream bubble in-place (prevents end "flash").
let turnFinalized = false;
let turnStreamedEnough = false;
try { try {
const transport = new OclawWsChatTransport({ const transport = new OclawWsChatTransport({
tokenProvider: () => token, tokenProvider: () => token,
@ -3064,6 +3070,7 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
// Immediate assistant placeholder bubble with dynamic status. // Immediate assistant placeholder bubble with dynamic status.
ensureStreamBubble(); ensureStreamBubble();
_startDynamicStreamStatus(); _startDynamicStreamStatus();
let wsLastActivityAt = Date.now();
const wsSendPromise = transport.sendSessionSend({ const wsSendPromise = transport.sendSessionSend({
sessionId: activeId, sessionId: activeId,
text: userText, text: userText,
@ -3078,6 +3085,7 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
memoryMode: String(localStorage.getItem(CHAT_MEMORY_MODE_KEY) || "default"), memoryMode: String(localStorage.getItem(CHAT_MEMORY_MODE_KEY) || "default"),
signal: abortController.signal, signal: abortController.signal,
onEvent: async (frame) => { onEvent: async (frame) => {
wsLastActivityAt = Date.now();
const eventName = String((frame && frame.event) || ""); const eventName = String((frame && frame.event) || "");
const payload = (frame && frame.payload) || {}; const payload = (frame && frame.payload) || {};
if (eventName === "agent.event") { if (eventName === "agent.event") {
@ -3164,6 +3172,8 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
streamTextBuffer = ""; streamTextBuffer = "";
chatRunId = null; chatRunId = null;
const streamedEnough = hasRealStreamText || chatStreamSegments.length > 0 || sawWsChatEvent; const streamedEnough = hasRealStreamText || chatStreamSegments.length > 0 || sawWsChatEvent;
turnFinalized = true;
turnStreamedEnough = streamedEnough;
chatStreamSegments = []; chatStreamSegments = [];
// Avoid end-of-turn flash: only reload history when stream had no usable content. // Avoid end-of-turn flash: only reload history when stream had no usable content.
if (!streamedEnough) { if (!streamedEnough) {
@ -3180,6 +3190,8 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
streamDisplayShown = streamDisplayTarget; streamDisplayShown = streamDisplayTarget;
const ok = await appendFinalAssistant(payload.message, chatStream); const ok = await appendFinalAssistant(payload.message, chatStream);
if (!ok) _markStreamTerminal("error", "aborted"); if (!ok) _markStreamTerminal("error", "aborted");
turnFinalized = true;
turnStreamedEnough = hasRealStreamText || chatStreamSegments.length > 0 || sawWsChatEvent;
chatStream = ""; chatStream = "";
streamTextBuffer = ""; streamTextBuffer = "";
chatRunId = null; chatRunId = null;
@ -3190,6 +3202,8 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
sawWsTerminalEvent = true; sawWsTerminalEvent = true;
_stopDynamicStreamStatus(); _stopDynamicStreamStatus();
_markStreamTerminal("error", String(payload.errorMessage || "chat error")); _markStreamTerminal("error", String(payload.errorMessage || "chat error"));
turnFinalized = true;
turnStreamedEnough = hasRealStreamText || chatStreamSegments.length > 0 || sawWsChatEvent;
chatStream = ""; chatStream = "";
streamTextBuffer = ""; streamTextBuffer = "";
chatRunId = null; chatRunId = null;
@ -3198,12 +3212,25 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
} }
}, },
}); });
doneMeta = await Promise.race([ let wsWatchdog = 0;
wsSendPromise, const wsInactivityTimeoutPromise = new Promise((_, reject) => {
new Promise((_, reject) => { wsWatchdog = setInterval(() => {
setTimeout(() => reject(new Error(`ws_send_timeout:${WS_CHAT_SEND_TIMEOUT_MS}`)), WS_CHAT_SEND_TIMEOUT_MS); if (Date.now() - wsLastActivityAt > WS_CHAT_SEND_TIMEOUT_MS) {
}), if (wsWatchdog) {
]); clearInterval(wsWatchdog);
wsWatchdog = 0;
}
reject(new Error(`ws_send_timeout:${WS_CHAT_SEND_TIMEOUT_MS}`));
}
}, 1000);
});
const wsSendObserved = wsSendPromise.finally(() => {
if (wsWatchdog) {
clearInterval(wsWatchdog);
wsWatchdog = 0;
}
});
doneMeta = await Promise.race([wsSendObserved, wsInactivityTimeoutPromise]);
if (doneMeta && typeof doneMeta === "object") doneMeta.__transport = "ws"; if (doneMeta && typeof doneMeta === "object") doneMeta.__transport = "ws";
if (doneMeta && typeof doneMeta === "object") { if (doneMeta && typeof doneMeta === "object") {
const startToRunning = const startToRunning =
@ -3231,11 +3258,17 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
emsg.includes("closed") || emsg.includes("closed") ||
emsg.includes("timeout"); emsg.includes("timeout");
if (wsLikeFailure) { if (wsLikeFailure) {
const wsLikelyAlreadyProducedReply = sawWsTerminalEvent || hasRealStreamText || sawWsChatEvent; // Only short-circuit when a terminal event was already received.
if (wsLikelyAlreadyProducedReply) { // If we only saw partial deltas and WS drops, we must attempt recovery.
setTimeout(() => { if (sawWsTerminalEvent) {
loadMessagesForActive().catch(() => {}); // If the stream bubble was already finalized with usable content, do NOT
}, 350); // repaint messagesEl (prevents end-of-turn flash). Only reload when we
// have nothing usable and need to recover from persisted history.
if (!turnFinalized || !turnStreamedEnough) {
setTimeout(() => {
loadMessagesForActive().catch(() => {});
}, 350);
}
return doneMeta; return doneMeta;
} }
// WS timeout may happen while backend still computes and persists final reply. // WS timeout may happen while backend still computes and persists final reply.

View file

@ -260,18 +260,102 @@ async def run_agent_turn_via_bridge(
on_tool_ui=on_tool_ui, on_tool_ui=on_tool_ui,
) )
async def _heartbeat_during_turn(stop_evt: asyncio.Event) -> None:
# Keep WS active while model/tool pipeline is still preparing first token.
# Some proxies/load balancers close idle WS in 20-30s without downstream frames.
while not stop_evt.is_set():
try:
await asyncio.wait_for(stop_evt.wait(), timeout=4.0)
break
except asyncio.TimeoutError:
rid = run_id_holder.get("run_id") or ""
if not rid:
continue
with conn._abort_lock:
if rid in conn._aborted_run_ids:
continue
try:
await conn.emit_agent_event(
run_id=rid,
stream="lifecycle",
data={"phase": "running", "event": "heartbeat"},
)
except Exception:
# Best-effort keepalive; never fail the turn on heartbeat errors.
pass
hb_stop = asyncio.Event()
hb_task: asyncio.Task[Any] | None = asyncio.create_task(_heartbeat_during_turn(hb_stop))
try: try:
result = await asyncio.to_thread(_run_turn_sync) result = await asyncio.to_thread(_run_turn_sync)
except Exception as exc: except Exception as exc:
hb_stop.set()
if hb_task is not None:
try:
await hb_task
except Exception:
pass
err_text = str(exc or "agent_failed")
user_facing_error = (
"本轮执行失败:工具不可用或执行异常。"
f"\n\n错误信息:{err_text}\n\n请重试,或改用其它可用工具。"
)
fail_msg: dict[str, Any] = {
"role": "assistant",
"content": user_facing_error,
"timestamp": now_ms(),
}
try:
store.add_message(
session_id=str(session_id),
role="assistant",
content=str(user_facing_error),
turn_uuid=str(run_id_holder.get("run_id") or "") or None,
event_type="assistant_text",
)
except Exception:
pass
if send_response: if send_response:
await conn.send_res(req_id, ok=False, error=error_shape("UNAVAILABLE", str(exc or "agent_failed"))) await conn.send_res(
req_id,
ok=True,
payload={
"runId": run_id_holder.get("run_id") or "",
"acceptedAt": now_ms(),
"mode": "sync_direct",
"taskId": "",
"traceId": "",
"reply": user_facing_error,
"selectedSpecialist": str(p.get("specialist") or "generalist"),
"interactionMode": str(p.get("interaction_mode") or "comprehensive"),
"dispatchReason": "execution_failed",
"managerSelectedSpecialist": str(p.get("specialist") or "generalist"),
"requestedSpecialist": str(p.get("specialist") or "generalist"),
"dynamicAgentUsed": False,
"dynamicAgentName": "",
"relayPointerCount": 0,
"relayEnvelopePresent": False,
"relayEnvelopePointerCount": 0,
"relayTtlTurnCount": int(marker_turn_count),
"relayTtlSessionCount": int(marker_session_count),
"relayTtlKeepCount": int(marker_keep_count),
"status": "failed",
"error": err_text,
},
)
await conn.emit_agent_event( await conn.emit_agent_event(
run_id=run_id_holder.get("run_id") or "", run_id=run_id_holder.get("run_id") or "",
stream="lifecycle", stream="lifecycle",
# Always emit a terminal lifecycle event even on failure. # Always emit a terminal lifecycle event even on failure.
data={"phase": "end", "status": "error", "error": str(exc or "agent_failed")}, data={"phase": "end", "status": "error", "error": str(exc or "agent_failed")},
) )
await conn.emit_chat_event(run_id=run_id_holder.get("run_id") or "", state="error", error=str(exc or "agent_failed")) await conn.emit_chat_event(
run_id=run_id_holder.get("run_id") or "",
state="final",
reply=user_facing_error,
message=fail_msg,
session_key=str(session_id),
)
try: try:
await conn.send_event( await conn.send_event(
"session.marker", "session.marker",
@ -287,8 +371,18 @@ async def run_agent_turn_via_bridge(
) )
except Exception: except Exception:
pass pass
try:
await conn.send_event("session.message", {"sessionKey": str(session_id), "message": fail_msg})
except Exception:
pass
return return
finally: finally:
hb_stop.set()
if hb_task is not None:
try:
await hb_task
except Exception:
pass
_rid0 = str(run_id_holder.get("run_id") or "") _rid0 = str(run_id_holder.get("run_id") or "")
if _rid0: if _rid0:
with conn._abort_lock: with conn._abort_lock:
@ -296,6 +390,25 @@ async def run_agent_turn_via_bridge(
conn._aborted_run_ids.discard(_rid0) conn._aborted_run_ids.discard(_rid0)
rid = str(getattr(result, "run_id", "") or run_id_holder.get("run_id") or "") rid = str(getattr(result, "run_id", "") or run_id_holder.get("run_id") or "")
run_status = "success"
run_last_error_code = ""
run_stop_reason = ""
run_error_detail = ""
try:
rr = store.oclaw_run_get(run_id=rid, tenant_id=tenant_id or None) if rid else None
if rr is not None:
run_status = str(getattr(rr, "status", "") or "success").strip().lower() or "success"
payload = getattr(rr, "payload", {}) or {}
if isinstance(payload, dict):
run_last_error_code = str(payload.get("last_error_code") or "").strip()
run_stop_reason = str(payload.get("stop_reason") or "").strip()
if rid and run_status == "failed":
attempts = store.oclaw_attempt_list(run_id=rid, limit=30)
if attempts:
latest = attempts[-1] if isinstance(attempts[-1], dict) else {}
run_error_detail = str((latest or {}).get("reason") or "").strip()
except Exception:
pass
ttft_payload: dict[str, Any] | None = None ttft_payload: dict[str, Any] | None = None
try: try:
ft = first_token_ms_holder.get("ms") ft = first_token_ms_holder.get("ms")
@ -362,6 +475,10 @@ async def run_agent_turn_via_bridge(
"relayTtlSessionCount": int(getattr(result, "relay_ttl_session_count", 0) or 0), "relayTtlSessionCount": int(getattr(result, "relay_ttl_session_count", 0) or 0),
"relayTtlKeepCount": int(getattr(result, "relay_ttl_keep_count", 0) or 0), "relayTtlKeepCount": int(getattr(result, "relay_ttl_keep_count", 0) or 0),
"ttft": ttft_payload if isinstance(ttft_payload, dict) else None, "ttft": ttft_payload if isinstance(ttft_payload, dict) else None,
"status": run_status,
"lastErrorCode": run_last_error_code,
"stopReason": run_stop_reason,
"errorDetail": run_error_detail,
}, },
) )
await conn.emit_agent_event( await conn.emit_agent_event(
@ -369,7 +486,7 @@ async def run_agent_turn_via_bridge(
stream="lifecycle", stream="lifecycle",
data={ data={
"phase": "end", "phase": "end",
"status": "ok", "status": "error" if run_status == "failed" else "ok",
"reply": str(getattr(result, "reply_text", "") or ""), "reply": str(getattr(result, "reply_text", "") or ""),
"elapsedMs": int(getattr(result, "elapsed_ms", 0) or 0), "elapsedMs": int(getattr(result, "elapsed_ms", 0) or 0),
"mode": str(getattr(result, "mode", "sync_direct") or "sync_direct"), "mode": str(getattr(result, "mode", "sync_direct") or "sync_direct"),
@ -382,26 +499,49 @@ async def run_agent_turn_via_bridge(
with buf_lock: with buf_lock:
final_text = "".join(token_chunks) final_text = "".join(token_chunks)
final_msg: dict[str, Any] = {"role": "assistant", "content": final_text, "timestamp": now_ms()} final_msg: dict[str, Any] = {"role": "assistant", "content": final_text, "timestamp": now_ms()}
try: if run_status == "failed" and not str(final_text or "").strip():
persisted = store.get_messages(session_id=session_id, limit=12) err_code = run_last_error_code or "unknown_error"
for m in reversed(list(persisted or [])): stop_reason = run_stop_reason or "failed"
if str(getattr(m, "role", "") or "").lower() != "assistant": detail_line = f"\n详细原因:{run_error_detail}" if run_error_detail else ""
continue final_text = (
content = str(getattr(m, "content", "") or "") "本轮执行失败,已提前结束。"
if not content.strip(): f"\n\n错误代码:{err_code}"
continue f"\n停止原因:{stop_reason}"
final_text = content f"{detail_line}"
final_msg = { "\n\n请重试,或减少输入复杂度后再试。"
"id": int(getattr(m, "id", 0) or 0), )
"role": "assistant", final_msg = {"role": "assistant", "content": final_text, "timestamp": now_ms()}
"content": content, try:
"timestamp": str(getattr(m, "timestamp", "") or ""), store.add_message(
"tool_calls": getattr(m, "tool_calls", None), session_id=str(session_id),
"attachments": getattr(m, "attachments", None), role="assistant",
} content=str(final_text),
break turn_uuid=str(rid or "") or None,
except Exception: event_type="assistant_text",
pass )
except Exception:
pass
elif run_status != "failed":
try:
persisted = store.get_messages(session_id=session_id, limit=12)
for m in reversed(list(persisted or [])):
if str(getattr(m, "role", "") or "").lower() != "assistant":
continue
content = str(getattr(m, "content", "") or "")
if not content.strip():
continue
final_text = content
final_msg = {
"id": int(getattr(m, "id", 0) or 0),
"role": "assistant",
"content": content,
"timestamp": str(getattr(m, "timestamp", "") or ""),
"tool_calls": getattr(m, "tool_calls", None),
"attachments": getattr(m, "attachments", None),
}
break
except Exception:
pass
await conn.emit_chat_event(run_id=rid, state="final", reply=str(final_text or ""), message=final_msg, session_key=str(session_id)) await conn.emit_chat_event(run_id=rid, state="final", reply=str(final_text or ""), message=final_msg, session_key=str(session_id))
try: try:

View file

@ -24,6 +24,12 @@ def _model_id_suggests_gemini(model: str | None) -> bool:
return "gemini" in (model or "").lower() return "gemini" in (model or "").lower()
def _is_minimax_compat(model: str | None, base_url: str | None) -> bool:
m = str(model or "").strip().lower()
b = str(base_url or "").strip().lower()
return ("minimax" in m) or ("minimax" in b)
def _find_thought_signature_in_obj(o: Any) -> str | None: def _find_thought_signature_in_obj(o: Any) -> str | None:
if isinstance(o, dict): if isinstance(o, dict):
for k in ("thought_signature", "thoughtSignature"): for k in ("thought_signature", "thoughtSignature"):
@ -126,12 +132,34 @@ class OpenAIChatModel(ChatModel):
except Exception as exc: except Exception as exc:
logger.warning("replay_policy apply failed (%s); continuing without rewrite", exc) logger.warning("replay_policy apply failed (%s); continuing without rewrite", exc)
minimax_compat = _is_minimax_compat(self.model, self.base_url)
use_tools = bool(tools) use_tools = bool(tools)
cleaned_msgs: list[dict[str, Any]] = [] cleaned_msgs: list[dict[str, Any]] = []
for m in norm_msgs or []: for m in norm_msgs or []:
if not isinstance(m, dict): if not isinstance(m, dict):
continue continue
if str(m.get("role") or "") == "tool": role = str(m.get("role") or "")
if minimax_compat and role == "assistant" and isinstance(m.get("tool_calls"), list):
# MiniMax OpenAI-compat can reject long-history tool_use/tool_result replays.
# Keep text semantics, but remove historical wire-level tool_calls.
mm = dict(m)
mm.pop("tool_calls", None)
cleaned_msgs.append(mm)
continue
if minimax_compat and role == "tool":
# Downgrade history tool_result rows to plain assistant text context,
# avoiding strict tool_result sequence validation on replay.
tcid = str(m.get("tool_call_id") or m.get("call_id") or "").strip()
tname = str(m.get("name") or "").strip()
raw = str(m.get("content") or "")
prefix = "[tool_result replay]"
if tname:
prefix += f" name={tname}"
if tcid:
prefix += f" id={tcid}"
cleaned_msgs.append({"role": "assistant", "content": f"{prefix}\n{raw}".strip()})
continue
if role == "tool":
mm = dict(m) mm = dict(m)
for k in ("tool_call_id", "call_id"): for k in ("tool_call_id", "call_id"):
if k in mm: if k in mm:

View file

@ -28,6 +28,7 @@ ALL_ATTEMPT_ERROR_CODES = (
"control_interrupted", "control_interrupted",
"auth_invalid_credentials", "auth_invalid_credentials",
"input_invalid_request", "input_invalid_request",
"tool_replay_protocol_mismatch",
"context_overflow", "context_overflow",
"tool_loop_guard", "tool_loop_guard",
"tool_execution_failed", "tool_execution_failed",
@ -83,6 +84,8 @@ def _classify_attempt_error(exc: Exception) -> tuple[str, str, bool]:
return ("control_interrupted", raw[:500], False) return ("control_interrupted", raw[:500], False)
if "api_key" in low or "invalid api key" in low or "unauthorized" in low or "401" in low or "forbidden" in low or "403" in low: if "api_key" in low or "invalid api key" in low or "unauthorized" in low or "401" in low or "forbidden" in low or "403" in low:
return ("auth_invalid_credentials", raw[:500], False) return ("auth_invalid_credentials", raw[:500], False)
if "invalid tool_result sequence" in low or "unexpected tool_use_id" in low:
return ("tool_replay_protocol_mismatch", raw[:500], True)
if "invalid_request" in low or "bad_request" in low or "400" in low: if "invalid_request" in low or "bad_request" in low or "400" in low:
return ("input_invalid_request", raw[:500], False) return ("input_invalid_request", raw[:500], False)
if "context_length" in low or "token limit" in low or "max context" in low: if "context_length" in low or "token limit" in low or "max context" in low:

View file

@ -61,6 +61,7 @@ DEFAULT_RETRYABLE_ERROR_CODES = (
"provider_rate_limited", "provider_rate_limited",
"provider_temporary_error", "provider_temporary_error",
"provider_unavailable", "provider_unavailable",
"tool_replay_protocol_mismatch",
"context_overflow", "context_overflow",
"tool_execution_failed", "tool_execution_failed",
) )