security: strengthen session handling

This commit is contained in:
Rohit P
2026-07-24 18:44:39 -07:00
parent 9cc2761998
commit 8ee0a0d082
6 changed files with 277 additions and 45 deletions
+138 -32
View File
@@ -12,6 +12,7 @@ import json
import os import os
import re import re
import uuid import uuid
from collections import deque
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
from typing import Any, Optional from typing import Any, Optional
@@ -39,25 +40,25 @@ def _origin_allowed(origin: str | None) -> bool:
return origin is None or bool(_ALLOWED_ORIGIN_RE.match(origin)) return origin is None or bool(_ALLOWED_ORIGIN_RE.match(origin))
# Caps on a single inbound `user_message` frame. The loopback socket is unauthenticated # Caps on inbound WebSocket traffic. The loopback socket is unauthenticated (any local
# (any local process can reach it), so an oversized frame is a cheap way to spike memory — # process can reach it), so bound frames, messages, and per-connection request rate before
# these bound one message's text and its attachments before we build content / start a turn. # building model content or starting a turn.
_WS_MAX_FRAME_BYTES = 16 * 1024 * 1024
_WS_RATE_LIMIT_COUNT = 30
_WS_RATE_LIMIT_WINDOW_SECONDS = 10.0
_MAX_MESSAGE_TEXT_CHARS = 200_000 _MAX_MESSAGE_TEXT_CHARS = 200_000
_MAX_ATTACHMENTS = 25 _MAX_ATTACHMENTS_BYTES = 15_000_000 # leaves JSON overhead below the 16 MiB frame cap
_MAX_ATTACHMENTS_BYTES = 32 * 1024 * 1024 # ~32 MB total across a message's attachments
def _attachments_size(attachments: list) -> int: def _json_value_size(value: Any) -> int:
"""Approximate on-wire byte size of a message's attachments (data URLs dominate).""" """Conservative UTF-8 size of parsed JSON without allocating another giant string."""
total = 0 if isinstance(value, str):
for a in attachments: return len(value.encode("utf-8"))
if isinstance(a, str): if isinstance(value, dict):
total += len(a) return sum(_json_value_size(k) + _json_value_size(v) for k, v in value.items())
elif isinstance(a, dict): if isinstance(value, list):
for v in a.values(): return sum(_json_value_size(v) for v in value)
if isinstance(v, str): return 8 # numbers, booleans, null, separators
total += len(v)
return total
# Brand colors for the connector badge riding the ✓ (UX-DECISIONS §30). The GUI owns the # Brand colors for the connector badge riding the ✓ (UX-DECISIONS §30). The GUI owns the
@@ -145,7 +146,13 @@ _CONNECT_FAILED_DETAIL = (
"Close this tab and try again from OpenWorker." "Close this tab and try again from OpenWorker."
) )
from ..attachments import build_user_content from ..attachments import (
MAX_ATTACHMENTS as _MAX_ATTACHMENTS,
MAX_IMAGE_CHARS,
MAX_PDF_CHARS,
MAX_TEXT_CHARS,
build_user_content,
)
from ..engine import ApprovalOutcome from ..engine import ApprovalOutcome
from ..inbox import VIS_INBOX, VIS_INLINE, args_preview from ..inbox import VIS_INBOX, VIS_INLINE, args_preview
from ..permissions import Mode from ..permissions import Mode
@@ -1618,9 +1625,8 @@ def create_app(manager: SessionManager) -> FastAPI:
} }
async def run_turn(content, *, retry: bool = False) -> None: async def run_turn(content, *, retry: bool = False) -> None:
manager.mark_running( # The receive loop atomically claims this session before scheduling the task.
session_id # Keeping the claim outside prevents two back-to-back frames from both starting.
) # busy → self-wakes steer instead of colliding
try: try:
events = engine.retry() if retry else engine.run(content) events = engine.retry() if retry else engine.run(content)
async for event in events: async for event in events:
@@ -1641,10 +1647,48 @@ def create_app(manager: SessionManager) -> FastAPI:
# This socket is now a live view of the session; background turns (channel delivery, # This socket is now a live view of the session; background turns (channel delivery,
# self-wake, durable resume) broadcast here too, not just locally driven run_turns. # self-wake, durable resume) broadcast here too, not just locally driven run_turns.
manager.register_session_client(session_id, ws.send_json) manager.register_session_client(session_id, ws.send_json)
inbound_times: deque[float] = deque()
async def reject_input(reason: str) -> None:
# Input validation failures are not provider failures and must not offer "Retry"
# or flush an in-progress assistant stream in the GUI.
await ws.send_json({"type": "input_rejected", "data": {"error": reason}})
async def claim_turn(*, retry: bool = False, content=None) -> None:
if not manager.try_mark_running(session_id):
await reject_input(
"This session is already running a turn. Wait for it to finish or stop it."
)
return
asyncio.create_task(run_turn(content, retry=retry))
try: try:
while True: while True:
try:
message = await ws.receive_json() message = await ws.receive_json()
except (json.JSONDecodeError, UnicodeDecodeError):
await reject_input("Invalid WebSocket message: expected JSON.")
continue
now = asyncio.get_running_loop().time()
while (
inbound_times
and now - inbound_times[0] > _WS_RATE_LIMIT_WINDOW_SECONDS
):
inbound_times.popleft()
if len(inbound_times) >= _WS_RATE_LIMIT_COUNT:
await reject_input("Too many WebSocket messages; reconnect and try again.")
await ws.close(code=1008)
return
inbound_times.append(now)
if not isinstance(message, dict):
await reject_input("Invalid WebSocket message: expected an object.")
continue
kind = message.get("type") kind = message.get("type")
if not isinstance(kind, str):
await reject_input("Invalid WebSocket message: missing string type.")
continue
if kind == "approval": if kind == "approval":
_resolve_pending(message.get("decision", "deny")) _resolve_pending(message.get("decision", "deny"))
elif kind == "directory_response": elif kind == "directory_response":
@@ -1674,22 +1718,33 @@ def create_app(manager: SessionManager) -> FastAPI:
elif kind == "retry": elif kind == "retry":
# Re-run after a provider error (engine guards on the error-notice # Re-run after a provider error (engine guards on the error-notice
# tail, so a stray frame is a no-op that still ends with turn_done). # tail, so a stray frame is a no-op that still ends with turn_done).
if not manager.is_running(session_id): await claim_turn(retry=True)
asyncio.create_task(run_turn(None, retry=True))
elif kind == "set_mode": elif kind == "set_mode":
try: try:
engine.permissions.mode = Mode(message.get("mode")) engine.permissions.mode = Mode(message.get("mode"))
except ValueError: except (TypeError, ValueError):
pass pass
elif kind == "set_model": elif kind == "set_model":
await _apply_model(message.get("model")) model = message.get("model")
if model is not None and not isinstance(model, str):
await reject_input("Invalid model: expected a string.")
else:
await _apply_model(model)
elif kind == "user_message": elif kind == "user_message":
text = (message.get("text") or "").strip() raw_text = message.get("text")
attachments = message.get("attachments") or [] if raw_text is None:
raw_text = ""
if not isinstance(raw_text, str):
await reject_input("Invalid message text: expected a string.")
continue
text = raw_text.strip()
raw_attachments = message.get("attachments")
attachments = [] if raw_attachments is None else raw_attachments
# Reject an oversized frame instead of buffering it into a turn. Send a # Reject an oversized frame instead of buffering it into a turn. Send a
# visible error so the surface can tell the user, and drop the message. # visible error so the surface can tell the user, and drop the message.
if not isinstance(attachments, list): if not isinstance(attachments, list):
attachments = [] await reject_input("Invalid attachments: expected a list.")
continue
reject = None reject = None
if len(text) > _MAX_MESSAGE_TEXT_CHARS: if len(text) > _MAX_MESSAGE_TEXT_CHARS:
reject = ( reject = (
@@ -1701,18 +1756,69 @@ def create_app(manager: SessionManager) -> FastAPI:
f"Too many attachments ({len(attachments)}; " f"Too many attachments ({len(attachments)}; "
f"limit {_MAX_ATTACHMENTS})." f"limit {_MAX_ATTACHMENTS})."
) )
elif _attachments_size(attachments) > _MAX_ATTACHMENTS_BYTES: elif any(not isinstance(a, dict) for a in attachments):
reject = "Attachments too large (limit 32 MB per message)." reject = "Invalid attachment: expected an object."
elif _json_value_size(attachments) > _MAX_ATTACHMENTS_BYTES:
reject = "Attachments too large (limit 15 MB per message)."
else:
for attachment in attachments:
attachment_kind = attachment.get("kind")
name = attachment.get("name")
mime = attachment.get("mime")
if attachment_kind not in {"image", "pdf", "text"}:
reject = "Invalid attachment kind."
elif name is not None and (
not isinstance(name, str) or len(name) > 1024
):
reject = "Invalid attachment name."
elif mime is not None and (
not isinstance(mime, str) or len(mime) > 255
):
reject = "Invalid attachment MIME type."
elif attachment_kind == "image":
data = attachment.get("data_url")
if (
not isinstance(data, str)
or not data.startswith("data:image/")
or ";base64," not in data
or len(data) > MAX_IMAGE_CHARS
):
reject = "Invalid or oversized image attachment."
elif attachment_kind == "pdf":
data = attachment.get("data_url")
if (
not isinstance(data, str)
or not data.startswith(
"data:application/pdf;base64,"
)
or len(data) > MAX_PDF_CHARS
):
reject = "Invalid or oversized PDF attachment."
else:
body = attachment.get("text")
if (
not isinstance(body, str)
or len(body) > MAX_TEXT_CHARS
):
reject = "Invalid or oversized text attachment."
if reject is not None: if reject is not None:
await ws.send_json({"type": "error", "data": {"error": reject}}) break
if reject is not None:
await reject_input(reject)
continue continue
# The composer sends its visible model with every message — the FIRST # The composer sends its visible model with every message — the FIRST
# one binds the session (race-proof across reconnects; see api.ts # one binds the session (race-proof across reconnects; see api.ts
# Session.userMessage), later ones may switch it (notice persisted). # Session.userMessage), later ones may switch it (notice persisted).
await _apply_model(message.get("model")) model = message.get("model")
if model is not None and not isinstance(model, str):
await reject_input("Invalid model: expected a string.")
continue
await _apply_model(model)
if text or attachments: if text or attachments:
content = build_user_content(text, attachments) content = build_user_content(text, attachments)
asyncio.create_task(run_turn(content)) await claim_turn(content=content)
else:
await reject_input(f"Unknown WebSocket message type: {kind}.")
except WebSocketDisconnect: except WebSocketDisconnect:
pass pass
finally: finally:
+10 -6
View File
@@ -2465,6 +2465,13 @@ class SessionManager:
def mark_running(self, session_id: str) -> None: def mark_running(self, session_id: str) -> None:
self._running_sessions.add(session_id) self._running_sessions.add(session_id)
def try_mark_running(self, session_id: str) -> bool:
"""Atomically claim an idle session for one turn on the server event loop."""
if session_id in self._running_sessions:
return False
self._running_sessions.add(session_id)
return True
def mark_idle(self, session_id: str) -> None: def mark_idle(self, session_id: str) -> None:
self._running_sessions.discard(session_id) self._running_sessions.discard(session_id)
# Every turn path (WS, background delivery, durable resume) marks idle when it # Every turn path (WS, background delivery, durable resume) marks idle when it
@@ -2488,15 +2495,12 @@ class SessionManager:
by self-wake and channel-subscription delivery. `source` is the display-only MessageSource by self-wake and channel-subscription delivery. `source` is the display-only MessageSource
sidecar for connector messages (framed `message` stays the model-facing text). sidecar for connector messages (framed `message` stays the model-facing text).
""" """
if self.is_running(session_id):
engine = self._engines.get(session_id)
if engine is not None:
engine.queue_steering(message, source)
return
engine = self.get_engine(session_id) engine = self.get_engine(session_id)
if engine is None: if engine is None:
return return
self.mark_running(session_id) if not self.try_mark_running(session_id):
engine.queue_steering(message, source)
return
try: try:
async for event in engine.run(message, source=source): async for event in engine.run(message, source=source):
# Stream every event to any socket viewing this session, so a background turn # Stream every event to any socket viewing this session, so a background turn
+4 -2
View File
@@ -10,7 +10,7 @@ from pathlib import Path
from ..config import load_config from ..config import load_config
from ..permissions import Mode from ..permissions import Mode
from ..secrets import state_dir from ..secrets import state_dir
from .app import create_app from .app import _WS_MAX_FRAME_BYTES, create_app
from .manager import SessionManager from .manager import SessionManager
@@ -149,7 +149,9 @@ def main(argv=None) -> None:
_exit_when_orphaned() _exit_when_orphaned()
app = build_app(args.cwd, args.model, args.mode) app = build_app(args.cwd, args.model, args.mode)
uvicorn.run(app, host=args.host, port=args.port) uvicorn.run(
app, host=args.host, port=args.port, ws_max_size=_WS_MAX_FRAME_BYTES
)
if __name__ == "__main__": if __name__ == "__main__":
+6
View File
@@ -695,6 +695,12 @@ export function App() {
{ kind: "notice", tone: "warn", text: "Error: " + (d.error || "unknown"), retriable: true }, { kind: "notice", tone: "warn", text: "Error: " + (d.error || "unknown"), retriable: true },
]); ]);
break; break;
case "input_rejected":
setItems((p) => [
...p,
{ kind: "notice", tone: "warn", text: d.error || "That message was rejected." },
]);
break;
case "turn_done": case "turn_done":
setRunning(false); setRunning(false);
refreshSessions(); refreshSessions();
+1
View File
@@ -15,6 +15,7 @@ export type EventType =
| "iteration_end" | "iteration_end"
| "turn_end" | "turn_end"
| "error" | "error"
| "input_rejected"
| "interrupted" | "interrupted"
| "model_changed" | "model_changed"
| "turn_done"; | "turn_done";
+117 -4
View File
@@ -325,19 +325,22 @@ def test_ws_simple_turn(tmp_path):
def test_ws_rejects_oversized_message(tmp_path): def test_ws_rejects_oversized_message(tmp_path):
from coworker.server import app as app_mod from coworker.server import app as app_mod
from coworker.attachments import MAX_ATTACHMENTS
client = _client(tmp_path, [_text("should not run")]) client = _client(tmp_path, [_text("should not run")])
with client.websocket_connect("/ws/session/big") as ws: with client.websocket_connect("/ws/session/big") as ws:
assert ws.receive_json()["type"] == "ready" assert ws.receive_json()["type"] == "ready"
# Oversized text → single error frame, no turn runs. # Oversized text → single input-rejected frame, no turn runs.
ws.send_json( ws.send_json(
{"type": "user_message", "text": "x" * (app_mod._MAX_MESSAGE_TEXT_CHARS + 1)} {"type": "user_message", "text": "x" * (app_mod._MAX_MESSAGE_TEXT_CHARS + 1)}
) )
evt = ws.receive_json() evt = ws.receive_json()
assert evt["type"] == "error" and "too long" in evt["data"]["error"].lower() assert evt["type"] == "input_rejected"
assert "too long" in evt["data"]["error"].lower()
# Too many attachments → error frame. # The ingress cap is the same cap the attachment builder enforces.
assert app_mod._MAX_ATTACHMENTS == MAX_ATTACHMENTS
ws.send_json( ws.send_json(
{ {
"type": "user_message", "type": "user_message",
@@ -346,13 +349,123 @@ def test_ws_rejects_oversized_message(tmp_path):
} }
) )
evt = ws.receive_json() evt = ws.receive_json()
assert evt["type"] == "error" and "attachment" in evt["data"]["error"].lower() assert evt["type"] == "input_rejected"
assert "attachment" in evt["data"]["error"].lower()
# A normal message still works afterwards (the socket wasn't torn down). # A normal message still works afterwards (the socket wasn't torn down).
ws.send_json({"type": "user_message", "text": "hello"}) ws.send_json({"type": "user_message", "text": "hello"})
assert "turn_done" in _drain(ws) assert "turn_done" in _drain(ws)
def test_ws_rejects_malformed_payloads_without_killing_socket(tmp_path):
client = _client(tmp_path, [_text("normal")])
with client.websocket_connect("/ws/session/malformed") as ws:
assert ws.receive_json()["type"] == "ready"
invalid = [
[],
{"type": "user_message", "text": ["not", "text"]},
{"type": "user_message", "text": "x", "attachments": {}},
{
"type": "user_message",
"text": "x",
"attachments": [{"kind": "image", "data_url": "https://example.com/x"}],
},
{"type": "set_model", "model": {"unexpected": True}},
{"type": "unknown"},
]
for payload in invalid:
ws.send_json(payload)
evt = ws.receive_json()
assert evt["type"] == "input_rejected"
ws.send_json({"type": "user_message", "text": "still works"})
assert "turn_done" in _drain(ws)
def test_ws_allows_only_one_inflight_turn_per_session(tmp_path):
import threading
import time
class SlowProvider(ProviderClient):
def __init__(self):
self._lock = threading.Lock()
self.active = 0
self.max_active = 0
def complete(self, *, model, messages, tools=None, **settings):
with self._lock:
self.active += 1
self.max_active = max(self.max_active, self.active)
try:
time.sleep(0.08)
return _text("done")
finally:
with self._lock:
self.active -= 1
def capabilities(self, model):
return ModelCapabilities()
provider = SlowProvider()
manager = SessionManager(workspace=tmp_path, provider=provider)
client = TestClient(create_app(manager))
with client.websocket_connect("/ws/session/serialized") as ws:
assert ws.receive_json()["type"] == "ready"
ws.send_json({"type": "user_message", "text": "first"})
ws.send_json({"type": "user_message", "text": "second"})
types = []
while "turn_done" not in types:
types.append(ws.receive_json()["type"])
assert "input_rejected" in types
assert provider.max_active == 1
engine = manager._engines["serialized"]
user_messages = [m for m in engine.messages if m.get("role") == "user"]
assert [m["content"] for m in user_messages] == ["first"]
def test_ws_rate_limits_inbound_frames(tmp_path):
from coworker.server import app as app_mod
from starlette.websockets import WebSocketDisconnect
client = _client(tmp_path, [])
with pytest.raises(WebSocketDisconnect):
with client.websocket_connect("/ws/session/rate") as ws:
assert ws.receive_json()["type"] == "ready"
for _ in range(app_mod._WS_RATE_LIMIT_COUNT):
ws.send_json({"type": "unknown"})
assert ws.receive_json()["type"] == "input_rejected"
ws.send_json({"type": "unknown"})
assert ws.receive_json()["type"] == "input_rejected"
ws.receive_json()
def test_server_sets_explicit_websocket_frame_limit(tmp_path, monkeypatch):
import sys
from types import SimpleNamespace
from coworker.server import run as server_run
seen = {}
fake_app = object()
monkeypatch.setattr(server_run, "_ensure_ca_bundle", lambda: None)
monkeypatch.setattr(server_run, "_exit_when_orphaned", lambda: None)
monkeypatch.setattr(server_run, "build_app", lambda *args: fake_app)
monkeypatch.setitem(
sys.modules,
"uvicorn",
SimpleNamespace(run=lambda app, **kwargs: seen.update(app=app, **kwargs)),
)
server_run.main(["--cwd", str(tmp_path), "--port", "8766"])
assert seen["app"] is fake_app
assert seen["ws_max_size"] == server_run._WS_MAX_FRAME_BYTES
def test_ws_error_persists_notice_and_retry_reruns(tmp_path): def test_ws_error_persists_notice_and_retry_reruns(tmp_path):
class FlakyProvider(ProviderClient): class FlakyProvider(ProviderClient):
def __init__(self): def __init__(self):