mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 03:30:48 +08:00
Make Chat abort actually stop the agent turn.
Wire should_stop from aborted runs, snapshot abort before clearing flags so finals are not emitted after stop, and keep fail-path turn_uuid aligned for bubble splitting. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
23ecce4755
commit
3cdf550298
7 changed files with 450 additions and 62 deletions
|
|
@ -5304,6 +5304,10 @@ ${autoLimit ? `<div style="margin-top:8px;"><span class="muted">auto-added claus
|
||||||
}
|
}
|
||||||
if (eventName === "session.turn_started") {
|
if (eventName === "session.turn_started") {
|
||||||
turnAcceptedAtMs = Number(payload.acceptedAt || Date.now()) || Date.now();
|
turnAcceptedAtMs = Number(payload.acceptedAt || Date.now()) || Date.now();
|
||||||
|
if (payload.runId) {
|
||||||
|
chatRunId = String(payload.runId);
|
||||||
|
currentAbortMeta = { sessionId: String(activeId || ""), runId: String(chatRunId || "") };
|
||||||
|
}
|
||||||
_setStreamStatusBase("start");
|
_setStreamStatusBase("start");
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -51,18 +51,40 @@ def _chat_abort_handler(opts: dict[str, Any]) -> None:
|
||||||
_bad(respond, "invalid chat.abort params")
|
_bad(respond, "invalid chat.abort params")
|
||||||
return
|
return
|
||||||
run_id = params.get("runId")
|
run_id = params.get("runId")
|
||||||
if not isinstance(run_id, str) or not run_id.strip():
|
session_key = params.get("sessionKey") or params.get("key")
|
||||||
_bad(respond, "invalid chat.abort params: runId required")
|
rid = run_id.strip() if isinstance(run_id, str) and run_id.strip() else ""
|
||||||
|
sid = session_key.strip() if isinstance(session_key, str) and session_key.strip() else ""
|
||||||
|
if not rid and not sid:
|
||||||
|
_bad(respond, "invalid chat.abort params: runId or sessionKey required")
|
||||||
return
|
return
|
||||||
aborted = False
|
aborted = False
|
||||||
|
aborted_run_ids: list[str] = []
|
||||||
if isinstance(context, dict):
|
if isinstance(context, dict):
|
||||||
abort_fn = context.get("abort_chat_run")
|
if rid:
|
||||||
if callable(abort_fn):
|
abort_fn = context.get("abort_chat_run")
|
||||||
try:
|
if callable(abort_fn):
|
||||||
aborted = bool(abort_fn(run_id.strip()))
|
try:
|
||||||
except Exception:
|
aborted = bool(abort_fn(rid))
|
||||||
aborted = False
|
if aborted:
|
||||||
_ok(respond, {"runId": run_id.strip(), "aborted": aborted})
|
aborted_run_ids = [rid]
|
||||||
|
except Exception:
|
||||||
|
aborted = False
|
||||||
|
if not aborted and sid:
|
||||||
|
abort_session_fn = context.get("abort_chat_session")
|
||||||
|
if callable(abort_session_fn):
|
||||||
|
try:
|
||||||
|
aborted = bool(abort_session_fn(sid))
|
||||||
|
except Exception:
|
||||||
|
aborted = False
|
||||||
|
_ok(
|
||||||
|
respond,
|
||||||
|
{
|
||||||
|
"runId": rid,
|
||||||
|
"sessionKey": sid,
|
||||||
|
"aborted": aborted,
|
||||||
|
"abortedRunIds": aborted_run_ids,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _chat_send_handler(opts: dict[str, Any]) -> None:
|
def _chat_send_handler(opts: dict[str, Any]) -> None:
|
||||||
|
|
|
||||||
|
|
@ -35,6 +35,18 @@ def build_gateway_context(
|
||||||
aborted_run_ids.add(rid)
|
aborted_run_ids.add(rid)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def _abort_chat_session(session_key: str) -> bool:
|
||||||
|
sid = str(session_key or "").strip()
|
||||||
|
if not sid:
|
||||||
|
return False
|
||||||
|
with abort_lock:
|
||||||
|
rids = [rid for rid, s in active_run_session.items() if str(s) == sid]
|
||||||
|
if not rids:
|
||||||
|
return False
|
||||||
|
for rid in rids:
|
||||||
|
aborted_run_ids.add(rid)
|
||||||
|
return True
|
||||||
|
|
||||||
def _enqueue_chat_send(session_key: str, message: str, run_id: str | None, params: dict[str, Any]) -> bool:
|
def _enqueue_chat_send(session_key: str, message: str, run_id: str | None, params: dict[str, Any]) -> bool:
|
||||||
sid = str(session_key or "").strip()
|
sid = str(session_key or "").strip()
|
||||||
txt = str(message or "").strip()
|
txt = str(message or "").strip()
|
||||||
|
|
@ -151,6 +163,7 @@ def build_gateway_context(
|
||||||
context = build_common_gateway_context(store=store)
|
context = build_common_gateway_context(store=store)
|
||||||
context.update({
|
context.update({
|
||||||
"abort_chat_run": _abort_chat_run,
|
"abort_chat_run": _abort_chat_run,
|
||||||
|
"abort_chat_session": _abort_chat_session,
|
||||||
"enqueue_chat_send": _enqueue_chat_send,
|
"enqueue_chat_send": _enqueue_chat_send,
|
||||||
"run_agent": _run_agent,
|
"run_agent": _run_agent,
|
||||||
"enqueue_session_send": _enqueue_session_send,
|
"enqueue_session_send": _enqueue_session_send,
|
||||||
|
|
|
||||||
|
|
@ -20,6 +20,27 @@ from svc.persistence.assistant_store import get_assistant_store
|
||||||
_LOG = logging.getLogger(__name__)
|
_LOG = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _latest_turn_uuid(store: SqliteStore, session_id: str) -> str:
|
||||||
|
"""Recover this turn's uuid from recent persisted messages (prefer latest user)."""
|
||||||
|
sid = str(session_id or "").strip()
|
||||||
|
if not sid:
|
||||||
|
return ""
|
||||||
|
try:
|
||||||
|
rows = store.get_messages(sid, limit=40)
|
||||||
|
except Exception:
|
||||||
|
return ""
|
||||||
|
latest_any = ""
|
||||||
|
for m in reversed(list(rows or [])):
|
||||||
|
tu = str(getattr(m, "turn_uuid", "") or "").strip()
|
||||||
|
if not tu:
|
||||||
|
continue
|
||||||
|
if str(getattr(m, "role", "") or "").strip().lower() == "user":
|
||||||
|
return tu
|
||||||
|
if not latest_any:
|
||||||
|
latest_any = tu
|
||||||
|
return latest_any
|
||||||
|
|
||||||
|
|
||||||
def _persisted_chat_attachments_nonempty(raw: Any) -> bool:
|
def _persisted_chat_attachments_nonempty(raw: Any) -> bool:
|
||||||
"""True when chat_message.attachments has at least one JSON object (list or dict or encoded string)."""
|
"""True when chat_message.attachments has at least one JSON object (list or dict or encoded string)."""
|
||||||
if raw is None:
|
if raw is None:
|
||||||
|
|
@ -207,6 +228,8 @@ async def run_agent_turn_via_bridge(
|
||||||
marker_turn_count = 0
|
marker_turn_count = 0
|
||||||
marker_session_count = 0
|
marker_session_count = 0
|
||||||
marker_keep_count = 0
|
marker_keep_count = 0
|
||||||
|
# Snapshot abort before finally clears ``_aborted_run_ids`` (otherwise final still emits).
|
||||||
|
was_aborted = False
|
||||||
|
|
||||||
def _schedule(coro: Any) -> None:
|
def _schedule(coro: Any) -> None:
|
||||||
try:
|
try:
|
||||||
|
|
@ -214,15 +237,21 @@ async def run_agent_turn_via_bridge(
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def should_stop() -> bool:
|
||||||
|
rid = str(run_id_holder.get("run_id") or "")
|
||||||
|
if not rid:
|
||||||
|
return False
|
||||||
|
with conn._abort_lock:
|
||||||
|
return rid in conn._aborted_run_ids
|
||||||
|
|
||||||
def on_token(tok: str) -> None:
|
def on_token(tok: str) -> None:
|
||||||
|
# Raise so streaming LLM / loop exits instead of only suppressing WS deltas.
|
||||||
|
if should_stop():
|
||||||
|
raise RuntimeError("generation interrupted by user")
|
||||||
token_text = str(tok)
|
token_text = str(tok)
|
||||||
rid = run_id_holder.get("run_id") or ""
|
rid = run_id_holder.get("run_id") or ""
|
||||||
if first_token_ms_holder["ms"] is None and token_text:
|
if first_token_ms_holder["ms"] is None and token_text:
|
||||||
first_token_ms_holder["ms"] = now_ms()
|
first_token_ms_holder["ms"] = now_ms()
|
||||||
if rid:
|
|
||||||
with conn._abort_lock:
|
|
||||||
if rid in conn._aborted_run_ids:
|
|
||||||
return
|
|
||||||
if rid and not conn._is_webchat_client:
|
if rid and not conn._is_webchat_client:
|
||||||
_schedule(conn.emit_agent_event(run_id=rid, stream="assistant", data={"event": "delta", "delta": token_text}))
|
_schedule(conn.emit_agent_event(run_id=rid, stream="assistant", data={"event": "delta", "delta": token_text}))
|
||||||
with buf_lock:
|
with buf_lock:
|
||||||
|
|
@ -237,19 +266,17 @@ async def run_agent_turn_via_bridge(
|
||||||
)
|
)
|
||||||
|
|
||||||
def on_progress(text: str) -> None:
|
def on_progress(text: str) -> None:
|
||||||
|
if should_stop():
|
||||||
|
return
|
||||||
rid = run_id_holder.get("run_id") or ""
|
rid = run_id_holder.get("run_id") or ""
|
||||||
if rid:
|
if rid:
|
||||||
with conn._abort_lock:
|
|
||||||
if rid in conn._aborted_run_ids:
|
|
||||||
return
|
|
||||||
_schedule(conn.emit_agent_event(run_id=rid, stream="lifecycle", data={"phase": "running", "message": str(text)}))
|
_schedule(conn.emit_agent_event(run_id=rid, stream="lifecycle", data={"phase": "running", "message": str(text)}))
|
||||||
|
|
||||||
def on_tool_ui(name: str, payload: dict[str, Any]) -> None:
|
def on_tool_ui(name: str, payload: dict[str, Any]) -> None:
|
||||||
|
if should_stop():
|
||||||
|
return
|
||||||
rid = run_id_holder.get("run_id") or ""
|
rid = run_id_holder.get("run_id") or ""
|
||||||
if rid:
|
if rid:
|
||||||
with conn._abort_lock:
|
|
||||||
if rid in conn._aborted_run_ids:
|
|
||||||
return
|
|
||||||
pl = {"runId": rid, "sessionKey": str(session_id), "name": str(name), "payload": dict(payload or {}), "ts": now_ms()}
|
pl = {"runId": rid, "sessionKey": str(session_id), "name": str(name), "payload": dict(payload or {}), "ts": now_ms()}
|
||||||
_schedule(conn.send_event("session.tool", pl))
|
_schedule(conn.send_event("session.tool", pl))
|
||||||
|
|
||||||
|
|
@ -334,6 +361,7 @@ async def run_agent_turn_via_bridge(
|
||||||
on_token=on_token,
|
on_token=on_token,
|
||||||
on_progress=on_progress,
|
on_progress=on_progress,
|
||||||
on_tool_ui=on_tool_ui,
|
on_tool_ui=on_tool_ui,
|
||||||
|
should_stop=should_stop,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _heartbeat_during_turn(stop_evt: asyncio.Event) -> None:
|
async def _heartbeat_during_turn(stop_evt: asyncio.Event) -> None:
|
||||||
|
|
@ -344,12 +372,11 @@ async def run_agent_turn_via_bridge(
|
||||||
await asyncio.wait_for(stop_evt.wait(), timeout=4.0)
|
await asyncio.wait_for(stop_evt.wait(), timeout=4.0)
|
||||||
break
|
break
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
|
if should_stop():
|
||||||
|
continue
|
||||||
rid = run_id_holder.get("run_id") or ""
|
rid = run_id_holder.get("run_id") or ""
|
||||||
if not rid:
|
if not rid:
|
||||||
continue
|
continue
|
||||||
with conn._abort_lock:
|
|
||||||
if rid in conn._aborted_run_ids:
|
|
||||||
continue
|
|
||||||
try:
|
try:
|
||||||
await conn.emit_agent_event(
|
await conn.emit_agent_event(
|
||||||
run_id=rid,
|
run_id=rid,
|
||||||
|
|
@ -360,33 +387,144 @@ async def run_agent_turn_via_bridge(
|
||||||
# Best-effort keepalive; never fail the turn on heartbeat errors.
|
# Best-effort keepalive; never fail the turn on heartbeat errors.
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def _interrupt_reply_text() -> str:
|
||||||
|
return "已中断回答。" if str(lang or "").startswith("zh") else "Response stopped."
|
||||||
|
|
||||||
|
async def _emit_aborted_terminal(*, rid: str) -> None:
|
||||||
|
stopped = _interrupt_reply_text()
|
||||||
|
abort_turn = _latest_turn_uuid(store, str(session_id)) or uuid.uuid4().hex
|
||||||
|
abort_msg: dict[str, Any] = {
|
||||||
|
"role": "assistant",
|
||||||
|
"content": stopped,
|
||||||
|
"timestamp": now_ms(),
|
||||||
|
"turn_uuid": abort_turn,
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
store.add_message(
|
||||||
|
session_id=str(session_id),
|
||||||
|
role="assistant",
|
||||||
|
content=stopped,
|
||||||
|
turn_uuid=abort_turn,
|
||||||
|
event_type="assistant_text",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
_LOG.exception(
|
||||||
|
"turn_runner_aborted_reply_persist_failed session_id=%s run_id=%s",
|
||||||
|
str(session_id),
|
||||||
|
str(rid or ""),
|
||||||
|
)
|
||||||
|
if send_response:
|
||||||
|
try:
|
||||||
|
await conn.send_res(
|
||||||
|
req_id,
|
||||||
|
ok=True,
|
||||||
|
payload={
|
||||||
|
"runId": rid,
|
||||||
|
"acceptedAt": now_ms(),
|
||||||
|
"mode": "sync_direct",
|
||||||
|
"taskId": "",
|
||||||
|
"traceId": "",
|
||||||
|
"reply": stopped,
|
||||||
|
"selectedSpecialist": str(p.get("specialist") or "generalist"),
|
||||||
|
"interactionMode": str(p.get("interaction_mode") or "comprehensive"),
|
||||||
|
"dispatchReason": "aborted",
|
||||||
|
"executionMode": execution_mode,
|
||||||
|
"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": "aborted",
|
||||||
|
"interrupted": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
await conn.emit_agent_event(
|
||||||
|
run_id=rid,
|
||||||
|
stream="lifecycle",
|
||||||
|
data={"phase": "end", "status": "aborted"},
|
||||||
|
)
|
||||||
|
await conn.emit_chat_event(
|
||||||
|
run_id=rid,
|
||||||
|
state="aborted",
|
||||||
|
reply=stopped,
|
||||||
|
message=abort_msg,
|
||||||
|
session_key=str(session_id),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
await conn.send_event(
|
||||||
|
"session.marker",
|
||||||
|
{
|
||||||
|
"runId": rid,
|
||||||
|
"sessionKey": str(session_id),
|
||||||
|
"action": "turn_reclaimed",
|
||||||
|
"reclaimedTurnPointers": int(marker_turn_count),
|
||||||
|
"relayTtlTurnCount": int(marker_turn_count),
|
||||||
|
"relayTtlSessionCount": int(marker_session_count),
|
||||||
|
"relayTtlKeepCount": int(marker_keep_count),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
await conn.send_event("session.message", {"sessionKey": str(session_id), "message": abort_msg})
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
hb_stop = asyncio.Event()
|
hb_stop = asyncio.Event()
|
||||||
hb_task: asyncio.Task[Any] | None = asyncio.create_task(_heartbeat_during_turn(hb_stop))
|
hb_task: asyncio.Task[Any] | None = asyncio.create_task(_heartbeat_during_turn(hb_stop))
|
||||||
|
result: Any = None
|
||||||
|
turn_exc: BaseException | None = None
|
||||||
try:
|
try:
|
||||||
result = await asyncio.to_thread(_run_turn_sync)
|
result = await asyncio.to_thread(_run_turn_sync)
|
||||||
except Exception as exc:
|
except BaseException as exc:
|
||||||
|
turn_exc = exc
|
||||||
|
finally:
|
||||||
hb_stop.set()
|
hb_stop.set()
|
||||||
if hb_task is not None:
|
if hb_task is not None:
|
||||||
try:
|
try:
|
||||||
await hb_task
|
await hb_task
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
err_text = str(exc or "agent_failed")
|
_rid0 = str(run_id_holder.get("run_id") or "")
|
||||||
|
if _rid0:
|
||||||
|
with conn._abort_lock:
|
||||||
|
if _rid0 in conn._aborted_run_ids:
|
||||||
|
was_aborted = True
|
||||||
|
conn._active_run_session.pop(_rid0, None)
|
||||||
|
conn._aborted_run_ids.discard(_rid0)
|
||||||
|
|
||||||
|
rid = str(getattr(result, "run_id", "") or run_id_holder.get("run_id") or "")
|
||||||
|
if turn_exc is not None:
|
||||||
|
exc_low = str(turn_exc or "").lower()
|
||||||
|
if was_aborted or ("interrupted" in exc_low and "user" in exc_low):
|
||||||
|
was_aborted = True
|
||||||
|
await _emit_aborted_terminal(rid=rid)
|
||||||
|
return
|
||||||
|
err_text = str(turn_exc or "agent_failed")
|
||||||
user_facing_error = (
|
user_facing_error = (
|
||||||
"本轮执行失败:工具不可用或执行异常。"
|
"本轮执行失败:工具不可用或执行异常。"
|
||||||
f"\n\n错误信息:{err_text}\n\n请重试,或改用其它可用工具。"
|
f"\n\n错误信息:{err_text}\n\n请重试,或改用其它可用工具。"
|
||||||
)
|
)
|
||||||
|
fail_turn = _latest_turn_uuid(store, str(session_id)) or uuid.uuid4().hex
|
||||||
fail_msg: dict[str, Any] = {
|
fail_msg: dict[str, Any] = {
|
||||||
"role": "assistant",
|
"role": "assistant",
|
||||||
"content": user_facing_error,
|
"content": user_facing_error,
|
||||||
"timestamp": now_ms(),
|
"timestamp": now_ms(),
|
||||||
|
"turn_uuid": fail_turn,
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
store.add_message(
|
store.add_message(
|
||||||
session_id=str(session_id),
|
session_id=str(session_id),
|
||||||
role="assistant",
|
role="assistant",
|
||||||
content=str(user_facing_error),
|
content=str(user_facing_error),
|
||||||
turn_uuid=None,
|
turn_uuid=fail_turn,
|
||||||
event_type="assistant_text",
|
event_type="assistant_text",
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|
@ -428,7 +566,7 @@ async def run_agent_turn_via_bridge(
|
||||||
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(turn_exc or "agent_failed")},
|
||||||
)
|
)
|
||||||
await conn.emit_chat_event(
|
await conn.emit_chat_event(
|
||||||
run_id=run_id_holder.get("run_id") or "",
|
run_id=run_id_holder.get("run_id") or "",
|
||||||
|
|
@ -457,20 +595,7 @@ async def run_agent_turn_via_bridge(
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
return
|
return
|
||||||
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 "")
|
|
||||||
if _rid0:
|
|
||||||
with conn._abort_lock:
|
|
||||||
conn._active_run_session.pop(_rid0, None)
|
|
||||||
conn._aborted_run_ids.discard(_rid0)
|
|
||||||
|
|
||||||
rid = str(getattr(result, "run_id", "") or run_id_holder.get("run_id") or "")
|
|
||||||
run_status = "success"
|
run_status = "success"
|
||||||
run_last_error_code = ""
|
run_last_error_code = ""
|
||||||
run_stop_reason = ""
|
run_stop_reason = ""
|
||||||
|
|
@ -488,8 +613,15 @@ async def run_agent_turn_via_bridge(
|
||||||
if attempts:
|
if attempts:
|
||||||
latest = attempts[-1] if isinstance(attempts[-1], dict) else {}
|
latest = attempts[-1] if isinstance(attempts[-1], dict) else {}
|
||||||
run_error_detail = str((latest or {}).get("reason") or "").strip()
|
run_error_detail = str((latest or {}).get("reason") or "").strip()
|
||||||
|
if not run_last_error_code:
|
||||||
|
run_last_error_code = str((latest or {}).get("error_code") or "").strip()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
if (not was_aborted) and run_last_error_code == "control_interrupted":
|
||||||
|
was_aborted = True
|
||||||
|
if was_aborted:
|
||||||
|
await _emit_aborted_terminal(rid=rid)
|
||||||
|
return
|
||||||
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")
|
||||||
|
|
@ -513,24 +645,6 @@ async def run_agent_turn_via_bridge(
|
||||||
ttft_payload = _compute_ttft(rows, accepted_ms=int(accepted_ms))
|
ttft_payload = _compute_ttft(rows, accepted_ms=int(accepted_ms))
|
||||||
except Exception:
|
except Exception:
|
||||||
ttft_payload = None
|
ttft_payload = None
|
||||||
with conn._abort_lock:
|
|
||||||
if rid and rid in conn._aborted_run_ids:
|
|
||||||
try:
|
|
||||||
await conn.send_event(
|
|
||||||
"session.marker",
|
|
||||||
{
|
|
||||||
"runId": rid,
|
|
||||||
"sessionKey": str(session_id),
|
|
||||||
"action": "turn_reclaimed",
|
|
||||||
"reclaimedTurnPointers": int(marker_turn_count),
|
|
||||||
"relayTtlTurnCount": int(marker_turn_count),
|
|
||||||
"relayTtlSessionCount": int(marker_session_count),
|
|
||||||
"relayTtlKeepCount": int(marker_keep_count),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
return
|
|
||||||
if send_response:
|
if send_response:
|
||||||
await conn.send_res(
|
await conn.send_res(
|
||||||
req_id,
|
req_id,
|
||||||
|
|
@ -581,8 +695,15 @@ async def run_agent_turn_via_bridge(
|
||||||
with buf_lock:
|
with buf_lock:
|
||||||
final_text = "".join(token_chunks)
|
final_text = "".join(token_chunks)
|
||||||
stream_snapshot = str(final_text or "").strip()
|
stream_snapshot = str(final_text or "").strip()
|
||||||
turn_for_persist = str(getattr(result, "turn_uuid", "") or "").strip()
|
turn_for_persist = str(getattr(result, "turn_uuid", "") or "").strip() or _latest_turn_uuid(
|
||||||
final_msg: dict[str, Any] = {"role": "assistant", "content": final_text, "timestamp": now_ms(), "turn_uuid": turn_for_persist or ""}
|
store, str(session_id)
|
||||||
|
)
|
||||||
|
final_msg: dict[str, Any] = {
|
||||||
|
"role": "assistant",
|
||||||
|
"content": final_text,
|
||||||
|
"timestamp": now_ms(),
|
||||||
|
"turn_uuid": turn_for_persist or "",
|
||||||
|
}
|
||||||
if run_status == "failed" and not stream_snapshot:
|
if run_status == "failed" and not stream_snapshot:
|
||||||
err_code = run_last_error_code or "unknown_error"
|
err_code = run_last_error_code or "unknown_error"
|
||||||
stop_reason = run_stop_reason or "failed"
|
stop_reason = run_stop_reason or "failed"
|
||||||
|
|
|
||||||
|
|
@ -924,6 +924,7 @@ def _chat_with_empty_body_retry(
|
||||||
on_progress: Optional[Callable[[str], None]],
|
on_progress: Optional[Callable[[str], None]],
|
||||||
progress_label: str = "oclaw: think",
|
progress_label: str = "oclaw: think",
|
||||||
allow_dsml_text_tools: bool = False,
|
allow_dsml_text_tools: bool = False,
|
||||||
|
should_stop: Optional[Callable[[], bool]] = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
# Empty assistant body can occur transiently at upstream gateways.
|
# Empty assistant body can occur transiently at upstream gateways.
|
||||||
# Retry until non-empty (bounded by retry count and total timeout).
|
# Retry until non-empty (bounded by retry count and total timeout).
|
||||||
|
|
@ -932,8 +933,16 @@ def _chat_with_empty_body_retry(
|
||||||
retry_total_timeout_ms = _safe_nonneg_int(os.getenv("AIA_EMPTY_ASSISTANT_RETRY_TOTAL_TIMEOUT_MS"), 30_000, max_value=300_000)
|
retry_total_timeout_ms = _safe_nonneg_int(os.getenv("AIA_EMPTY_ASSISTANT_RETRY_TOTAL_TIMEOUT_MS"), 30_000, max_value=300_000)
|
||||||
started = time.perf_counter()
|
started = time.perf_counter()
|
||||||
retries_done = 0
|
retries_done = 0
|
||||||
resp = model.chat(msgs, llm_tools, on_token=on_token)
|
|
||||||
|
def _on_token_guarded(tok: str) -> None:
|
||||||
|
_check_stop(should_stop)
|
||||||
|
if on_token:
|
||||||
|
on_token(tok)
|
||||||
|
|
||||||
|
token_cb = _on_token_guarded if (on_token is not None or should_stop is not None) else None
|
||||||
|
resp = model.chat(msgs, llm_tools, on_token=token_cb)
|
||||||
while True:
|
while True:
|
||||||
|
_check_stop(should_stop)
|
||||||
content = str(getattr(resp, "content", "") or "")
|
content = str(getattr(resp, "content", "") or "")
|
||||||
reasoning = str(getattr(resp, "reasoning_content", "") or "")
|
reasoning = str(getattr(resp, "reasoning_content", "") or "")
|
||||||
tool_calls = list(getattr(resp, "tool_calls", []) or [])
|
tool_calls = list(getattr(resp, "tool_calls", []) or [])
|
||||||
|
|
@ -962,14 +971,14 @@ def _chat_with_empty_body_retry(
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
retries_done += 1
|
retries_done += 1
|
||||||
resp = model.chat(repair_msgs, llm_tools, on_token=on_token)
|
resp = model.chat(repair_msgs, llm_tools, on_token=token_cb)
|
||||||
continue
|
continue
|
||||||
if on_progress:
|
if on_progress:
|
||||||
on_progress(f"{progress_label} retry-empty ({retries_done + 1}/{retry_max})…")
|
on_progress(f"{progress_label} retry-empty ({retries_done + 1}/{retry_max})…")
|
||||||
if retry_delay_ms > 0:
|
if retry_delay_ms > 0:
|
||||||
time.sleep(float(retry_delay_ms) / 1000.0)
|
time.sleep(float(retry_delay_ms) / 1000.0)
|
||||||
retries_done += 1
|
retries_done += 1
|
||||||
resp = model.chat(msgs, llm_tools, on_token=on_token)
|
resp = model.chat(msgs, llm_tools, on_token=token_cb)
|
||||||
|
|
||||||
|
|
||||||
def _extract_dsml_invoke_names(text: str) -> list[str]:
|
def _extract_dsml_invoke_names(text: str) -> list[str]:
|
||||||
|
|
@ -1516,6 +1525,7 @@ def run_oclaw_direct_loop(
|
||||||
on_progress=on_progress,
|
on_progress=on_progress,
|
||||||
progress_label="oclaw: think",
|
progress_label="oclaw: think",
|
||||||
allow_dsml_text_tools=allow_dsml_text_tools,
|
allow_dsml_text_tools=allow_dsml_text_tools,
|
||||||
|
should_stop=should_stop,
|
||||||
)
|
)
|
||||||
assistant_text = str(getattr(resp, "content", "") or "")
|
assistant_text = str(getattr(resp, "content", "") or "")
|
||||||
reasoning_text = str(getattr(resp, "reasoning_content", "") or "")
|
reasoning_text = str(getattr(resp, "reasoning_content", "") or "")
|
||||||
|
|
@ -1673,6 +1683,7 @@ def run_oclaw_direct_loop(
|
||||||
on_progress=on_progress,
|
on_progress=on_progress,
|
||||||
progress_label="oclaw: finalize",
|
progress_label="oclaw: finalize",
|
||||||
allow_dsml_text_tools=False,
|
allow_dsml_text_tools=False,
|
||||||
|
should_stop=should_stop,
|
||||||
)
|
)
|
||||||
step = _persist_assistant_step(
|
step = _persist_assistant_step(
|
||||||
store=store,
|
store=store,
|
||||||
|
|
|
||||||
164
tests/test_turn_runner_abort.py
Normal file
164
tests/test_turn_runner_abort.py
Normal file
|
|
@ -0,0 +1,164 @@
|
||||||
|
"""WS turn abort must stop the agent and not emit a full final reply."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DummyResult:
|
||||||
|
run_id: str = "run-abort-1"
|
||||||
|
reply_text: str = "FULL_REPLY_SHOULD_NOT_EMIT"
|
||||||
|
elapsed_ms: int = 10
|
||||||
|
turn_uuid: str = "tu-1"
|
||||||
|
mode: str = "sync_direct"
|
||||||
|
trace_id: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class DummyLock:
|
||||||
|
def __enter__(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc, tb) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class DummyStore:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.messages: list[dict[str, Any]] = []
|
||||||
|
|
||||||
|
def get_messages(self, session_id: str, limit: int = 200) -> list[Any]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
def add_message(self, **kwargs: Any) -> None:
|
||||||
|
self.messages.append(dict(kwargs))
|
||||||
|
|
||||||
|
def oclaw_run_get(self, **kwargs: Any) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def get_setting(self, key: str) -> str:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
class DummyConn:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.auth_ctx = {"tenant_id": "t", "user_id": "u", "username": "n", "lang": "zh"}
|
||||||
|
self._abort_lock = DummyLock()
|
||||||
|
self._active_run_session: dict[str, str] = {}
|
||||||
|
self._aborted_run_ids: set[str] = set()
|
||||||
|
self._is_webchat_client = True
|
||||||
|
self._subscribed_sessions_changed = False
|
||||||
|
self.chat_calls: list[dict[str, Any]] = []
|
||||||
|
self.events: list[tuple[str, Any]] = []
|
||||||
|
|
||||||
|
async def emit_agent_event(self, **kwargs: Any) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def emit_chat_event(self, **kwargs: Any) -> None:
|
||||||
|
self.chat_calls.append(dict(kwargs))
|
||||||
|
|
||||||
|
async def send_event(self, event: str, payload: Any) -> None:
|
||||||
|
self.events.append((str(event), payload))
|
||||||
|
|
||||||
|
async def send_res(self, *args: Any, **kwargs: Any) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _run_turn(monkeypatch: pytest.MonkeyPatch, *, gateway: Any, conn: DummyConn, run_id: str) -> DummyStore:
|
||||||
|
from interfaces.ws import turn_runner
|
||||||
|
|
||||||
|
store = DummyStore()
|
||||||
|
monkeypatch.setattr(turn_runner, "SqliteStore", lambda _p: store)
|
||||||
|
monkeypatch.setattr(turn_runner, "db_path", lambda: "dummy.sqlite")
|
||||||
|
monkeypatch.setattr(turn_runner, "OclawGateway", lambda store: gateway)
|
||||||
|
monkeypatch.setattr(turn_runner, "build_gateway_executor", lambda *a, **k: object())
|
||||||
|
monkeypatch.setattr(turn_runner, "get_assistant_store", lambda: store)
|
||||||
|
monkeypatch.setattr(turn_runner, "persist_assistant_text_if_turn_missing", lambda **k: False)
|
||||||
|
|
||||||
|
asyncio.run(
|
||||||
|
turn_runner.run_agent_turn_via_bridge(
|
||||||
|
conn=conn,
|
||||||
|
req_id="req-1",
|
||||||
|
p={"message": "hello", "runId": run_id, "lang": "zh"},
|
||||||
|
session_id="s1",
|
||||||
|
send_response=False,
|
||||||
|
normalize_ws_attachments=lambda _a: [],
|
||||||
|
validate_relay_share_envelope=lambda _e: (True, "", {}),
|
||||||
|
now_ms=lambda: 1,
|
||||||
|
error_shape=lambda c, m: {"code": c, "message": m},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return store
|
||||||
|
|
||||||
|
|
||||||
|
def test_abort_before_finish_emits_aborted_not_final(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
run_id = "run-abort-1"
|
||||||
|
seen_should_stop: list[bool] = []
|
||||||
|
|
||||||
|
class Gateway:
|
||||||
|
def handle_turn(self, **kwargs: Any) -> DummyResult:
|
||||||
|
assert callable(kwargs.get("should_stop"))
|
||||||
|
seen_should_stop.append(bool(kwargs["should_stop"]()))
|
||||||
|
return DummyResult(run_id=run_id, reply_text="FULL_REPLY_SHOULD_NOT_EMIT")
|
||||||
|
|
||||||
|
conn = DummyConn()
|
||||||
|
conn._aborted_run_ids.add(run_id)
|
||||||
|
store = _run_turn(monkeypatch, gateway=Gateway(), conn=conn, run_id=run_id)
|
||||||
|
|
||||||
|
states = [c.get("state") for c in conn.chat_calls]
|
||||||
|
assert "aborted" in states
|
||||||
|
assert "final" not in states
|
||||||
|
assert seen_should_stop == [True]
|
||||||
|
assert any("已中断" in str(m.get("content") or "") for m in store.messages)
|
||||||
|
assert not any("FULL_REPLY" in str(m.get("content") or "") for m in store.messages)
|
||||||
|
|
||||||
|
|
||||||
|
def test_on_token_raises_when_aborted_mid_stream(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
run_id = "run-abort-2"
|
||||||
|
conn = DummyConn()
|
||||||
|
|
||||||
|
class Gateway:
|
||||||
|
def handle_turn(self, **kwargs: Any) -> DummyResult:
|
||||||
|
on_token = kwargs["on_token"]
|
||||||
|
on_token("partial")
|
||||||
|
conn._aborted_run_ids.add(run_id)
|
||||||
|
with pytest.raises(RuntimeError, match="interrupted"):
|
||||||
|
on_token("more")
|
||||||
|
raise RuntimeError("generation interrupted by user")
|
||||||
|
|
||||||
|
store = _run_turn(monkeypatch, gateway=Gateway(), conn=conn, run_id=run_id)
|
||||||
|
states = [c.get("state") for c in conn.chat_calls]
|
||||||
|
assert "aborted" in states
|
||||||
|
assert "final" not in states
|
||||||
|
assert any("已中断" in str(m.get("content") or "") for m in store.messages)
|
||||||
|
|
||||||
|
|
||||||
|
def test_abort_chat_session_marks_active_runs(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
from interfaces.ws import server_methods_bridge
|
||||||
|
|
||||||
|
monkeypatch.setattr(server_methods_bridge, "get_assistant_store", lambda: DummyStore())
|
||||||
|
lock = DummyLock()
|
||||||
|
active: dict[str, str] = {"r1": "s1", "r2": "s2"}
|
||||||
|
aborted: set[str] = set()
|
||||||
|
|
||||||
|
ctx = server_methods_bridge.build_gateway_context(
|
||||||
|
conn_id="c1",
|
||||||
|
subscribed_sessions_changed=False,
|
||||||
|
subscribed_message_keys=set(),
|
||||||
|
abort_lock=lock,
|
||||||
|
active_run_session=active,
|
||||||
|
aborted_run_ids=aborted,
|
||||||
|
run_agent_turn=lambda *a, **k: None,
|
||||||
|
normalize_ws_attachments=lambda _a: [],
|
||||||
|
validate_relay_share_envelope=lambda _e: (True, "", {}),
|
||||||
|
now_ms=lambda: 1,
|
||||||
|
)
|
||||||
|
assert ctx["abort_chat_session"]("s1") is True
|
||||||
|
assert "r1" in aborted
|
||||||
|
assert "r2" not in aborted
|
||||||
|
assert ctx["abort_chat_run"]("r2") is True
|
||||||
|
assert "r2" in aborted
|
||||||
53
tests/test_turn_runner_turn_uuid.py
Normal file
53
tests/test_turn_runner_turn_uuid.py
Normal file
|
|
@ -0,0 +1,53 @@
|
||||||
|
"""turn_uuid recovery / final payload alignment for WS turn runner."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from interfaces.ws.turn_runner import _latest_turn_uuid
|
||||||
|
from svc.persistence.sqlite_store import SqliteStore
|
||||||
|
|
||||||
|
|
||||||
|
def _seed_session(store: SqliteStore) -> str:
|
||||||
|
t = store.create_tenant("T")
|
||||||
|
tid = str(t["id"])
|
||||||
|
store.create_user_account(
|
||||||
|
tenant_id=tid,
|
||||||
|
username="u",
|
||||||
|
display_name="U",
|
||||||
|
role="owner",
|
||||||
|
password_hash="x",
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
uid = str(store.get_user_by_username(tenant_id=tid, username="u")["id"])
|
||||||
|
s = store.create_session_for_user(title="s", tenant_id=tid, user_id=uid)
|
||||||
|
return str(s.id)
|
||||||
|
|
||||||
|
|
||||||
|
def test_latest_turn_uuid_prefers_latest_user(tmp_path) -> None:
|
||||||
|
store = SqliteStore(str(tmp_path / "t.sqlite"))
|
||||||
|
sid = _seed_session(store)
|
||||||
|
older = uuid.uuid4().hex
|
||||||
|
newer = uuid.uuid4().hex
|
||||||
|
store.add_message(session_id=sid, role="user", content="a", turn_uuid=older, event_type="user_text")
|
||||||
|
store.add_message(
|
||||||
|
session_id=sid, role="assistant", content="ok", turn_uuid=older, event_type="assistant_text"
|
||||||
|
)
|
||||||
|
store.add_message(session_id=sid, role="user", content="b", turn_uuid=newer, event_type="user_text")
|
||||||
|
assert _latest_turn_uuid(store, sid) == newer
|
||||||
|
|
||||||
|
|
||||||
|
def test_latest_turn_uuid_falls_back_to_assistant(tmp_path) -> None:
|
||||||
|
store = SqliteStore(str(tmp_path / "u.sqlite"))
|
||||||
|
sid = _seed_session(store)
|
||||||
|
tu = uuid.uuid4().hex
|
||||||
|
store.add_message(
|
||||||
|
session_id=sid, role="assistant", content="partial", turn_uuid=tu, event_type="assistant_text"
|
||||||
|
)
|
||||||
|
assert _latest_turn_uuid(store, sid) == tu
|
||||||
|
|
||||||
|
|
||||||
|
def test_latest_turn_uuid_empty_session(tmp_path) -> None:
|
||||||
|
store = SqliteStore(str(tmp_path / "e.sqlite"))
|
||||||
|
sid = _seed_session(store)
|
||||||
|
assert _latest_turn_uuid(store, sid) == ""
|
||||||
Loading…
Add table
Add a link
Reference in a new issue