mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 03:30:48 +08:00
完成 WebSocket 与 workspace 架构优化收尾,并修复全量回归阻塞。
本次将握手鉴权、Origin 校验、限流、重连补偿、发送背压与观测字段打通,同时修复 gateway 异常路径与备份清理兼容问题,确保工具历史压缩语义一致并恢复全量测试通过。 Made-with: Cursor
This commit is contained in:
parent
09836dddfa
commit
7253f6795d
12 changed files with 536 additions and 24 deletions
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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" &&
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 "<root>/oclaw/data".
|
||||
nested = (PROJECT_ROOT / "oclaw" / "data").resolve()
|
||||
if nested.exists():
|
||||
return nested
|
||||
return (PROJECT_ROOT / "data").resolve()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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=(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue