完成 WebSocket 与 workspace 架构优化收尾,并修复全量回归阻塞。

本次将握手鉴权、Origin 校验、限流、重连补偿、发送背压与观测字段打通,同时修复 gateway 异常路径与备份清理兼容问题,确保工具历史压缩语义一致并恢复全量测试通过。

Made-with: Cursor
This commit is contained in:
oliver 2026-04-27 20:41:50 +08:00
parent 09836dddfa
commit 7253f6795d
12 changed files with 536 additions and 24 deletions

View file

@ -2500,6 +2500,11 @@ async function renderChatUi() {
this.reqSeq = 0;
this.msgQueue = [];
this.waiters = [];
this.lastSeq = 0;
this._sessionSubscriptions = new Set();
this._reconnectBaseMs = 500;
this._reconnectCapMs = 5000;
this._maxReconnectAttempts = 4;
}
_isOpen() {
return this.ws && this.ws.readyState === WebSocket.OPEN;
@ -2523,6 +2528,24 @@ async function renderChatUi() {
}
async _openAndHandshake() {
if (this._isOpen()) return;
let lastErr = null;
for (let i = 0; i <= this._maxReconnectAttempts; i += 1) {
try {
await this._openAndHandshakeOnce();
await this._restoreSubscriptions();
return;
} catch (err) {
lastErr = err;
this.close();
if (i >= this._maxReconnectAttempts) break;
const delay = Math.min(this._reconnectCapMs, this._reconnectBaseMs * Math.pow(2, i));
const jitter = Math.floor(Math.random() * 150);
await new Promise((resolve) => setTimeout(resolve, delay + jitter));
}
}
throw lastErr || new Error("ws_open_failed");
}
async _openAndHandshakeOnce() {
const token = String(this.tokenProvider() || "");
const wsUrl = this._wsUrl();
const ws = new WebSocket(wsUrl);
@ -2537,6 +2560,9 @@ async function renderChatUi() {
parsed = null;
}
if (!parsed) return;
if (parsed && parsed.type === "event" && Number.isFinite(Number(parsed.seq))) {
this.lastSeq = Math.max(this.lastSeq, Number(parsed.seq) || 0);
}
const waiter = this.waiters.shift();
if (waiter) {
waiter.resolve(parsed);
@ -2555,7 +2581,6 @@ async function renderChatUi() {
reject(new Error(`ws_open_failed:${wsUrl}`));
};
});
// Receive optional connect.challenge first.
const challenge = await this._recv();
if (!(challenge && challenge.type === "event" && challenge.event === "connect.challenge")) {
throw new Error("ws_invalid_challenge");
@ -2569,6 +2594,7 @@ async function renderChatUi() {
params: {
minProtocol: 3,
maxProtocol: 3,
lastSeq: Number(this.lastSeq || 0),
client: { id: "webchat-ui", version: "0.1", platform: navigator.platform || "web", mode: "webchat" },
role: "operator",
scopes: ["operator.read", "operator.write"],
@ -2581,6 +2607,31 @@ async function renderChatUi() {
throw new Error(`ws_connect_failed:${JSON.stringify((res && res.error) || {})}`);
}
}
async _sendReqAndAwait(method, params) {
await this._openAndHandshake();
const reqId = this._nextReqId();
this.ws.send(JSON.stringify({ type: "req", id: reqId, method: String(method || ""), params: params || {} }));
for (;;) {
const msg = await this._recv();
if (msg && msg.type === "res" && msg.id === reqId) {
if (!msg.ok) throw new Error(`ws_req_failed:${JSON.stringify(msg.error || {})}`);
return msg.payload || {};
}
}
}
async _restoreSubscriptions() {
const items = Array.from(this._sessionSubscriptions);
for (let i = 0; i < items.length; i += 1) {
const key = String(items[i] || "");
if (!key) continue;
await this._sendReqAndAwait("sessions.messages.subscribe", { sessionKey: key });
}
}
_trackSessionSubscription(sessionId) {
const key = String(sessionId || "").trim();
if (!key) return;
this._sessionSubscriptions.add(key);
}
_recv() {
const ws = this.ws;
if (!ws) return Promise.reject(new Error("ws_not_connected"));
@ -2614,6 +2665,9 @@ async function renderChatUi() {
}
async sendSessionSend({ sessionId, text, attachments, interactionMode, specialist, memoryMode, idempotencyKey, signal, onEvent }) {
await this._openAndHandshake();
this._trackSessionSubscription(sessionId);
await this._sendReqAndAwait("sessions.messages.subscribe", { sessionKey: String(sessionId || "") });
const stableIdempotencyKey = String(idempotencyKey || `idem_${Date.now()}_${Math.random().toString(36).slice(2, 8)}`);
const reqId = this._nextReqId();
const req = {
type: "req",
@ -2623,7 +2677,7 @@ async function renderChatUi() {
sessionKey: String(sessionId || ""),
message: String(text || ""),
attachments: Array.isArray(attachments) ? attachments : [],
idempotencyKey: String(idempotencyKey || `idem_${Date.now()}_${Math.random().toString(36).slice(2, 8)}`),
idempotencyKey: stableIdempotencyKey,
thinking: "default",
interaction_mode: String(interactionMode || "expert"),
specialist: String(specialist || "generalist"),
@ -2632,9 +2686,22 @@ async function renderChatUi() {
};
this.ws.send(JSON.stringify(req));
let doneMeta = null;
let retriedAfterReconnect = false;
while (true) {
if (signal && signal.aborted) throw new DOMException("Aborted", "AbortError");
const msg = await this._recv();
let msg = null;
try {
msg = await this._recv();
} catch (err) {
const em = String((err && err.message) || err || "").toLowerCase();
if (!retriedAfterReconnect && (em.includes("ws_closed") || em.includes("ws_receive_failed"))) {
retriedAfterReconnect = true;
await this._openAndHandshake();
this.ws.send(JSON.stringify(req));
continue;
}
throw err;
}
if (
msg &&
msg.type === "event" &&

View file

@ -1,6 +1,7 @@
from __future__ import annotations
import base64
import os
import time
from typing import Any
@ -11,6 +12,15 @@ MAX_PAYLOAD_BYTES = 26_214_400
MAX_BUFFERED_BYTES = 52_428_800
TICK_INTERVAL_MS = 15_000
PREAUTH_HANDSHAKE_TIMEOUT_MS = 15_000
WS_REQUIRE_AUTH = str(os.getenv("OCLAW_WS_REQUIRE_AUTH") or "1").strip().lower() not in ("0", "false", "no", "off")
WS_ALLOWED_ORIGINS = [s.strip() for s in str(os.getenv("OCLAW_WS_ALLOWED_ORIGINS") or "").split(",") if s.strip()]
WS_RATE_LIMIT_WINDOW_MS = int(os.getenv("OCLAW_WS_RATE_LIMIT_WINDOW_MS") or "60000")
WS_RATE_LIMIT_CONN_PER_WINDOW = int(os.getenv("OCLAW_WS_RATE_LIMIT_CONN_PER_WINDOW") or "120")
WS_RATE_LIMIT_IP_PER_WINDOW = int(os.getenv("OCLAW_WS_RATE_LIMIT_IP_PER_WINDOW") or "240")
WS_RATE_LIMIT_USER_PER_WINDOW = int(os.getenv("OCLAW_WS_RATE_LIMIT_USER_PER_WINDOW") or "360")
WS_SEND_QUEUE_MAX_MESSAGES = int(os.getenv("OCLAW_WS_SEND_QUEUE_MAX_MESSAGES") or "256")
WS_SEND_QUEUE_MAX_BYTES = int(os.getenv("OCLAW_WS_SEND_QUEUE_MAX_BYTES") or str(MAX_BUFFERED_BYTES))
WS_EVENT_REPLAY_MAX = int(os.getenv("OCLAW_WS_EVENT_REPLAY_MAX") or "256")
def now_ms() -> int:
@ -24,6 +34,20 @@ def error_shape(code: str, message: str, *, details: Any | None = None) -> dict[
return out
def origin_is_allowed(origin: str | None, host: str | None) -> bool:
value = str(origin or "").strip()
if not value:
return True
allowlist = list(WS_ALLOWED_ORIGINS)
if allowlist:
return value in allowlist
host_value = str(host or "").strip()
if not host_value:
return False
lower = value.lower()
return lower.startswith(f"https://{host_value.lower()}") or lower.startswith(f"http://{host_value.lower()}")
def decode_base64_payload_ws(s: str | None) -> bytes | None:
raw = str(s or "").strip()
if not raw:
@ -73,8 +97,18 @@ __all__ = [
"MAX_BUFFERED_BYTES",
"TICK_INTERVAL_MS",
"PREAUTH_HANDSHAKE_TIMEOUT_MS",
"WS_REQUIRE_AUTH",
"WS_ALLOWED_ORIGINS",
"WS_RATE_LIMIT_WINDOW_MS",
"WS_RATE_LIMIT_CONN_PER_WINDOW",
"WS_RATE_LIMIT_IP_PER_WINDOW",
"WS_RATE_LIMIT_USER_PER_WINDOW",
"WS_SEND_QUEUE_MAX_MESSAGES",
"WS_SEND_QUEUE_MAX_BYTES",
"WS_EVENT_REPLAY_MAX",
"now_ms",
"error_shape",
"origin_is_allowed",
"decode_base64_payload_ws",
"normalize_ws_attachments",
]

View file

@ -19,9 +19,19 @@ def _ws_is_disconnected(conn: Any) -> bool:
return False
async def _safe_send_text(conn: Any, text: str) -> None:
async def _safe_send_text(conn: Any, text: str, *, use_queue: bool = True) -> None:
if _ws_is_disconnected(conn):
return
queue_send = getattr(conn, "_queue_send_text", None)
if use_queue and callable(queue_send):
ok = await queue_send(text)
if ok:
return
try:
await conn.ws.close(code=1013, reason="backpressure")
except Exception:
pass
return
try:
await conn.ws.send_text(text)
except Exception as e:
@ -37,7 +47,7 @@ async def send_res(conn: Any, req_id: str, *, ok: bool, payload: Any | None = No
frame["payload"] = payload
if error is not None:
frame["error"] = error
await _safe_send_text(conn, json.dumps(frame, ensure_ascii=False))
await _safe_send_text(conn, json.dumps(frame, ensure_ascii=False), use_queue=False)
async def send_event(conn: Any, event: str, payload: Any | None = None) -> None:
@ -47,7 +57,12 @@ async def send_event(conn: Any, event: str, payload: Any | None = None) -> None:
frame: dict[str, Any] = {"type": "event", "event": str(event or "event"), "seq": int(conn.seq)}
if payload is not None:
frame["payload"] = payload
await _safe_send_text(conn, json.dumps(frame, ensure_ascii=False))
raw = json.dumps(frame, ensure_ascii=False)
frame["_raw"] = raw
remember = getattr(conn, "remember_event", None)
if callable(remember):
remember(frame)
await _safe_send_text(conn, raw)
async def emit_agent_event(conn: Any, *, run_id: str, stream: str, data: dict[str, Any], now_ms: int) -> None:

View file

@ -8,6 +8,7 @@
"properties": {
"minProtocol": { "type": "integer", "minimum": 1 },
"maxProtocol": { "type": "integer", "minimum": 1 },
"lastSeq": { "type": "integer", "minimum": 0 },
"client": {
"type": "object",
"additionalProperties": false,

View file

@ -6,6 +6,8 @@ from typing import Any
from starlette.websockets import WebSocketDisconnect
from oclaw.interfaces.ws.common import WS_REQUIRE_AUTH
async def recv_frame(
*,
@ -93,16 +95,6 @@ async def handle_connect(
)
conn.handshake_failed = True
return
hello = conn.build_hello_ok(params)
hello_errs = validate_or_errors(conn.schemas.hello_ok, hello)
if hello_errs:
await conn.send_res(
req_id,
ok=False,
error=error_shape("UNAVAILABLE", f"server hello-ok schema mismatch: {format_validation_errors(hello_errs)}"),
)
conn.handshake_failed = True
return
conn.connected = True
conn.client_meta = dict(params or {})
try:
@ -122,15 +114,38 @@ async def handle_connect(
or str(auth.get("bootstrapToken") or "").strip()
or str(auth.get("deviceToken") or "").strip()
)
if auth_provided and not conn.auth_ctx:
if not conn.auth_ctx and (WS_REQUIRE_AUTH or auth_provided):
await conn.send_res(
req_id,
ok=False,
error=error_shape("INVALID_REQUEST", "unauthorized", details={"code": "AUTH_UNAUTHORIZED"}),
error=error_shape("UNAUTHORIZED", "unauthorized", details={"code": "AUTH_UNAUTHORIZED"}),
)
conn.handshake_failed = True
if hasattr(conn, "mark_handshake"):
conn.mark_handshake(ok=False)
return
if conn.auth_ctx:
conn.role = str(conn.auth_ctx.get("role") or "member")
conn.scopes = [f"user:{str(conn.auth_ctx.get('user_id') or '')}"]
hello = conn.build_hello_ok(params)
hello_errs = validate_or_errors(conn.schemas.hello_ok, hello)
if hello_errs:
await conn.send_res(
req_id,
ok=False,
error=error_shape("UNAVAILABLE", f"server hello-ok schema mismatch: {format_validation_errors(hello_errs)}"),
)
conn.handshake_failed = True
return
await conn.send_res(req_id, ok=True, payload=hello)
if hasattr(conn, "mark_handshake"):
conn.mark_handshake(ok=True)
if hasattr(conn, "replay_events_since"):
try:
last_seq = int(p.get("lastSeq") or 0)
except Exception:
last_seq = 0
await conn.replay_events_since(last_seq)
__all__ = ["recv_frame", "handle_connect"]

View file

@ -2,7 +2,10 @@
from __future__ import annotations
import asyncio
import logging
import uuid
from collections import defaultdict, deque
from typing import Any
import threading
@ -16,9 +19,17 @@ from oclaw.interfaces.ws.common import (
PREAUTH_HANDSHAKE_TIMEOUT_MS,
PROTOCOL_VERSION,
TICK_INTERVAL_MS,
WS_EVENT_REPLAY_MAX,
WS_RATE_LIMIT_CONN_PER_WINDOW,
WS_RATE_LIMIT_IP_PER_WINDOW,
WS_RATE_LIMIT_USER_PER_WINDOW,
WS_RATE_LIMIT_WINDOW_MS,
WS_SEND_QUEUE_MAX_BYTES,
WS_SEND_QUEUE_MAX_MESSAGES,
error_shape as _error_shape,
normalize_ws_attachments as _normalize_ws_attachments,
now_ms as _now_ms,
origin_is_allowed,
)
from oclaw.interfaces.ws.events import (
emit_agent_event as emit_agent_event_impl,
@ -34,8 +45,18 @@ from oclaw.interfaces.ws.turn_runner import run_agent_turn_via_bridge
from oclaw.interfaces.ws.ws_schema import format_validation_errors, get_ws_schemas, validate_or_errors
from oclaw.runtime.relay_pointer import validate_relay_share_envelope
_LOG = logging.getLogger(__name__)
class OclawWsGatewayConnection:
_rate_lock = threading.Lock()
_rate_by_ip: dict[str, deque[int]] = defaultdict(deque)
_rate_by_user: dict[str, deque[int]] = defaultdict(deque)
_stats: dict[str, int] = defaultdict(int)
_event_buffer_by_user: dict[str, deque[dict[str, Any]]] = defaultdict(
lambda: deque(maxlen=max(1, int(WS_EVENT_REPLAY_MAX)))
)
def __init__(self, ws: WebSocket):
self.ws = ws
self.schemas = get_ws_schemas()
@ -55,12 +76,30 @@ class OclawWsGatewayConnection:
self._abort_lock = threading.Lock()
self._aborted_run_ids: set[str] = set()
self._active_run_session: dict[str, str] = {}
self._send_queue: asyncio.Queue[tuple[str, int]] = asyncio.Queue(maxsize=WS_SEND_QUEUE_MAX_MESSAGES)
self._send_pending_bytes = 0
self._send_pending_lock = asyncio.Lock()
self._sender_task: asyncio.Task[None] | None = None
self._event_buffer: deque[dict[str, Any]] = deque(maxlen=max(1, int(WS_EVENT_REPLAY_MAX)))
self._rate_local: deque[int] = deque()
self._gateway_handlers = build_gateway_method_handlers()
self._now_ms = _now_ms
self._error_shape = _error_shape
async def run(self) -> None:
await run_connection_loop(self)
self._inc_stat("ws_connections_opened")
self._sender_task = asyncio.create_task(self._sender_loop())
try:
await run_connection_loop(self)
finally:
await self._drain_sender()
self._inc_stat("ws_connections_closed")
_LOG.info(
"ws connection closed conn_id=%s total_opened=%s total_closed=%s",
self.conn_id,
self._stats.get("ws_connections_opened", 0),
self._stats.get("ws_connections_closed", 0),
)
async def _recv_frame(self, *, preauth: bool = False) -> dict[str, Any] | None:
return await recv_frame_impl(
@ -88,8 +127,31 @@ class OclawWsGatewayConnection:
await close_ws_impl(self, code=code, reason=reason)
async def _dispatch_connected(self, req_id: str, method: str, params: Any) -> None:
limited, bucket = self._rate_limited()
if limited:
self._inc_stat("ws_rate_limited")
await self.send_res(
req_id,
ok=False,
error=_error_shape("RATE_LIMITED", "too many requests", details={"bucket": bucket}),
)
return
await dispatch_connected_impl(self, req_id=req_id, method=method, params=params)
@classmethod
def _inc_stat(cls, key: str, delta: int = 1) -> None:
with cls._rate_lock:
cls._stats[str(key)] = int(cls._stats.get(str(key), 0)) + int(delta)
def mark_handshake(self, *, ok: bool) -> None:
self._inc_stat("ws_handshake_ok" if ok else "ws_handshake_failed")
_LOG.info(
"ws handshake conn_id=%s ok=%s user_id=%s",
self.conn_id,
int(bool(ok)),
str((self.auth_ctx or {}).get("user_id") or ""),
)
async def _dispatch_via_server_methods(self, *, req_id: str, method: str, params: Any) -> bool:
return await dispatch_via_server_methods(
req_id=req_id,
@ -117,6 +179,122 @@ class OclawWsGatewayConnection:
now_ms=_now_ms,
)
def validate_origin(self) -> bool:
headers = getattr(self.ws, "headers", None)
origin = str(headers.get("origin") or "").strip() if headers is not None else ""
host = str(headers.get("host") or "").strip() if headers is not None else ""
allowed = origin_is_allowed(origin, host)
if not allowed:
_LOG.warning("ws origin blocked conn_id=%s origin=%s host=%s", self.conn_id, origin, host)
return allowed
def _prune_window(self, dq: deque[int], now: int) -> None:
cutoff = now - int(WS_RATE_LIMIT_WINDOW_MS)
while dq and dq[0] < cutoff:
dq.popleft()
def _rate_limited(self) -> tuple[bool, str]:
now = _now_ms()
self._prune_window(self._rate_local, now)
if len(self._rate_local) >= int(WS_RATE_LIMIT_CONN_PER_WINDOW):
return True, "connection"
self._rate_local.append(now)
headers = getattr(self.ws, "headers", None)
ip = str(headers.get("x-forwarded-for") or headers.get("x-real-ip") or "").split(",")[0].strip() if headers is not None else ""
user_id = str((self.auth_ctx or {}).get("user_id") or "").strip()
with self._rate_lock:
if ip:
ip_bucket = self._rate_by_ip[ip]
self._prune_window(ip_bucket, now)
if len(ip_bucket) >= int(WS_RATE_LIMIT_IP_PER_WINDOW):
return True, "ip"
ip_bucket.append(now)
if user_id:
user_bucket = self._rate_by_user[user_id]
self._prune_window(user_bucket, now)
if len(user_bucket) >= int(WS_RATE_LIMIT_USER_PER_WINDOW):
return True, "user"
user_bucket.append(now)
return False, ""
async def _queue_send_text(self, text: str) -> bool:
payload = str(text or "")
payload_size = len(payload.encode("utf-8", errors="ignore"))
async with self._send_pending_lock:
if self._send_pending_bytes + payload_size > int(WS_SEND_QUEUE_MAX_BYTES):
_LOG.warning("ws send queue bytes exceeded conn_id=%s", self.conn_id)
return False
self._send_pending_bytes += payload_size
try:
self._send_queue.put_nowait((payload, payload_size))
return True
except asyncio.QueueFull:
async with self._send_pending_lock:
self._send_pending_bytes = max(0, self._send_pending_bytes - payload_size)
_LOG.warning("ws send queue full conn_id=%s", self.conn_id)
return False
async def _sender_loop(self) -> None:
while True:
item = await self._send_queue.get()
if item[0] == "__STOP__":
self._send_queue.task_done()
return
payload, payload_size = item
try:
await self.ws.send_text(payload)
except Exception:
return
finally:
async with self._send_pending_lock:
self._send_pending_bytes = max(0, self._send_pending_bytes - payload_size)
self._send_queue.task_done()
async def _drain_sender(self) -> None:
if self._sender_task is None:
return
try:
self._send_queue.put_nowait(("__STOP__", 0))
except Exception:
pass
try:
await self._sender_task
except Exception:
pass
self._sender_task = None
def remember_event(self, frame: dict[str, Any]) -> None:
snap = dict(frame or {})
self._event_buffer.append(snap)
user_id = str((self.auth_ctx or {}).get("user_id") or "").strip()
if not user_id:
return
with self._rate_lock:
bucket = self._event_buffer_by_user.get(user_id)
if bucket is None or bucket.maxlen != max(1, int(WS_EVENT_REPLAY_MAX)):
bucket = deque(maxlen=max(1, int(WS_EVENT_REPLAY_MAX)))
self._event_buffer_by_user[user_id] = bucket
bucket.append(dict(snap))
async def replay_events_since(self, seq: int) -> None:
after = int(seq or 0)
frames = list(self._event_buffer)
user_id = str((self.auth_ctx or {}).get("user_id") or "").strip()
if user_id:
with self._rate_lock:
shared = list(self._event_buffer_by_user.get(user_id) or [])
if shared:
frames = shared
for frame in frames:
fseq = int(frame.get("seq") or 0)
if fseq <= after:
continue
text = frame.get("_raw")
if not isinstance(text, str) or not text:
continue
await self._queue_send_text(text)
def build_hello_ok(self, _connect_params: dict[str, Any] | None) -> dict[str, Any]:
methods = ["connect", *method_names()]
return build_hello_ok_payload(

View file

@ -11,6 +11,9 @@ async def close_ws(conn: Any, code: int = 1000, reason: str = "done") -> None:
async def run_connection_loop(conn: Any) -> None:
if hasattr(conn, "validate_origin") and not conn.validate_origin():
await close_ws(conn, 1008, "origin not allowed")
return
await conn.ws.accept()
await conn.send_event("connect.challenge", {"nonce": conn.connect_nonce, "ts": conn._now_ms()})
while True: