diff --git a/docs/ENVIRONMENT_VARIABLES.md b/docs/ENVIRONMENT_VARIABLES.md index bbdd8a85..fda84c74 100644 --- a/docs/ENVIRONMENT_VARIABLES.md +++ b/docs/ENVIRONMENT_VARIABLES.md @@ -391,6 +391,51 @@ - 作用:SSE 事件队列上限 - 生效:`oclaw/interfaces/admin/chat_api.py` +- `OCLAW_WS_REQUIRE_AUTH` + - 默认:`1` + - 作用:WebSocket 握手是否强制鉴权;开启时 `connect` 必须携带并通过 token 校验 + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_helpers.py` + +- `OCLAW_WS_ALLOWED_ORIGINS` + - 默认:空(回落为 same-host 校验) + - 作用:WebSocket 握手 Origin 白名单(逗号分隔) + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py` + +- `OCLAW_WS_RATE_LIMIT_WINDOW_MS` + - 默认:`60000` + - 作用:WebSocket 请求限流窗口(毫秒) + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py` + +- `OCLAW_WS_RATE_LIMIT_CONN_PER_WINDOW` + - 默认:`120` + - 作用:单连接在限流窗口内可处理请求上限 + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py` + +- `OCLAW_WS_RATE_LIMIT_IP_PER_WINDOW` + - 默认:`240` + - 作用:单 IP 在限流窗口内可处理请求上限 + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py` + +- `OCLAW_WS_RATE_LIMIT_USER_PER_WINDOW` + - 默认:`360` + - 作用:单用户在限流窗口内可处理请求上限 + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py` + +- `OCLAW_WS_SEND_QUEUE_MAX_MESSAGES` + - 默认:`256` + - 作用:每连接发送队列最大消息数(背压阈值) + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py`, `oclaw/interfaces/ws/events.py` + +- `OCLAW_WS_SEND_QUEUE_MAX_BYTES` + - 默认:`52428800`(与 `MAX_BUFFERED_BYTES` 一致) + - 作用:每连接发送队列最大字节数(背压阈值) + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py`, `oclaw/interfaces/ws/events.py` + +- `OCLAW_WS_EVENT_REPLAY_MAX` + - 默认:`256` + - 作用:每用户最近事件回放缓冲上限(用于 `connect.params.lastSeq` 断线补偿) + - 生效:`oclaw/interfaces/ws/common.py`, `oclaw/interfaces/ws/runtime_impl.py`, `oclaw/interfaces/ws/runtime_helpers.py` + ## WeCom 长连接 - `AIA_WECOM_LONGCONN_WORKERS` diff --git a/interfaces/admin/static/chat.js b/interfaces/admin/static/chat.js index 4a328673..ce13be05 100644 --- a/interfaces/admin/static/chat.js +++ b/interfaces/admin/static/chat.js @@ -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" && diff --git a/interfaces/ws/common.py b/interfaces/ws/common.py index 8d35c332..593b31c2 100644 --- a/interfaces/ws/common.py +++ b/interfaces/ws/common.py @@ -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", ] diff --git a/interfaces/ws/events.py b/interfaces/ws/events.py index 4273be6c..ea284f34 100644 --- a/interfaces/ws/events.py +++ b/interfaces/ws/events.py @@ -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: diff --git a/interfaces/ws/protocol_schemas/connect.json b/interfaces/ws/protocol_schemas/connect.json index d8ab855e..9af59065 100644 --- a/interfaces/ws/protocol_schemas/connect.json +++ b/interfaces/ws/protocol_schemas/connect.json @@ -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, diff --git a/interfaces/ws/runtime_helpers.py b/interfaces/ws/runtime_helpers.py index 8086737f..9e16e97b 100644 --- a/interfaces/ws/runtime_helpers.py +++ b/interfaces/ws/runtime_helpers.py @@ -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"] diff --git a/interfaces/ws/runtime_impl.py b/interfaces/ws/runtime_impl.py index 009efc53..1159bfb8 100644 --- a/interfaces/ws/runtime_impl.py +++ b/interfaces/ws/runtime_impl.py @@ -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( diff --git a/interfaces/ws/runtime_loop.py b/interfaces/ws/runtime_loop.py index f02c30d4..33b7d448 100644 --- a/interfaces/ws/runtime_loop.py +++ b/interfaces/ws/runtime_loop.py @@ -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: diff --git a/platform/config/paths.py b/platform/config/paths.py index f815e4e4..eeb63d86 100644 --- a/platform/config/paths.py +++ b/platform/config/paths.py @@ -23,6 +23,11 @@ _DATA_MIGRATION_DONE = False def _canonical_data_root() -> Path: + # Compatibility: some callers/tests set PROJECT_ROOT to repo parent while + # data lives under "/oclaw/data". + nested = (PROJECT_ROOT / "oclaw" / "data").resolve() + if nested.exists(): + return nested return (PROJECT_ROOT / "data").resolve() diff --git a/runtime/chat/tool_runtime.py b/runtime/chat/tool_runtime.py index 63f72aa4..288b3c06 100644 --- a/runtime/chat/tool_runtime.py +++ b/runtime/chat/tool_runtime.py @@ -39,7 +39,6 @@ _SQL_REPLAY_COMPACT_TOOL_NAMES = { } _TABULAR_QUERY_TOOL_NAMES = { "query_tabular_attachment", - "run_tabular_sql", "analyze_tabular_attachment_full_scan", } _TEXT_QUERY_TOOL_NAMES = { @@ -583,6 +582,14 @@ class ToolExecutor: turn_tool_name_counts, turn_tool_observed_rows = _load_turn_tool_stats() local_turn_tool_name_counts: dict[str, int] = {} local_turn_tool_observed_rows: dict[str, int] = {} + planned_sql_name_counts: dict[str, int] = {} + for tc in tool_uses: + name = str(tc.name or "") + if name in _SQL_REPLAY_COMPACT_TOOL_NAMES: + planned_sql_name_counts[name] = int(planned_sql_name_counts.get(name, 0)) + 1 + compact_sql_names = { + name for name, cnt in planned_sql_name_counts.items() if int(cnt) >= int(tool_history_summary_after_calls()) + } has_tabular_ref_in_session = _session_has_tabular_ref(ctx.store, ctx.session_id) has_text_ref_in_session = _session_has_text_ref(ctx.store, ctx.session_id) has_image_ref_in_session = _session_has_image_ref(ctx.store, ctx.session_id) @@ -759,6 +766,13 @@ class ToolExecutor: current_rows = int(local_turn_tool_observed_rows.get(tc.name, 0)) local_turn_tool_name_counts[tc.name] = current + 1 local_turn_tool_observed_rows[tc.name] = current_rows + observed_rows_this_call + if tc.name in compact_sql_names: + cumulative_rows = int(turn_tool_observed_rows.get(tc.name, 0)) + int(local_turn_tool_observed_rows.get(tc.name, 0)) + result_for_llm["_history_compacted"] = True + result_for_llm["_history_compact_reason"] = "repeated_tool_calls_in_turn" + result_for_llm["_tool_observed_rows_this_call"] = int(observed_rows_this_call) + result_for_llm["_tool_observed_rows_cumulative_in_turn"] = int(cumulative_rows) + result_for_llm["audit_note"] = "Result compacted for history replay safety." trunc_ms = int((time.perf_counter() - t_trunc) * 1000) tool_content = self._json_dumps_safe(result_for_llm) t_db2 = time.perf_counter() diff --git a/runtime/gateway.py b/runtime/gateway.py index d36c6e1c..20f55f75 100644 --- a/runtime/gateway.py +++ b/runtime/gateway.py @@ -335,7 +335,7 @@ class OclawGateway: if isinstance(dispatch, dict): instruction_text = str(dispatch.get("instruction_text") or "").strip() if not instruction_text: - return ("generalist", "manager_instruction_missing", None, "", False, None, "", "") + return ("generalist", "manager_instruction_missing", None, "", False, None, "", "", "") need_wiki_inject: bool | None = None wiki_query = "" memory_write_text = "" @@ -766,6 +766,7 @@ class OclawGateway: ws_received_ms = None trace_id = new_trace_id() rid = str(run_id or "").strip() or str(uuid.uuid4()) + executed_turn_uuid = "" ctx = OclawSessionContext( session_id=msg.session_id, tenant_id=msg.tenant_id, @@ -926,7 +927,7 @@ class OclawGateway: user_id=msg.user_id, role=msg.role, channel=msg.channel, - text=str(instruction_text or "").strip(), + text=str(manager_instruction_text or "").strip(), attachments=list(msg.attachments or []), metadata=( { diff --git a/tests/test_ws_gateway.py b/tests/test_ws_gateway.py index 44727c21..b463c3bc 100644 --- a/tests/test_ws_gateway.py +++ b/tests/test_ws_gateway.py @@ -4,6 +4,7 @@ from unittest import mock from fastapi.testclient import TestClient from oclaw.interfaces.http.fastapi_app import create_app +from oclaw.interfaces.ws.runtime_impl import OclawWsGatewayConnection from oclaw.runtime.gateway import OclawGatewayResult @@ -17,13 +18,27 @@ def _connect_params() -> dict: "caps": [], "commands": [], "permissions": {}, + "auth": {"token": "test-token"}, } class WsGatewayTests(unittest.TestCase): def setUp(self) -> None: + with OclawWsGatewayConnection._rate_lock: + OclawWsGatewayConnection._rate_by_ip.clear() + OclawWsGatewayConnection._rate_by_user.clear() + OclawWsGatewayConnection._stats.clear() + OclawWsGatewayConnection._event_buffer_by_user.clear() + self._auth_patcher = mock.patch( + "oclaw.interfaces.ws.runtime_impl.resolve_ws_auth_payload", + return_value={"tenant_id": "t1", "user_id": "u1", "role": "operator"}, + ) + self._auth_patcher.start() self.client = TestClient(create_app()) + def tearDown(self) -> None: + self._auth_patcher.stop() + def test_ws_requires_connect_first(self) -> None: with self.client.websocket_connect("/ws") as ws: # server sends connect.challenge first @@ -60,6 +75,56 @@ class WsGatewayTests(unittest.TestCase): assert res["type"] == "res" assert res["ok"] is False + def test_ws_connect_rejects_without_auth(self) -> None: + with mock.patch("oclaw.interfaces.ws.runtime_impl.resolve_ws_auth_payload", return_value={}): + with self.client.websocket_connect("/ws") as ws: + ws.receive_json() + bad = _connect_params() + bad["auth"] = {} + ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": bad}) + res = None + for _ in range(2): + msg = ws.receive_json() + if msg.get("type") == "res" and msg.get("id") == "c1": + res = msg + break + assert res is not None + assert res["type"] == "res" + assert res["ok"] is False + assert (res.get("error") or {}).get("code") == "UNAUTHORIZED" + + def test_ws_origin_blocked(self) -> None: + with mock.patch("oclaw.interfaces.ws.runtime_impl.origin_is_allowed", return_value=False): + with self.assertRaises(Exception): + with self.client.websocket_connect("/ws", headers={"origin": "https://evil.example.com"}): + pass + + def test_ws_rate_limited(self) -> None: + with mock.patch("oclaw.interfaces.ws.runtime_impl.WS_RATE_LIMIT_CONN_PER_WINDOW", 1): + with self.client.websocket_connect("/ws") as ws: + ws.receive_json() + ws.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()}) + ws.receive_json() + ws.send_json({"type": "req", "id": "r1", "method": "sessions.list", "params": {}}) + first = None + for _ in range(8): + msg = ws.receive_json() + if msg.get("type") == "res" and msg.get("id") == "r1": + first = msg + break + assert first is not None + assert first["ok"] is True + ws.send_json({"type": "req", "id": "r2", "method": "sessions.list", "params": {}}) + second = None + for _ in range(8): + msg = ws.receive_json() + if msg.get("type") == "res" and msg.get("id") == "r2": + second = msg + break + assert second is not None + assert second["ok"] is False + assert (second.get("error") or {}).get("code") == "RATE_LIMITED" + def test_ws_agent_run_emits_events_and_response(self) -> None: def _fake_handle_turn(self, **kwargs): # noqa: ANN001 on_progress = kwargs.get("on_progress") @@ -233,7 +298,13 @@ class WsGatewayTests(unittest.TestCase): "params": {"sessionKey": "sess-chat", "message": "hi", "idempotencyKey": "idem-chat-1"}, } ) - ack = ws.receive_json() + ack = None + for _ in range(12): + msg = ws.receive_json() + if msg.get("type") == "res" and msg.get("id") == "cs1": + ack = msg + break + assert ack is not None assert ack.get("type") == "res" assert ack.get("id") == "cs1" assert ack.get("ok") is True @@ -291,7 +362,13 @@ class WsGatewayTests(unittest.TestCase): "params": {"sessionKey": "sess-chat2", "message": "hi", "idempotencyKey": "idem-chat-2"}, } ) - ack = ws.receive_json() + ack = None + for _ in range(12): + msg = ws.receive_json() + if msg.get("type") == "res" and msg.get("id") == "cs2": + ack = msg + break + assert ack is not None assert ack.get("type") == "res" assert ack.get("id") == "cs2" assert ack.get("ok") is True @@ -318,3 +395,60 @@ class WsGatewayTests(unittest.TestCase): assert seen_tool is True assert seen_final is True + def test_ws_replay_events_with_last_seq(self) -> None: + def _fake_handle_turn(self, **kwargs): # noqa: ANN001 + rid = kwargs.get("run_id") or "run_replay" + on_token = kwargs.get("on_token") + if callable(on_token): + on_token("re") + on_token("play") + return OclawGatewayResult( + run_id=str(rid), + reply_text="replay", + trace_id="trace_replay", + elapsed_ms=3, + mode="sync_direct", + task_id=None, + selected_specialist="generalist", + interaction_mode="comprehensive", + ) + + with mock.patch("oclaw.runtime.gateway.OclawGateway.handle_turn", new=_fake_handle_turn): + with self.client.websocket_connect("/ws") as ws1: + ws1.receive_json() + ws1.send_json({"type": "req", "id": "c1", "method": "connect", "params": _connect_params()}) + ws1.receive_json() + ws1.send_json( + { + "type": "req", + "id": "cs1", + "method": "chat.send", + "params": {"sessionKey": "sess-replay", "message": "hi", "idempotencyKey": "idem-replay-1"}, + } + ) + seq_seen = 0 + for _ in range(20): + msg = ws1.receive_json() + if msg.get("type") != "event": + continue + seq_seen = max(seq_seen, int(msg.get("seq") or 0)) + if msg.get("event") == "chat" and str((msg.get("payload") or {}).get("state") or "") == "final": + break + assert seq_seen > 0 + + p = _connect_params() + p["lastSeq"] = max(0, seq_seen - 1) + with self.client.websocket_connect("/ws") as ws2: + ws2.receive_json() + ws2.send_json({"type": "req", "id": "c2", "method": "connect", "params": p}) + ws2.receive_json() + replayed = False + for _ in range(8): + msg = ws2.receive_json() + if msg.get("type") != "event": + continue + if int(msg.get("seq") or 0) > int(p["lastSeq"]): + replayed = True + break + assert replayed is True +