From 3cdf550298ecd2c7600cc3ecaa5ef101a6e8f177 Mon Sep 17 00:00:00 2001 From: oliver Date: Thu, 23 Jul 2026 10:16:07 +0800 Subject: [PATCH] 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 --- interfaces/admin/static/chat.js | 4 + interfaces/gateway/server_methods/chat.py | 40 +++- interfaces/ws/server_methods_bridge.py | 13 ++ interfaces/ws/turn_runner.py | 221 +++++++++++++++++----- runtime/direct_loop.py | 17 +- tests/test_turn_runner_abort.py | 164 ++++++++++++++++ tests/test_turn_runner_turn_uuid.py | 53 ++++++ 7 files changed, 450 insertions(+), 62 deletions(-) create mode 100644 tests/test_turn_runner_abort.py create mode 100644 tests/test_turn_runner_turn_uuid.py diff --git a/interfaces/admin/static/chat.js b/interfaces/admin/static/chat.js index 616912df..318bb268 100644 --- a/interfaces/admin/static/chat.js +++ b/interfaces/admin/static/chat.js @@ -5304,6 +5304,10 @@ ${autoLimit ? `
auto-added claus } if (eventName === "session.turn_started") { turnAcceptedAtMs = Number(payload.acceptedAt || Date.now()) || Date.now(); + if (payload.runId) { + chatRunId = String(payload.runId); + currentAbortMeta = { sessionId: String(activeId || ""), runId: String(chatRunId || "") }; + } _setStreamStatusBase("start"); return; } diff --git a/interfaces/gateway/server_methods/chat.py b/interfaces/gateway/server_methods/chat.py index 1a2ead97..751e0813 100644 --- a/interfaces/gateway/server_methods/chat.py +++ b/interfaces/gateway/server_methods/chat.py @@ -51,18 +51,40 @@ def _chat_abort_handler(opts: dict[str, Any]) -> None: _bad(respond, "invalid chat.abort params") return run_id = params.get("runId") - if not isinstance(run_id, str) or not run_id.strip(): - _bad(respond, "invalid chat.abort params: runId required") + session_key = params.get("sessionKey") or params.get("key") + 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 aborted = False + aborted_run_ids: list[str] = [] if isinstance(context, dict): - abort_fn = context.get("abort_chat_run") - if callable(abort_fn): - try: - aborted = bool(abort_fn(run_id.strip())) - except Exception: - aborted = False - _ok(respond, {"runId": run_id.strip(), "aborted": aborted}) + if rid: + abort_fn = context.get("abort_chat_run") + if callable(abort_fn): + try: + aborted = bool(abort_fn(rid)) + if 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: diff --git a/interfaces/ws/server_methods_bridge.py b/interfaces/ws/server_methods_bridge.py index 07234ade..72439af4 100644 --- a/interfaces/ws/server_methods_bridge.py +++ b/interfaces/ws/server_methods_bridge.py @@ -35,6 +35,18 @@ def build_gateway_context( aborted_run_ids.add(rid) 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: sid = str(session_key or "").strip() txt = str(message or "").strip() @@ -151,6 +163,7 @@ def build_gateway_context( context = build_common_gateway_context(store=store) context.update({ "abort_chat_run": _abort_chat_run, + "abort_chat_session": _abort_chat_session, "enqueue_chat_send": _enqueue_chat_send, "run_agent": _run_agent, "enqueue_session_send": _enqueue_session_send, diff --git a/interfaces/ws/turn_runner.py b/interfaces/ws/turn_runner.py index adbff092..90dc59ad 100644 --- a/interfaces/ws/turn_runner.py +++ b/interfaces/ws/turn_runner.py @@ -20,6 +20,27 @@ from svc.persistence.assistant_store import get_assistant_store _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: """True when chat_message.attachments has at least one JSON object (list or dict or encoded string).""" if raw is None: @@ -207,6 +228,8 @@ async def run_agent_turn_via_bridge( marker_turn_count = 0 marker_session_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: try: @@ -214,15 +237,21 @@ async def run_agent_turn_via_bridge( except Exception: 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: + # 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) rid = run_id_holder.get("run_id") or "" if first_token_ms_holder["ms"] is None and token_text: 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: _schedule(conn.emit_agent_event(run_id=rid, stream="assistant", data={"event": "delta", "delta": token_text})) with buf_lock: @@ -237,19 +266,17 @@ async def run_agent_turn_via_bridge( ) def on_progress(text: str) -> None: + if should_stop(): + return rid = run_id_holder.get("run_id") or "" 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)})) def on_tool_ui(name: str, payload: dict[str, Any]) -> None: + if should_stop(): + return rid = run_id_holder.get("run_id") or "" 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()} _schedule(conn.send_event("session.tool", pl)) @@ -334,6 +361,7 @@ async def run_agent_turn_via_bridge( on_token=on_token, on_progress=on_progress, on_tool_ui=on_tool_ui, + should_stop=should_stop, ) 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) break except asyncio.TimeoutError: + if should_stop(): + continue 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, @@ -360,33 +387,144 @@ async def run_agent_turn_via_bridge( # Best-effort keepalive; never fail the turn on heartbeat errors. 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_task: asyncio.Task[Any] | None = asyncio.create_task(_heartbeat_during_turn(hb_stop)) + result: Any = None + turn_exc: BaseException | None = None try: result = await asyncio.to_thread(_run_turn_sync) - except Exception as exc: + except BaseException as exc: + turn_exc = exc + finally: hb_stop.set() if hb_task is not None: try: await hb_task except Exception: 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 = ( "本轮执行失败:工具不可用或执行异常。" 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] = { "role": "assistant", "content": user_facing_error, "timestamp": now_ms(), + "turn_uuid": fail_turn, } try: store.add_message( session_id=str(session_id), role="assistant", content=str(user_facing_error), - turn_uuid=None, + turn_uuid=fail_turn, event_type="assistant_text", ) except Exception: @@ -428,7 +566,7 @@ async def run_agent_turn_via_bridge( run_id=run_id_holder.get("run_id") or "", stream="lifecycle", # 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( run_id=run_id_holder.get("run_id") or "", @@ -457,20 +595,7 @@ async def run_agent_turn_via_bridge( except Exception: pass 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_last_error_code = "" run_stop_reason = "" @@ -488,8 +613,15 @@ async def run_agent_turn_via_bridge( if attempts: latest = attempts[-1] if isinstance(attempts[-1], dict) else {} 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: 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 try: 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)) except Exception: 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: await conn.send_res( req_id, @@ -581,8 +695,15 @@ async def run_agent_turn_via_bridge( with buf_lock: final_text = "".join(token_chunks) stream_snapshot = str(final_text or "").strip() - turn_for_persist = str(getattr(result, "turn_uuid", "") or "").strip() - final_msg: dict[str, Any] = {"role": "assistant", "content": final_text, "timestamp": now_ms(), "turn_uuid": turn_for_persist or ""} + turn_for_persist = str(getattr(result, "turn_uuid", "") or "").strip() or _latest_turn_uuid( + 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: err_code = run_last_error_code or "unknown_error" stop_reason = run_stop_reason or "failed" diff --git a/runtime/direct_loop.py b/runtime/direct_loop.py index 29699082..7efed6f5 100644 --- a/runtime/direct_loop.py +++ b/runtime/direct_loop.py @@ -924,6 +924,7 @@ def _chat_with_empty_body_retry( on_progress: Optional[Callable[[str], None]], progress_label: str = "oclaw: think", allow_dsml_text_tools: bool = False, + should_stop: Optional[Callable[[], bool]] = None, ) -> Any: # Empty assistant body can occur transiently at upstream gateways. # 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) started = time.perf_counter() 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: + _check_stop(should_stop) content = str(getattr(resp, "content", "") or "") reasoning = str(getattr(resp, "reasoning_content", "") or "") tool_calls = list(getattr(resp, "tool_calls", []) or []) @@ -962,14 +971,14 @@ def _chat_with_empty_body_retry( } ] 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 if on_progress: on_progress(f"{progress_label} retry-empty ({retries_done + 1}/{retry_max})…") if retry_delay_ms > 0: time.sleep(float(retry_delay_ms) / 1000.0) 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]: @@ -1516,6 +1525,7 @@ def run_oclaw_direct_loop( on_progress=on_progress, progress_label="oclaw: think", allow_dsml_text_tools=allow_dsml_text_tools, + should_stop=should_stop, ) assistant_text = str(getattr(resp, "content", "") or "") reasoning_text = str(getattr(resp, "reasoning_content", "") or "") @@ -1673,6 +1683,7 @@ def run_oclaw_direct_loop( on_progress=on_progress, progress_label="oclaw: finalize", allow_dsml_text_tools=False, + should_stop=should_stop, ) step = _persist_assistant_step( store=store, diff --git a/tests/test_turn_runner_abort.py b/tests/test_turn_runner_abort.py new file mode 100644 index 00000000..d55cb6c0 --- /dev/null +++ b/tests/test_turn_runner_abort.py @@ -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 diff --git a/tests/test_turn_runner_turn_uuid.py b/tests/test_turn_runner_turn_uuid.py new file mode 100644 index 00000000..840a43ba --- /dev/null +++ b/tests/test_turn_runner_turn_uuid.py @@ -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) == ""