Fix run cancellation and live activity stability

This commit is contained in:
wuyang6
2026-05-14 12:23:53 +08:00
parent 105d287ebd
commit 22960c5e02
6 changed files with 117 additions and 43 deletions
+37 -25
View File
@@ -443,6 +443,28 @@ class RunManager:
record = self._runs.get(run_id) record = self._runs.get(run_id)
return self.cancel_record(record) return self.cancel_record(record)
def cancel_session(self, account_key: str, session_id: str) -> list[str]:
"""Cancel every non-terminal run for a session.
The UI can hold a stale run_id after refresh/replay while a newer
request is queued behind the same session lock. Cancelling only the
stale id leaves the queued run alive and makes the composer look
impossible to stop. Session-level cancel is the user-facing intent.
"""
with self._lock:
records = [
record
for record in self._runs.values()
if record.account_key == account_key
and record.session_id == session_id
and record.status not in {'completed', 'failed', 'cancelled'}
]
cancelled: list[str] = []
for record in records:
if self.cancel_record(record):
cancelled.append(record.run_id)
return cancelled
def cancel_record(self, record: RunRecord | None) -> bool: def cancel_record(self, record: RunRecord | None) -> bool:
if record is None or record.status in {'completed', 'failed', 'cancelled'}: if record is None or record.status in {'completed', 'failed', 'cancelled'}:
return False return False
@@ -1635,40 +1657,30 @@ def create_app(state: AgentState) -> FastAPI:
safe_id = _safe_session_id(payload.session_id) safe_id = _safe_session_id(payload.session_id)
if safe_id is None: if safe_id is None:
raise HTTPException(status_code=400, detail='Invalid session id') raise HTTPException(status_code=400, detail='Invalid session id')
cancelled = ( account_key = state._account_key(payload.account_id)
state.run_manager.cancel_run(payload.run_id) cancelled_ids: set[str] = set()
if payload.run_id if payload.run_id and state.run_manager.cancel_run(payload.run_id):
else state.run_manager.cancel_latest( cancelled_ids.add(payload.run_id)
state._account_key(payload.account_id), cancelled_ids.update(state.run_manager.cancel_session(account_key, safe_id))
stored_cancelled = state.run_state_store.finish_active_for_session(
account_key,
safe_id, safe_id,
)
)
if cancelled:
run_id = payload.run_id
if not run_id:
snapshot = state.run_state_store.snapshot_latest(
state._account_key(payload.account_id),
safe_id,
)
run_id = str(snapshot.get('run_id') or '') if snapshot else ''
if run_id:
stored = state.run_state_store.snapshot_run(run_id)
elapsed_ms = None
started_at = stored.get('started_at') if stored else None
if isinstance(started_at, (int, float)):
elapsed_ms = max(0, int((time.time() - float(started_at)) * 1000))
state.run_state_store.finish(
run_id,
status='cancelled', status='cancelled',
elapsed_ms=elapsed_ms,
stage='用户已取消', stage='用户已取消',
) )
cancelled_ids.update(stored_cancelled)
cancelled = bool(cancelled_ids)
if cancelled:
_mark_session_interrupted( _mark_session_interrupted(
state.account_paths(payload.account_id)['sessions'], state.account_paths(payload.account_id)['sessions'],
safe_id, safe_id,
status='cancelled', status='cancelled',
) )
return {'session_id': safe_id, 'cancelled': cancelled} return {
'session_id': safe_id,
'cancelled': cancelled,
'cancelled_run_ids': sorted(cancelled_ids),
}
def _run_chat_payload( def _run_chat_payload(
request: ChatRequest, request: ChatRequest,
@@ -979,10 +979,6 @@ function summarizeLiveRunEvents(
if (event.type === "run_started") { if (event.type === "run_started") {
lines.push("后端已开始执行本轮任务。"); lines.push("后端已开始执行本轮任务。");
} }
if (event.type === "content_delta" && event.delta) {
const note = event.delta.trim();
if (note && !isNoisyLiveDelta(note)) lines.push(note);
}
if (event.type === "tool_start") { if (event.type === "tool_start") {
const stageNote = const stageNote =
typeof event.assistant_content === "string" typeof event.assistant_content === "string"
@@ -610,10 +610,6 @@ function summarizeRunEvents(runStatus: ClawActiveRunStatus) {
if (event.type === "run_started") { if (event.type === "run_started") {
lines.push("后端已开始执行本轮任务。"); lines.push("后端已开始执行本轮任务。");
} }
if (event.type === "content_delta" && event.delta) {
const note = event.delta.trim();
if (note && !isNoisyLiveDelta(note)) lines.push(note);
}
if (event.type === "tool_start") { if (event.type === "tool_start") {
const stageNote = const stageNote =
typeof event.assistant_content === "string" typeof event.assistant_content === "string"
@@ -860,7 +860,11 @@ const ComposerAction: FC = () => {
includePending: true, includePending: true,
includeActive: true, includeActive: true,
}); });
const [runtimeCancelling, setRuntimeCancelling] = useState(false);
const cancelSessionId = replayedRun.sessionId ?? currentSessionId; const cancelSessionId = replayedRun.sessionId ?? currentSessionId;
useEffect(() => {
if (!runtimeRunning) setRuntimeCancelling(false);
}, [runtimeRunning]);
return ( return (
<div className="aui-composer-action-wrapper relative flex items-center justify-between"> <div className="aui-composer-action-wrapper relative flex items-center justify-between">
<div className="flex items-center gap-1"> <div className="flex items-center gap-1">
@@ -890,12 +894,15 @@ const ComposerAction: FC = () => {
size="icon" size="icon"
className="aui-composer-cancel size-8 rounded-full" className="aui-composer-cancel size-8 rounded-full"
aria-label="Stop generating" aria-label="Stop generating"
disabled={!cancelSessionId} disabled={!cancelSessionId || runtimeCancelling}
onClick={() => { onClick={() => {
if (!cancelSessionId) return; if (!cancelSessionId) return;
void cancelLatestRun(cancelSessionId, replayedRun.runId).catch( setRuntimeCancelling(true);
() => undefined, void cancelLatestRun(cancelSessionId, replayedRun.runId)
); .catch(() => undefined)
.finally(() => {
window.setTimeout(() => setRuntimeCancelling(false), 1200);
});
}} }}
> >
<SquareIcon className="aui-composer-cancel-icon size-3 fill-current" /> <SquareIcon className="aui-composer-cancel-icon size-3 fill-current" />
+23 -4
View File
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import json import json
import os
from typing import Any, Iterator from typing import Any, Iterator
from urllib import error, request from urllib import error, request
@@ -31,6 +32,7 @@ ANTHROPIC_MESSAGES_MODEL_PREFIXES = (
ANTHROPIC_VERSION = '2023-06-01' ANTHROPIC_VERSION = '2023-06-01'
ANTHROPIC_MAX_TOKENS = 4096 ANTHROPIC_MAX_TOKENS = 4096
DEFAULT_MODEL_IDLE_TIMEOUT_SECONDS = 60.0
def _join_url(base_url: str, suffix: str) -> str: def _join_url(base_url: str, suffix: str) -> str:
@@ -274,6 +276,23 @@ class OpenAICompatClient:
def __init__(self, config: ModelConfig) -> None: def __init__(self, config: ModelConfig) -> None:
self.config = config self.config = config
def _request_timeout_seconds(self) -> float:
"""Return per-socket model timeout.
`ModelConfig.timeout_seconds` is the whole request budget used by the
product, but urllib applies it as an idle socket timeout. Keeping it at
one hour makes a bad upstream stream hold a session lock for an hour.
Use a shorter idle timeout by default; any normally streaming response
keeps extending this naturally because bytes continue to arrive.
"""
raw = os.environ.get('CLAW_MODEL_IDLE_TIMEOUT_SECONDS', '').strip()
try:
configured = float(raw) if raw else DEFAULT_MODEL_IDLE_TIMEOUT_SECONDS
except ValueError:
configured = DEFAULT_MODEL_IDLE_TIMEOUT_SECONDS
configured = max(5.0, configured)
return min(float(self.config.timeout_seconds), configured)
def complete( def complete(
self, self,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
@@ -351,7 +370,7 @@ class OpenAICompatClient:
method='POST', method='POST',
) )
try: try:
with request.urlopen(req, timeout=self.config.timeout_seconds) as response: with request.urlopen(req, timeout=self._request_timeout_seconds()) as response:
yield StreamEvent(type='message_start') yield StreamEvent(type='message_start')
for event_payload in self._iter_sse_payloads(response): for event_payload in self._iter_sse_payloads(response):
yield from self._parse_stream_payload(event_payload) yield from self._parse_stream_payload(event_payload)
@@ -380,7 +399,7 @@ class OpenAICompatClient:
method='POST', method='POST',
) )
try: try:
with request.urlopen(req, timeout=self.config.timeout_seconds) as response: with request.urlopen(req, timeout=self._request_timeout_seconds()) as response:
raw = response.read() raw = response.read()
except error.HTTPError as exc: except error.HTTPError as exc:
detail = exc.read().decode('utf-8', errors='replace') detail = exc.read().decode('utf-8', errors='replace')
@@ -521,7 +540,7 @@ class OpenAICompatClient:
method='POST', method='POST',
) )
try: try:
with request.urlopen(req, timeout=self.config.timeout_seconds) as response: with request.urlopen(req, timeout=self._request_timeout_seconds()) as response:
raw = response.read() raw = response.read()
except error.HTTPError as exc: except error.HTTPError as exc:
detail = exc.read().decode('utf-8', errors='replace') detail = exc.read().decode('utf-8', errors='replace')
@@ -561,7 +580,7 @@ class OpenAICompatClient:
method='POST', method='POST',
) )
try: try:
with request.urlopen(req, timeout=self.config.timeout_seconds) as response: with request.urlopen(req, timeout=self._request_timeout_seconds()) as response:
yield StreamEvent(type='message_start') yield StreamEvent(type='message_start')
tool_block_indexes: dict[int, int] = {} tool_block_indexes: dict[int, int] = {}
next_tool_index = 0 next_tool_index = 0
+44
View File
@@ -131,6 +131,50 @@ class RunStateStore:
finished_at=time.time(), finished_at=time.time(),
) )
def finish_active_for_session(
self,
account_key: str,
session_id: str,
*,
status: str,
stage: str | None = None,
) -> list[str]:
now = time.time()
with self._connect() as conn:
rows = conn.execute(
"""
select run_id, started_at
from run_states
where account_key = ?
and session_id = ?
and status in ('queued', 'running')
""",
(account_key, session_id),
).fetchall()
run_ids: list[str] = []
for row in rows:
run_ids.append(row['run_id'])
started_at = row['started_at']
elapsed_ms = (
max(0, int((now - float(started_at)) * 1000))
if isinstance(started_at, (int, float))
else None
)
conn.execute(
"""
update run_states
set status = ?,
current_stage = coalesce(?, current_stage),
cancellable = 0,
elapsed_ms = coalesce(?, elapsed_ms),
finished_at = ?,
updated_at = ?
where run_id = ?
""",
(status, stage, elapsed_ms, now, now, row['run_id']),
)
return run_ids
def record_event(self, run_id: str, event: dict[str, Any]) -> None: def record_event(self, run_id: str, event: dict[str, Any]) -> None:
recorded_at = event.get('recorded_at') recorded_at = event.get('recorded_at')
if not isinstance(recorded_at, (int, float)): if not isinstance(recorded_at, (int, float)):