diff --git a/CLAUDE.md b/CLAUDE.md index d42483682e..7e260997a6 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -114,7 +114,8 @@ curl http://localhost:3000/api/v1/health # backend (via web proxy) ```text src/synthorg/ - api/ # Litestar REST + WebSocket API (controllers, guards, channels, JWT + API key auth, approval gate integration, coordination endpoint, collaboration endpoint, settings endpoint, RFC 9457 structured errors (ErrorCategory, ErrorCode, ErrorDetail, ProblemDetail, CATEGORY_TITLES, category_title, category_type_uri, content negotiation)) + api/ # Litestar REST + WebSocket API (controllers, guards, channels, JWT + API key + WS ticket auth, approval gate integration, coordination endpoint, collaboration endpoint, settings endpoint, RFC 9457 structured errors (ErrorCategory, ErrorCode, ErrorDetail, ProblemDetail, CATEGORY_TITLES, category_title, category_type_uri, content negotiation)) + auth/ # Authentication subpackage (controller, service, middleware, JWT + API key + WS ticket store, models, config) budget/ # Cost tracking, budget enforcement (pre-flight/in-flight checks, auto-downgrade), billing periods, cost tiers, quota/subscription tracking, CFO cost optimization (anomaly detection, efficiency analysis, downgrade recommendations, approval decisions), spending reports, budget errors (BudgetExhaustedError, DailyLimitExceededError, QuotaExhaustedError) cli/ # Python CLI module (superseded by top-level cli/ Go binary) communication/ # Message bus, dispatcher, messenger, channels, delegation, loop prevention, conflict resolution @@ -195,7 +196,7 @@ site/ # Astro landing page (synthorg.io) - **Every module** with business logic MUST have: `from synthorg.observability import get_logger` then `logger = get_logger(__name__)` - **Never** use `import logging` / `logging.getLogger()` / `print()` in application code - **Variable name**: always `logger` (not `_logger`, not `log`) -- **Event names**: always use constants from the domain-specific module under `synthorg.observability.events` (e.g., `PROVIDER_CALL_START` from `events.provider`, `BUDGET_RECORD_ADDED` from `events.budget`, `CFO_ANOMALY_DETECTED` from `events.cfo`, `CONFLICT_DETECTED` from `events.conflict`, `MEETING_STARTED` from `events.meeting`, `MEETING_SCHEDULER_STARTED` from `events.meeting`, `MEETING_SCHEDULER_ERROR` from `events.meeting`, `MEETING_SCHEDULER_STOPPED` from `events.meeting`, `MEETING_PERIODIC_TRIGGERED` from `events.meeting`, `MEETING_EVENT_TRIGGERED` from `events.meeting`, `MEETING_PARTICIPANTS_RESOLVED` from `events.meeting`, `MEETING_NO_PARTICIPANTS` from `events.meeting`, `MEETING_NOT_FOUND` from `events.meeting`, `CLASSIFICATION_START` from `events.classification`, `CONSOLIDATION_START` from `events.consolidation`, `ORG_MEMORY_QUERY_START` from `events.org_memory`, `API_REQUEST_STARTED` from `events.api`, `API_REQUEST_COMPLETED` from `events.api`, `API_REQUEST_ERROR` from `events.api`, `API_ROUTE_NOT_FOUND` from `events.api`, `API_HEALTH_CHECK` from `events.api`, `API_COORDINATION_STARTED` from `events.api`, `API_COORDINATION_COMPLETED` from `events.api`, `API_COORDINATION_FAILED` from `events.api`, `API_COORDINATION_AGENT_RESOLVE_FAILED` from `events.api`, `API_CONTENT_NEGOTIATED` from `events.api`, `API_CORRELATION_FALLBACK` from `events.api`, `API_ACCEPT_PARSE_FAILED` from `events.api`, `CODE_RUNNER_EXECUTE_START` from `events.code_runner`, `DOCKER_EXECUTE_START` from `events.docker`, `MCP_INVOKE_START` from `events.mcp`, `SECURITY_EVALUATE_START` from `events.security`, `HR_HIRING_REQUEST_CREATED` from `events.hr`, `PERF_METRIC_RECORDED` from `events.performance`, `PERF_LLM_SAMPLE_STARTED` from `events.performance`, `PERF_LLM_SAMPLE_COMPLETED` from `events.performance`, `PERF_LLM_SAMPLE_FAILED` from `events.performance`, `PERF_OVERRIDE_SET` from `events.performance`, `PERF_OVERRIDE_CLEARED` from `events.performance`, `PERF_OVERRIDE_APPLIED` from `events.performance`, `PERF_OVERRIDE_EXPIRED` from `events.performance`, `TRUST_EVALUATE_START` from `events.trust`, `PROMOTION_EVALUATE_START` from `events.promotion`, `PROMPT_BUILD_START` from `events.prompt`, `MEMORY_RETRIEVAL_START` from `events.memory`, `MEMORY_BACKEND_CONNECTED` from `events.memory`, `MEMORY_ENTRY_STORED` from `events.memory`, `MEMORY_BACKEND_SYSTEM_ERROR` from `events.memory`, `MEMORY_RRF_FUSION_COMPLETE` from `events.memory`, `MEMORY_RRF_VALIDATION_FAILED` from `events.memory`, `AUTONOMY_ACTION_AUTO_APPROVED` from `events.autonomy`, `TIMEOUT_POLICY_EVALUATED` from `events.timeout`, `PERSISTENCE_AUDIT_ENTRY_SAVED` from `events.persistence`, `TASK_ENGINE_STARTED` from `events.task_engine`, `COORDINATION_STARTED` from `events.coordination`, `COORDINATION_FACTORY_BUILT` from `events.coordination`, `COMMUNICATION_DISPATCH_START` from `events.communication`, `COMPANY_STARTED` from `events.company`, `CONFIG_LOADED` from `events.config`, `CORRELATION_ID_CREATED` from `events.correlation`, `DECOMPOSITION_STARTED` from `events.decomposition`, `DELEGATION_STARTED` from `events.delegation`, `EXECUTION_LOOP_START` from `events.execution`, `CHECKPOINT_SAVED` from `events.checkpoint`, `PERSISTENCE_CHECKPOINT_SAVED` from `events.persistence`, `GIT_OPERATION_START` from `events.git`, `PARALLEL_GROUP_START` from `events.parallel`, `PERSONALITY_LOADED` from `events.personality`, `QUOTA_CHECKED` from `events.quota`, `ROLE_ASSIGNED` from `events.role`, `ROUTING_STARTED` from `events.routing`, `SANDBOX_EXECUTE_START` from `events.sandbox`, `TASK_CREATED` from `events.task`, `TASK_ASSIGNMENT_STARTED` from `events.task_assignment`, `TASK_ROUTING_STARTED` from `events.task_routing`, `TEMPLATE_LOADED` from `events.template`, `TOOL_INVOKE_START` from `events.tool`, `TOOL_OUTPUT_WITHHELD` from `events.tool`, `WORKSPACE_CREATED` from `events.workspace`, `APPROVAL_GATE_ESCALATION_DETECTED` from `events.approval_gate`, `APPROVAL_GATE_ESCALATION_FAILED` from `events.approval_gate`, `APPROVAL_GATE_INITIALIZED` from `events.approval_gate`, `APPROVAL_GATE_RISK_CLASSIFIED` from `events.approval_gate`, `APPROVAL_GATE_RISK_CLASSIFY_FAILED` from `events.approval_gate`, `APPROVAL_GATE_CONTEXT_PARKED` from `events.approval_gate`, `APPROVAL_GATE_CONTEXT_PARK_FAILED` from `events.approval_gate`, `APPROVAL_GATE_PARK_TASKLESS` from `events.approval_gate`, `APPROVAL_GATE_RESUME_STARTED` from `events.approval_gate`, `APPROVAL_GATE_CONTEXT_RESUMED` from `events.approval_gate`, `APPROVAL_GATE_RESUME_FAILED` from `events.approval_gate`, `APPROVAL_GATE_RESUME_DELETE_FAILED` from `events.approval_gate`, `APPROVAL_GATE_RESUME_TRIGGERED` from `events.approval_gate`, `APPROVAL_GATE_NO_PARKED_CONTEXT` from `events.approval_gate`, `APPROVAL_GATE_LOOP_WIRING_WARNING` from `events.approval_gate`, `STAGNATION_CHECK_PERFORMED` from `events.stagnation`, `STAGNATION_DETECTED` from `events.stagnation`, `STAGNATION_CORRECTION_INJECTED` from `events.stagnation`, `STAGNATION_TERMINATED` from `events.stagnation`, `PERSISTENCE_AGENT_STATE_SAVED` from `events.persistence`, `PERSISTENCE_AGENT_STATE_FETCHED` from `events.persistence`, `PERSISTENCE_AGENT_STATE_ACTIVE_QUERIED` from `events.persistence`, `PERSISTENCE_AGENT_STATE_DELETED` from `events.persistence`, `SETTINGS_VALUE_SET` from `events.settings`, `SETTINGS_VALUE_DELETED` from `events.settings`, `SETTINGS_VALUE_RESOLVED` from `events.settings`, `SETTINGS_CACHE_INVALIDATED` from `events.settings`, `SETTINGS_ENCRYPTION_ERROR` from `events.settings`, `SETTINGS_VALIDATION_FAILED` from `events.settings`, `SETTINGS_NOTIFICATION_PUBLISHED` from `events.settings`, `SETTINGS_NOTIFICATION_FAILED` from `events.settings`, `SETTINGS_FETCH_FAILED` from `events.settings`, `SETTINGS_SET_FAILED` from `events.settings`, `SETTINGS_DELETE_FAILED` from `events.settings`, `SETTINGS_NOT_FOUND` from `events.settings`, `SETTINGS_REGISTRY_DUPLICATE` from `events.settings`, `SETTINGS_CONFIG_PATH_MISS` from `events.settings`). Import directly: `from synthorg.observability.events. import EVENT_CONSTANT` +- **Event names**: always use constants from the domain-specific module under `synthorg.observability.events` (e.g., `PROVIDER_CALL_START` from `events.provider`, `BUDGET_RECORD_ADDED` from `events.budget`, `CFO_ANOMALY_DETECTED` from `events.cfo`, `CONFLICT_DETECTED` from `events.conflict`, `MEETING_STARTED` from `events.meeting`, `MEETING_SCHEDULER_STARTED` from `events.meeting`, `MEETING_SCHEDULER_ERROR` from `events.meeting`, `MEETING_SCHEDULER_STOPPED` from `events.meeting`, `MEETING_PERIODIC_TRIGGERED` from `events.meeting`, `MEETING_EVENT_TRIGGERED` from `events.meeting`, `MEETING_PARTICIPANTS_RESOLVED` from `events.meeting`, `MEETING_NO_PARTICIPANTS` from `events.meeting`, `MEETING_NOT_FOUND` from `events.meeting`, `CLASSIFICATION_START` from `events.classification`, `CONSOLIDATION_START` from `events.consolidation`, `ORG_MEMORY_QUERY_START` from `events.org_memory`, `API_REQUEST_STARTED` from `events.api`, `API_REQUEST_COMPLETED` from `events.api`, `API_REQUEST_ERROR` from `events.api`, `API_ROUTE_NOT_FOUND` from `events.api`, `API_HEALTH_CHECK` from `events.api`, `API_COORDINATION_STARTED` from `events.api`, `API_COORDINATION_COMPLETED` from `events.api`, `API_COORDINATION_FAILED` from `events.api`, `API_COORDINATION_AGENT_RESOLVE_FAILED` from `events.api`, `API_CONTENT_NEGOTIATED` from `events.api`, `API_CORRELATION_FALLBACK` from `events.api`, `API_ACCEPT_PARSE_FAILED` from `events.api`, `API_WS_TICKET_ISSUED` from `events.api`, `API_WS_TICKET_CONSUMED` from `events.api`, `API_WS_TICKET_EXPIRED` from `events.api`, `API_WS_TICKET_INVALID` from `events.api`, `API_WS_TICKET_CLEANUP` from `events.api`, `CODE_RUNNER_EXECUTE_START` from `events.code_runner`, `DOCKER_EXECUTE_START` from `events.docker`, `MCP_INVOKE_START` from `events.mcp`, `SECURITY_EVALUATE_START` from `events.security`, `HR_HIRING_REQUEST_CREATED` from `events.hr`, `PERF_METRIC_RECORDED` from `events.performance`, `PERF_LLM_SAMPLE_STARTED` from `events.performance`, `PERF_LLM_SAMPLE_COMPLETED` from `events.performance`, `PERF_LLM_SAMPLE_FAILED` from `events.performance`, `PERF_OVERRIDE_SET` from `events.performance`, `PERF_OVERRIDE_CLEARED` from `events.performance`, `PERF_OVERRIDE_APPLIED` from `events.performance`, `PERF_OVERRIDE_EXPIRED` from `events.performance`, `TRUST_EVALUATE_START` from `events.trust`, `PROMOTION_EVALUATE_START` from `events.promotion`, `PROMPT_BUILD_START` from `events.prompt`, `MEMORY_RETRIEVAL_START` from `events.memory`, `MEMORY_BACKEND_CONNECTED` from `events.memory`, `MEMORY_ENTRY_STORED` from `events.memory`, `MEMORY_BACKEND_SYSTEM_ERROR` from `events.memory`, `MEMORY_RRF_FUSION_COMPLETE` from `events.memory`, `MEMORY_RRF_VALIDATION_FAILED` from `events.memory`, `AUTONOMY_ACTION_AUTO_APPROVED` from `events.autonomy`, `TIMEOUT_POLICY_EVALUATED` from `events.timeout`, `PERSISTENCE_AUDIT_ENTRY_SAVED` from `events.persistence`, `TASK_ENGINE_STARTED` from `events.task_engine`, `COORDINATION_STARTED` from `events.coordination`, `COORDINATION_FACTORY_BUILT` from `events.coordination`, `COMMUNICATION_DISPATCH_START` from `events.communication`, `COMPANY_STARTED` from `events.company`, `CONFIG_LOADED` from `events.config`, `CORRELATION_ID_CREATED` from `events.correlation`, `DECOMPOSITION_STARTED` from `events.decomposition`, `DELEGATION_STARTED` from `events.delegation`, `EXECUTION_LOOP_START` from `events.execution`, `CHECKPOINT_SAVED` from `events.checkpoint`, `PERSISTENCE_CHECKPOINT_SAVED` from `events.persistence`, `GIT_OPERATION_START` from `events.git`, `PARALLEL_GROUP_START` from `events.parallel`, `PERSONALITY_LOADED` from `events.personality`, `QUOTA_CHECKED` from `events.quota`, `ROLE_ASSIGNED` from `events.role`, `ROUTING_STARTED` from `events.routing`, `SANDBOX_EXECUTE_START` from `events.sandbox`, `TASK_CREATED` from `events.task`, `TASK_ASSIGNMENT_STARTED` from `events.task_assignment`, `TASK_ROUTING_STARTED` from `events.task_routing`, `TEMPLATE_LOADED` from `events.template`, `TOOL_INVOKE_START` from `events.tool`, `TOOL_OUTPUT_WITHHELD` from `events.tool`, `WORKSPACE_CREATED` from `events.workspace`, `APPROVAL_GATE_ESCALATION_DETECTED` from `events.approval_gate`, `APPROVAL_GATE_ESCALATION_FAILED` from `events.approval_gate`, `APPROVAL_GATE_INITIALIZED` from `events.approval_gate`, `APPROVAL_GATE_RISK_CLASSIFIED` from `events.approval_gate`, `APPROVAL_GATE_RISK_CLASSIFY_FAILED` from `events.approval_gate`, `APPROVAL_GATE_CONTEXT_PARKED` from `events.approval_gate`, `APPROVAL_GATE_CONTEXT_PARK_FAILED` from `events.approval_gate`, `APPROVAL_GATE_PARK_TASKLESS` from `events.approval_gate`, `APPROVAL_GATE_RESUME_STARTED` from `events.approval_gate`, `APPROVAL_GATE_CONTEXT_RESUMED` from `events.approval_gate`, `APPROVAL_GATE_RESUME_FAILED` from `events.approval_gate`, `APPROVAL_GATE_RESUME_DELETE_FAILED` from `events.approval_gate`, `APPROVAL_GATE_RESUME_TRIGGERED` from `events.approval_gate`, `APPROVAL_GATE_NO_PARKED_CONTEXT` from `events.approval_gate`, `APPROVAL_GATE_LOOP_WIRING_WARNING` from `events.approval_gate`, `STAGNATION_CHECK_PERFORMED` from `events.stagnation`, `STAGNATION_DETECTED` from `events.stagnation`, `STAGNATION_CORRECTION_INJECTED` from `events.stagnation`, `STAGNATION_TERMINATED` from `events.stagnation`, `PERSISTENCE_AGENT_STATE_SAVED` from `events.persistence`, `PERSISTENCE_AGENT_STATE_FETCHED` from `events.persistence`, `PERSISTENCE_AGENT_STATE_ACTIVE_QUERIED` from `events.persistence`, `PERSISTENCE_AGENT_STATE_DELETED` from `events.persistence`, `SETTINGS_VALUE_SET` from `events.settings`, `SETTINGS_VALUE_DELETED` from `events.settings`, `SETTINGS_VALUE_RESOLVED` from `events.settings`, `SETTINGS_CACHE_INVALIDATED` from `events.settings`, `SETTINGS_ENCRYPTION_ERROR` from `events.settings`, `SETTINGS_VALIDATION_FAILED` from `events.settings`, `SETTINGS_NOTIFICATION_PUBLISHED` from `events.settings`, `SETTINGS_NOTIFICATION_FAILED` from `events.settings`, `SETTINGS_FETCH_FAILED` from `events.settings`, `SETTINGS_SET_FAILED` from `events.settings`, `SETTINGS_DELETE_FAILED` from `events.settings`, `SETTINGS_NOT_FOUND` from `events.settings`, `SETTINGS_REGISTRY_DUPLICATE` from `events.settings`, `SETTINGS_CONFIG_PATH_MISS` from `events.settings`). Import directly: `from synthorg.observability.events. import EVENT_CONSTANT` - **Structured kwargs**: always `logger.info(EVENT, key=value)` — never `logger.info("msg %s", val)` - **All error paths** must log at WARNING or ERROR with context before raising - **All state transitions** must log at INFO diff --git a/docs/architecture/index.md b/docs/architecture/index.md index 30425df68c..fefbe0a139 100644 --- a/docs/architecture/index.md +++ b/docs/architecture/index.md @@ -37,7 +37,7 @@ graph TB | **budget** | Cost management — cost tracking, budget enforcement (pre-flight/in-flight), auto-downgrade, quota/subscription, CFO optimizer, spending reports | | **hr** | Agent lifecycle — hiring, firing, onboarding, offboarding, registry, performance tracking, promotion/demotion | | **tools** | Tool system — registry, built-in tools (file system, git, sandbox, code runner), MCP bridge, role-based access | -| **api** | REST + WebSocket API — Litestar controllers, JWT + API key auth, guards, channels, RFC 9457 structured error responses | +| **api** | REST + WebSocket API — Litestar controllers, JWT + API key + WS ticket auth, guards, channels, RFC 9457 structured error responses | | **config** | Company configuration — YAML schema, loader, validation, defaults | | **templates** | Pre-built company templates — personality presets, template builder | | **persistence** | Operational data — pluggable backend protocol, SQLite implementation | diff --git a/docs/design/operations.md b/docs/design/operations.md index 888fbfdd09..d8f3b1fa62 100644 --- a/docs/design/operations.md +++ b/docs/design/operations.md @@ -960,7 +960,7 @@ future CLI tool are thin clients that call the API -- they contain no business l | Endpoint | Purpose | |----------|---------| | `/api/v1/health` | Health check, readiness | -| `/api/v1/auth` | Authentication: setup, login, password change | +| `/api/v1/auth` | Authentication: setup, login, password change, ws-ticket | | `/api/v1/company` | CRUD company config | | `/api/v1/agents` | List, hire, fire, modify agents | | `/api/v1/departments` | Department management | @@ -974,7 +974,8 @@ future CLI tool are thin clients that call the API -- they contain no business l | `/api/v1/approvals` | Pending human approvals queue | | `/api/v1/analytics` | Performance metrics, dashboards | | `/api/v1/providers` | Model provider status, config | -| `/api/v1/ws` | WebSocket for real-time updates | +| `/api/v1/ws` | WebSocket for real-time updates (ticket auth via `?ticket=`) | +| `POST /api/v1/auth/ws-ticket` | Exchange JWT for one-time WebSocket connection ticket | ### Error Response Format (RFC 9457) diff --git a/docs/roadmap/index.md b/docs/roadmap/index.md index f3b90528a5..1038bf7fe1 100644 --- a/docs/roadmap/index.md +++ b/docs/roadmap/index.md @@ -12,7 +12,7 @@ The SynthOrg core framework is complete. The following subsystems are built and - Security and approval system (rule engine, output scanning, progressive trust, autonomy levels, timeout policies) - Tool system (file system, git, code runner, MCP bridge, sandboxing, permissions) - HR engine (hiring, firing, onboarding, offboarding, registry, performance tracking, promotions) -- REST and WebSocket API (Litestar controllers, JWT + API key auth, WebSocket channels) +- REST and WebSocket API (Litestar controllers, JWT + API key + WS ticket auth, WebSocket channels) - Persistence layer (pluggable protocol, SQLite backend, repository protocols) - Observability (structured logging, correlation tracking, per-domain event constants) - Configuration (YAML loading, Pydantic validation, company templates with inheritance) diff --git a/docs/security.md b/docs/security.md index d137e72d6c..ce5f637a56 100644 --- a/docs/security.md +++ b/docs/security.md @@ -72,6 +72,7 @@ on resolution. - **Argon2id password hashing** (time_cost=3, memory_cost=64 MB, parallelism=4) - **Timing-attack prevention** — dummy hash computation for non-existent users - **Forced password change** — `must_change_password` flag blocks API access +- **One-time WebSocket tickets** — short-lived (30 s), single-use, cryptographically random tokens exchanged via ``POST /api/v1/auth/ws-ticket`` (requires valid JWT). Prevents long-lived JWT leakage by replacing it with an ephemeral ticket in the WebSocket query parameter. In-memory store, monotonic clock expiry, per-process scope. - **Rate limiting** — configurable per-deployment (default: 100 req/min) ### Security Headers diff --git a/src/synthorg/api/app.py b/src/synthorg/api/app.py index 27940336bd..40af765d4e 100644 --- a/src/synthorg/api/app.py +++ b/src/synthorg/api/app.py @@ -5,6 +5,8 @@ lifecycle hooks (startup/shutdown). """ +import asyncio +import contextlib import os import time from datetime import UTC, datetime @@ -53,6 +55,7 @@ API_APP_SHUTDOWN, API_APP_STARTUP, API_APPROVAL_PUBLISH_FAILED, + API_WS_TICKET_CLEANUP, ) from synthorg.persistence.config import PersistenceConfig, SQLiteConfig from synthorg.persistence.factory import create_backend @@ -146,6 +149,22 @@ def _on_meeting_event( return _on_meeting_event +async def _ticket_cleanup_loop(app_state: AppState) -> None: + """Periodically prune expired WS tickets (runs as background task).""" + while True: + await asyncio.sleep(60) + try: + app_state.ticket_store.cleanup_expired() + except MemoryError, RecursionError: + raise + except Exception: + logger.warning( + API_WS_TICKET_CLEANUP, + error="Periodic ticket cleanup failed", + exc_info=True, + ) + + def _build_lifecycle( # noqa: PLR0913 persistence: PersistenceBackend | None, message_bus: MessageBus | None, @@ -162,8 +181,22 @@ def _build_lifecycle( # noqa: PLR0913 Returns: A tuple of (on_startup, on_shutdown) callback lists. """ + _ticket_cleanup_task: asyncio.Task[None] | None = None + + def _on_cleanup_task_done(task: asyncio.Task[None]) -> None: + """Log unexpected cleanup-task death.""" + if task.cancelled(): + return + exc = task.exception() + if exc is not None: + logger.error( + API_WS_TICKET_CLEANUP, + error="Ticket cleanup task died unexpectedly", + exc_info=exc, + ) async def on_startup() -> None: + nonlocal _ticket_cleanup_task logger.info(API_APP_STARTUP, version=__version__) await _safe_startup( persistence, @@ -173,8 +206,19 @@ async def on_startup() -> None: meeting_scheduler, app_state, ) + _ticket_cleanup_task = asyncio.create_task( + _ticket_cleanup_loop(app_state), + name="ws-ticket-cleanup", + ) + _ticket_cleanup_task.add_done_callback(_on_cleanup_task_done) async def on_shutdown() -> None: + nonlocal _ticket_cleanup_task + if _ticket_cleanup_task is not None: + _ticket_cleanup_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await _ticket_cleanup_task + _ticket_cleanup_task = None logger.info(API_APP_SHUTDOWN, version=__version__) await _safe_shutdown( task_engine, @@ -609,19 +653,25 @@ def _build_middleware(api_config: ApiConfig) -> list[Middleware]: exclude=list(rl.exclude_paths), ) auth = api_config.auth - if auth.exclude_paths is None: - prefix = api_config.api_prefix - auth = auth.model_copy( - update={ - "exclude_paths": ( - f"^{prefix}/health$", - "^/docs", - "^/api$", - f"^{prefix}/auth/setup$", - f"^{prefix}/auth/login$", - ), - }, + prefix = api_config.api_prefix + ws_path = f"^{prefix}/ws$" + exclude_paths = ( + auth.exclude_paths + if auth.exclude_paths is not None + else ( + f"^{prefix}/health$", + "^/docs", + "^/api$", + f"^{prefix}/auth/setup$", + f"^{prefix}/auth/login$", ) + ) + # Always ensure the WS upgrade path is excluded — the WS handler + # performs its own ticket-based auth, so the JWT middleware must + # not run on the upgrade request. + if ws_path not in exclude_paths: + exclude_paths = (*exclude_paths, ws_path) + auth = auth.model_copy(update={"exclude_paths": exclude_paths}) auth_middleware = create_auth_middleware_class(auth) return [ auth_middleware, diff --git a/src/synthorg/api/auth/__init__.py b/src/synthorg/api/auth/__init__.py index 73105b3547..f7590eaf2f 100644 --- a/src/synthorg/api/auth/__init__.py +++ b/src/synthorg/api/auth/__init__.py @@ -3,6 +3,7 @@ from synthorg.api.auth.config import AuthConfig from synthorg.api.auth.models import ApiKey, AuthenticatedUser, AuthMethod, User from synthorg.api.auth.service import AuthService +from synthorg.api.auth.ticket_store import TicketLimitExceededError, WsTicketStore __all__ = [ "ApiKey", @@ -10,5 +11,7 @@ "AuthMethod", "AuthService", "AuthenticatedUser", + "TicketLimitExceededError", "User", + "WsTicketStore", ] diff --git a/src/synthorg/api/auth/controller.py b/src/synthorg/api/auth/controller.py index a3b37b2dbb..c2ec6722d6 100644 --- a/src/synthorg/api/auth/controller.py +++ b/src/synthorg/api/auth/controller.py @@ -1,5 +1,6 @@ -"""Authentication controller — setup, login, password change, me.""" +"""Authentication controller — setup, login, password change, me, ws-ticket.""" +import math import uuid from datetime import UTC, datetime from typing import Any, Self @@ -10,8 +11,9 @@ from pydantic import BaseModel, ConfigDict, Field, model_validator from synthorg.api.auth.config import AuthConfig -from synthorg.api.auth.models import AuthenticatedUser, User +from synthorg.api.auth.models import AuthenticatedUser, AuthMethod, User from synthorg.api.auth.service import AuthService # noqa: TC001 +from synthorg.api.auth.ticket_store import TicketLimitExceededError from synthorg.api.dto import ApiResponse from synthorg.api.errors import ConflictError, UnauthorizedError from synthorg.api.guards import HumanRole @@ -151,6 +153,20 @@ class UserInfoResponse(BaseModel): must_change_password: bool +class WsTicketResponse(BaseModel): + """One-time WebSocket connection ticket. + + Attributes: + ticket: Single-use, short-lived ticket string. + expires_in: Ticket lifetime in seconds. + """ + + model_config = ConfigDict(frozen=True) + + ticket: NotBlankStr + expires_in: int = Field(gt=0) + + _PWD_CHANGE_EXEMPT_SUFFIXES = ("/auth/change-password", "/auth/me") # ── Guards ──────────────────────────────────────────────────── @@ -411,3 +427,62 @@ async def me( ), ), ) + + @post( + "/ws-ticket", + status_code=200, + summary="Issue a one-time WebSocket connection ticket", + ) + async def ws_ticket( + self, + request: Request[Any, Any, Any], + ) -> Response[ApiResponse[WsTicketResponse]]: + """Exchange a valid JWT for a short-lived, single-use WS ticket. + + Issue a short-lived, single-use ticket for WebSocket connections. + The ticket is passed as a query parameter instead of the JWT, so + long-lived credentials never appear in URLs or server logs. + """ + auth_user = request.scope.get("user") + if not isinstance(auth_user, AuthenticatedUser): + logger.warning( + API_AUTH_FAILED, + reason="ws_ticket_auth_required", + path=str(request.url.path), + ) + msg = "Authentication required" + raise UnauthorizedError(msg) + + if auth_user.auth_method != AuthMethod.JWT: + logger.warning( + API_AUTH_FAILED, + reason="ws_ticket_requires_jwt", + auth_method=auth_user.auth_method.value, + user_id=auth_user.user_id, + ) + msg = "WebSocket tickets require JWT authentication" + raise UnauthorizedError(msg) + + app_state = request.app.state["app_state"] + ws_user = auth_user.model_copy( + update={"auth_method": AuthMethod.WS_TICKET}, + ) + try: + ticket = app_state.ticket_store.create(ws_user) + except TicketLimitExceededError: + logger.warning( + API_AUTH_FAILED, + reason="ws_ticket_limit_exceeded", + user_id=auth_user.user_id, + ) + msg = "Too many pending tickets — wait for existing tickets to expire" + raise ConflictError(msg) # noqa: B904 + + return Response( + content=ApiResponse( + data=WsTicketResponse( + ticket=ticket, + expires_in=max(1, math.ceil(app_state.ticket_store.ttl_seconds)), + ), + ), + ) diff --git a/src/synthorg/api/auth/models.py b/src/synthorg/api/auth/models.py index 6d8f8cf4ff..88bb57a12c 100644 --- a/src/synthorg/api/auth/models.py +++ b/src/synthorg/api/auth/models.py @@ -13,6 +13,7 @@ class AuthMethod(StrEnum): JWT = "jwt" API_KEY = "api_key" + WS_TICKET = "ws_ticket" class User(BaseModel): diff --git a/src/synthorg/api/auth/ticket_store.py b/src/synthorg/api/auth/ticket_store.py new file mode 100644 index 0000000000..f8b97955fe --- /dev/null +++ b/src/synthorg/api/auth/ticket_store.py @@ -0,0 +1,172 @@ +"""In-memory store for short-lived, single-use WebSocket tickets. + +Tickets are ephemeral — they do not survive a server restart, which +forces re-authentication (correct security behaviour). The store +uses ``time.monotonic()`` for expiry so it is immune to wall-clock +adjustments. + +.. note:: + + The store is per-process — if the ASGI server runs multiple + worker processes, a ticket issued by one worker cannot be + consumed by another. ``ServerConfig.workers`` must be ``1`` + (the default) for ticket auth to work correctly. +""" + +import math +import secrets +import time + +from pydantic import BaseModel, ConfigDict + +from synthorg.api.auth.models import AuthenticatedUser # noqa: TC001 +from synthorg.observability import get_logger +from synthorg.observability.events.api import ( + API_WS_TICKET_CLEANUP, + API_WS_TICKET_CONSUMED, + API_WS_TICKET_EXPIRED, + API_WS_TICKET_INVALID, + API_WS_TICKET_ISSUED, +) + +logger = get_logger(__name__) + + +class TicketLimitExceededError(Exception): + """Raised when a user exceeds the per-user pending ticket cap.""" + + +# 32 bytes → 256 bits of entropy, encoded as 43 URL-safe base64 chars. +_TOKEN_BYTES: int = 32 +_MAX_PENDING_PER_USER: int = 5 + + +class _TicketEntry(BaseModel): + """Internal record for a pending ticket. + + Attributes: + user: Authenticated identity captured at ticket creation. + expires_at: ``time.monotonic()`` deadline. + """ + + model_config = ConfigDict(frozen=True) + + user: AuthenticatedUser + expires_at: float + + +class WsTicketStore: + """In-memory store for one-time WebSocket auth tickets. + + Each ticket is a cryptographically random URL-safe token + (43 characters, 256-bit entropy). Tickets expire after + *ttl_seconds* or on first use, whichever comes first. + + Args: + ttl_seconds: Ticket lifetime in seconds (default 30). + + Raises: + ValueError: If *ttl_seconds* is not positive. + """ + + def __init__(self, ttl_seconds: float = 30.0) -> None: + if not math.isfinite(ttl_seconds) or ttl_seconds <= 0: + msg = f"ttl_seconds must be a finite positive number, got {ttl_seconds}" + raise ValueError(msg) + self._ttl = ttl_seconds + self._tickets: dict[str, _TicketEntry] = {} + + @property + def ttl_seconds(self) -> float: + """Configured ticket lifetime.""" + return self._ttl + + def create(self, user: AuthenticatedUser) -> str: + """Issue a new single-use ticket for *user*. + + Args: + user: Authenticated identity to bind to the ticket. + + Returns: + URL-safe random token string. + """ + now = time.monotonic() + user_pending = sum( + 1 + for e in self._tickets.values() + if e.user.user_id == user.user_id and now <= e.expires_at + ) + if user_pending >= _MAX_PENDING_PER_USER: + msg = f"Ticket limit exceeded for user {user.user_id}" + raise TicketLimitExceededError(msg) + + ticket = secrets.token_urlsafe(_TOKEN_BYTES) + entry = _TicketEntry( + user=user, + expires_at=time.monotonic() + self._ttl, + ) + self._tickets[ticket] = entry + logger.info( + API_WS_TICKET_ISSUED, + user_id=user.user_id, + username=user.username, + ttl_seconds=self._ttl, + ) + return ticket + + def validate_and_consume(self, ticket: str) -> AuthenticatedUser | None: + """Validate and consume a ticket (single-use). + + Atomically removes the ticket via ``dict.pop`` before + checking expiry. In the single-threaded asyncio event loop, + ``dict.pop`` cannot be interleaved with another coroutine, + so concurrent calls on the same ticket are safely serialised. + + Args: + ticket: Raw ticket string from the client. + + Returns: + The bound ``AuthenticatedUser``, or ``None``. + """ + entry = self._tickets.pop(ticket, None) + if entry is None: + logger.warning(API_WS_TICKET_INVALID, reason="not_found") + return None + + now = time.monotonic() + if now > entry.expires_at: + logger.warning( + API_WS_TICKET_EXPIRED, + user_id=entry.user.user_id, + overdue_seconds=round(now - entry.expires_at, 2), + ) + return None + + logger.info( + API_WS_TICKET_CONSUMED, + user_id=entry.user.user_id, + username=entry.user.username, + ) + return entry.user + + def cleanup_expired(self) -> int: + """Remove expired tickets. + + Called periodically by a background task to prevent + unbounded memory growth from tickets that are requested + but never consumed. + + Returns: + Number of entries removed. + """ + now = time.monotonic() + expired = [k for k, v in self._tickets.items() if now > v.expires_at] + for k in expired: + self._tickets.pop(k, None) + if expired: + logger.info( + API_WS_TICKET_CLEANUP, + removed=len(expired), + remaining=len(self._tickets), + ) + return len(expired) diff --git a/src/synthorg/api/controllers/ws.py b/src/synthorg/api/controllers/ws.py index 650bdc896e..4811cb75e5 100644 --- a/src/synthorg/api/controllers/ws.py +++ b/src/synthorg/api/controllers/ws.py @@ -1,8 +1,13 @@ """WebSocket handler for real-time event feeds. -Clients connect to ``/api/v1/ws`` and send JSON messages to -subscribe/unsubscribe from named channels with optional payload -filters. The server pushes ``WsEvent`` JSON on subscribed channels. +Clients connect to ``/api/v1/ws?ticket=`` and send JSON +messages to subscribe/unsubscribe from named channels with optional +payload filters. The server pushes ``WsEvent`` JSON on subscribed +channels. + +Authentication uses a one-time ticket obtained from +``POST /api/v1/auth/ws-ticket``. The ticket is validated and +consumed (single-use) before the connection is accepted. """ import json @@ -13,8 +18,9 @@ from litestar.exceptions import WebSocketDisconnect from litestar.handlers import websocket +from synthorg.api.auth.models import AuthenticatedUser # noqa: TC001 from synthorg.api.channels import ALL_CHANNELS -from synthorg.api.guards import require_read_access +from synthorg.api.guards import _READ_ROLES, HumanRole from synthorg.observability import get_logger from synthorg.observability.events.api import ( API_WS_CONNECTED, @@ -22,6 +28,7 @@ API_WS_INVALID_MESSAGE, API_WS_SEND_FAILED, API_WS_SUBSCRIBE, + API_WS_TICKET_INVALID, API_WS_TRANSPORT_ERROR, API_WS_UNKNOWN_ACTION, API_WS_UNSUBSCRIBE, @@ -34,73 +41,173 @@ _MAX_FILTER_VALUE_LEN: int = 256 _MAX_WS_MESSAGE_BYTES: int = 4096 +# Application-layer WS close codes (RFC 6455 §7.4.2: 4000-4999). +_WS_CLOSE_AUTH_FAILED: int = 4001 +_WS_CLOSE_FORBIDDEN: int = 4003 -@websocket("/ws", guards=[require_read_access]) + +async def _validate_ticket( + socket: WebSocket[Any, Any, Any], +) -> AuthenticatedUser | None: + """Validate the one-time ticket and return the user. + + Returns ``None`` and closes the socket if the ticket is + missing, invalid, or expired. + """ + ticket = socket.query_params.get("ticket") + if not ticket: + logger.warning(API_WS_TICKET_INVALID, reason="missing_ticket") + await socket.close(code=_WS_CLOSE_AUTH_FAILED, reason="Missing ticket") + return None + + app_state = socket.app.state["app_state"] + user: AuthenticatedUser | None = app_state.ticket_store.validate_and_consume( + ticket, + ) + if user is None: + await socket.close( + code=_WS_CLOSE_AUTH_FAILED, + reason="Invalid or expired ticket", + ) + return None + + return user + + +async def _check_ws_role( + socket: WebSocket[Any, Any, Any], + user: AuthenticatedUser, +) -> bool: + """Verify the user has a role permitted for WebSocket access. + + Returns ``True`` if the role is valid. On failure, closes the + socket with a forbidden code and returns ``False``. + """ + # Defense-in-depth: user.role is already validated as HumanRole by + # Pydantic, and _READ_ROLES == frozenset(HumanRole). These checks + # guard against future changes to the role model or read-role set. + try: + role = HumanRole(user.role) + except ValueError: + logger.warning( + API_WS_TICKET_INVALID, + reason="invalid_role", + role=str(user.role), + ) + await socket.close(code=_WS_CLOSE_FORBIDDEN, reason="Invalid role") + return False + + if role not in _READ_ROLES: + logger.warning( + API_WS_TICKET_INVALID, + reason="insufficient_role", + role=role.value, + ) + await socket.close( + code=_WS_CLOSE_FORBIDDEN, + reason="Insufficient permissions", + ) + return False + + return True + + +async def _on_event( + event_data: bytes, + subscribed: set[str], + filters: dict[str, dict[str, str]], + socket: WebSocket[Any, Any, Any], +) -> None: + """Filter and forward a single channel event to the client.""" + try: + event = json.loads(event_data) + except json.JSONDecodeError: + logger.warning( + API_WS_INVALID_MESSAGE, + data_preview=str(event_data)[:100], + source="channels_backend", + ) + return + except TypeError: + logger.error( + API_WS_INVALID_MESSAGE, + data_type=type(event_data).__name__, + reason="unexpected_type", + source="channels_backend", + exc_info=True, + ) + return + + if not isinstance(event, dict): + logger.warning( + API_WS_INVALID_MESSAGE, + data_preview=str(event_data)[:100], + reason="not_a_dict", + ) + return + + channel = event.get("channel", "") + if channel not in subscribed: + return + + channel_filters = filters.get(channel) + if channel_filters: + payload = event.get("payload", {}) + if not isinstance(payload, dict): + return + if not all(payload.get(k) == v for k, v in channel_filters.items()): + return + + try: + await socket.send_text(event_data.decode("utf-8")) + except WebSocketDisconnect: + logger.debug(API_WS_SEND_FAILED, reason="client_disconnected") + except Exception: + logger.warning(API_WS_SEND_FAILED, exc_info=True) + await socket.close(code=1011, reason="Internal error") + + +@websocket("/ws") async def ws_handler( socket: WebSocket[Any, Any, Any], channels_plugin: ChannelsPlugin, ) -> None: """Handle WebSocket connections with channel subscriptions. - Clients subscribe to named channels with optional payload - filters. The server pushes ``WsEvent`` JSON for matching - events only. + Authentication is performed via a one-time ticket passed as + ``?ticket=`` in the query string. The ticket is + validated and consumed before the connection is accepted. Protocol (JSON from client): ``{"action": "subscribe", "channels": ["tasks"], "filters": {"agent_id": "...", "project": "..."}}`` ``{"action": "unsubscribe", "channels": ["tasks"]}`` """ + user = await _validate_ticket(socket) + if user is None: + return + + if not await _check_ws_role(socket, user): + return + + socket.scope["user"] = user await socket.accept() - logger.info(API_WS_CONNECTED, client=str(socket.client)) + logger.info( + API_WS_CONNECTED, + client=str(socket.client), + user_id=user.user_id, + ) subscribed: set[str] = set() filters: dict[str, dict[str, str]] = {} subscriber = await channels_plugin.subscribe(list(ALL_CHANNELS)) - async def _on_event(event_data: bytes) -> None: - """Filter and forward events to the client.""" - try: - event = json.loads(event_data) - except json.JSONDecodeError, TypeError: - logger.warning( - API_WS_INVALID_MESSAGE, - data_preview=str(event_data)[:100], - source="channels_backend", - ) - return - - if not isinstance(event, dict): - logger.warning( - API_WS_INVALID_MESSAGE, - data_preview=str(event_data)[:100], - reason="not_a_dict", - ) - return - - channel = event.get("channel", "") - if channel not in subscribed: - return - - channel_filters = filters.get(channel) - if channel_filters: - payload = event.get("payload", {}) - if not all(payload.get(k) == v for k, v in channel_filters.items()): - return - - try: - await socket.send_text(event_data.decode("utf-8")) - except WebSocketDisconnect: - logger.debug(API_WS_SEND_FAILED, reason="client_disconnected") - except Exception: - logger.warning( - API_WS_SEND_FAILED, - exc_info=True, - ) + async def _event_callback(event_data: bytes) -> None: + await _on_event(event_data, subscribed, filters, socket) try: - async with subscriber.run_in_background(_on_event): + async with subscriber.run_in_background(_event_callback): await _receive_loop(socket, subscribed, filters) finally: await channels_plugin.unsubscribe(subscriber) @@ -121,81 +228,140 @@ async def _receive_loop( except WebSocketDisconnect: logger.debug(API_WS_DISCONNECTED, reason="client_disconnect") except Exception: + user = socket.scope.get("user") logger.error( API_WS_TRANSPORT_ERROR, + user_id=getattr(user, "user_id", "unknown"), + client=str(socket.client), exc_info=True, ) raise -def _handle_message( # noqa: C901, PLR0911 +def _parse_ws_message( data: str, - subscribed: set[str], - filters: dict[str, dict[str, str]], -) -> str: - """Parse and handle a single client message. - - Args: - data: Raw JSON string from the client. - subscribed: Mutable set of subscribed channel names. - filters: Mutable per-channel payload filters. - - Returns: - JSON acknowledgement or error string. - """ +) -> dict[str, Any] | str: + """Parse raw JSON from the client, returning a dict or an error string.""" if len(data.encode()) > _MAX_WS_MESSAGE_BYTES: return json.dumps({"error": "Message too large"}) try: msg = json.loads(data) - except json.JSONDecodeError, TypeError: - logger.warning( + except json.JSONDecodeError: + logger.warning(API_WS_INVALID_MESSAGE, data_preview=str(data)[:100]) + return json.dumps({"error": "Invalid JSON"}) + except TypeError: + logger.error( API_WS_INVALID_MESSAGE, - data_preview=str(data)[:100], + data_type=type(data).__name__, + reason="unexpected_type", + exc_info=True, ) return json.dumps({"error": "Invalid JSON"}) if not isinstance(msg, dict): return json.dumps({"error": "Expected JSON object"}) - action = msg.get("action") + return msg + + +def _validate_ws_fields( + msg: dict[str, Any], +) -> tuple[str, list[str], dict[str, Any] | None] | str: + """Extract and validate action, channels, and filters from a parsed message. + + Returns ``(action, channels, client_filters)`` on success, or a + JSON error string on validation failure. + """ + action = str(msg.get("action", "")) channels = msg.get("channels", []) - client_filters = msg.get("filters", {}) + # None = key absent (leave existing filters), {} = explicitly clear + raw_filters = msg.get("filters") + client_filters: dict[str, Any] | None = None + if raw_filters is not None: + if not isinstance(raw_filters, dict): + return json.dumps({"error": "filters must be an object"}) + client_filters = raw_filters if not isinstance(channels, list) or not all(isinstance(c, str) for c in channels): return json.dumps({"error": "channels must be a list of strings"}) - if not isinstance(client_filters, dict): - return json.dumps({"error": "filters must be an object"}) + + return (action, channels, client_filters) + + +def _handle_message( + data: str, + subscribed: set[str], + filters: dict[str, dict[str, str]], +) -> str: + """Parse, validate, and dispatch a single client message.""" + parsed = _parse_ws_message(data) + if isinstance(parsed, str): + return parsed + + fields = _validate_ws_fields(parsed) + if isinstance(fields, str): + return fields + + action, channels, client_filters = fields if action == "subscribe": - # Validate filter bounds to prevent memory abuse. - if len(client_filters) > _MAX_FILTER_KEYS or any( - len(str(v)) > _MAX_FILTER_VALUE_LEN for v in client_filters.values() - ): - return json.dumps({"error": "Filter bounds exceeded"}) - - valid = [c for c in channels if c in _ALL_CHANNELS_SET] - subscribed.update(valid) - for c in valid: - if client_filters: - filters[c] = dict(client_filters) - logger.debug( - API_WS_SUBSCRIBE, - channels=valid, - active=sorted(subscribed), - ) - return json.dumps({"action": "subscribed", "channels": sorted(subscribed)}) + return _handle_subscribe(channels, client_filters, subscribed, filters) if action == "unsubscribe": - subscribed -= set(channels) - for c in channels: - filters.pop(c, None) - logger.debug( - API_WS_UNSUBSCRIBE, - channels=channels, - active=sorted(subscribed), - ) - return json.dumps({"action": "unsubscribed", "channels": sorted(subscribed)}) + return _handle_unsubscribe(channels, subscribed, filters) logger.warning(API_WS_UNKNOWN_ACTION, action=str(action)[:64]) return json.dumps({"error": "Unknown action"}) + + +def _handle_subscribe( + channels: list[str], + client_filters: dict[str, Any] | None, + subscribed: set[str], + filters: dict[str, dict[str, str]], +) -> str: + """Process a subscribe action. + + Filter semantics: + ``None`` — filters key absent, leave existing filters unchanged. + ``{}`` — explicitly clear filters for the subscribed channels. + ``{...}`` — set new filters for the subscribed channels. + """ + if client_filters is not None and ( + len(client_filters) > _MAX_FILTER_KEYS + or any(len(str(v)) > _MAX_FILTER_VALUE_LEN for v in client_filters.values()) + ): + return json.dumps({"error": "Filter bounds exceeded"}) + + valid = [c for c in channels if c in _ALL_CHANNELS_SET] + subscribed.update(valid) + if client_filters is not None: + for c in valid: + if client_filters: + filters[c] = dict(client_filters) + else: + filters.pop(c, None) + logger.debug( + API_WS_SUBSCRIBE, + channels=valid, + active=sorted(subscribed), + ) + return json.dumps({"action": "subscribed", "channels": sorted(subscribed)}) + + +def _handle_unsubscribe( + channels: list[str], + subscribed: set[str], + filters: dict[str, dict[str, str]], +) -> str: + """Process an unsubscribe action.""" + subscribed -= set(channels) + for c in channels: + filters.pop(c, None) + logger.debug( + API_WS_UNSUBSCRIBE, + channels=channels, + active=sorted(subscribed), + ) + return json.dumps({"action": "unsubscribed", "channels": sorted(subscribed)}) diff --git a/src/synthorg/api/state.py b/src/synthorg/api/state.py index 60b30936db..51b6587308 100644 --- a/src/synthorg/api/state.py +++ b/src/synthorg/api/state.py @@ -7,6 +7,7 @@ from synthorg.api.approval_store import ApprovalStore # noqa: TC001 from synthorg.api.auth.service import AuthService # noqa: TC001 +from synthorg.api.auth.ticket_store import WsTicketStore from synthorg.api.errors import ServiceUnavailableError from synthorg.budget.tracker import CostTracker # noqa: TC001 from synthorg.communication.bus_protocol import MessageBus # noqa: TC001 @@ -43,6 +44,8 @@ class AppState: config: Root company configuration. approval_store: In-memory approval queue store. startup_time: ``time.monotonic()`` snapshot at app creation. + ticket_store: In-memory one-time WebSocket ticket store + (always available, no ``_require_service`` guard needed). """ __slots__ = ( @@ -58,6 +61,7 @@ class AppState: "_persistence", "_settings_service", "_task_engine", + "_ticket_store", "approval_store", "config", "startup_time", @@ -96,6 +100,7 @@ def __init__( # noqa: PLR0913 self._meeting_orchestrator = meeting_orchestrator self._meeting_scheduler = meeting_scheduler self._settings_service = settings_service + self._ticket_store = WsTicketStore() self.startup_time = startup_time def _require_service[T](self, service: T | None, name: str) -> T: @@ -236,6 +241,11 @@ def has_auth_service(self) -> bool: """Check whether the auth service is already configured.""" return self._auth_service is not None + @property + def ticket_store(self) -> WsTicketStore: + """Return the WebSocket ticket store (always available).""" + return self._ticket_store + def set_auth_service(self, service: AuthService) -> None: """Set the auth service (deferred initialisation). diff --git a/src/synthorg/observability/events/api.py b/src/synthorg/observability/events/api.py index 1ddffc48ff..8dc23330d9 100644 --- a/src/synthorg/observability/events/api.py +++ b/src/synthorg/observability/events/api.py @@ -50,3 +50,8 @@ API_CONTENT_NEGOTIATED: Final[str] = "api.content.negotiated" API_CORRELATION_FALLBACK: Final[str] = "api.correlation.fallback" API_ACCEPT_PARSE_FAILED: Final[str] = "api.accept.parse_failed" +API_WS_TICKET_ISSUED: Final[str] = "api.ws.ticket_issued" +API_WS_TICKET_CONSUMED: Final[str] = "api.ws.ticket_consumed" +API_WS_TICKET_EXPIRED: Final[str] = "api.ws.ticket_expired" +API_WS_TICKET_INVALID: Final[str] = "api.ws.ticket_invalid" +API_WS_TICKET_CLEANUP: Final[str] = "api.ws.ticket_cleanup" diff --git a/tests/unit/api/auth/test_controller.py b/tests/unit/api/auth/test_controller.py index f2417bcb7e..b8c7d904d5 100644 --- a/tests/unit/api/auth/test_controller.py +++ b/tests/unit/api/auth/test_controller.py @@ -325,3 +325,56 @@ def test_rejects_unknown_user_type(self) -> None: with pytest.raises(PermissionDeniedException): require_password_changed(connection, None) + + +@pytest.mark.timeout(30) +@pytest.mark.unit +class TestWsTicket: + def test_ws_ticket_returns_ticket_and_expires_in( + self, + test_client: TestClient[Any], + ) -> None: + response = test_client.post("/api/v1/auth/ws-ticket") + assert response.status_code == 200 + data = response.json()["data"] + assert "ticket" in data + assert isinstance(data["ticket"], str) + assert len(data["ticket"]) > 0 + assert data["expires_in"] == 30 + + def test_ws_ticket_requires_auth( + self, + bare_client: TestClient[Any], + ) -> None: + response = bare_client.post("/api/v1/auth/ws-ticket") + assert response.status_code == 401 + + def test_ws_ticket_with_observer_role( + self, + test_client: TestClient[Any], + ) -> None: + """All roles should be able to get a WS ticket.""" + response = test_client.post( + "/api/v1/auth/ws-ticket", + headers=make_auth_headers("observer"), + ) + assert response.status_code == 200 + data = response.json()["data"] + assert "ticket" in data + + def test_ws_ticket_is_consumable( + self, + test_client: TestClient[Any], + ) -> None: + """The returned ticket can be consumed by the ticket store.""" + response = test_client.post("/api/v1/auth/ws-ticket") + data = response.json()["data"] + ticket = data["ticket"] + + app_state = test_client.app.state["app_state"] + user = app_state.ticket_store.validate_and_consume(ticket) + assert user is not None + assert user.auth_method.value == "ws_ticket" + + # Single-use: second consume fails + assert app_state.ticket_store.validate_and_consume(ticket) is None diff --git a/tests/unit/api/auth/test_ticket_store.py b/tests/unit/api/auth/test_ticket_store.py new file mode 100644 index 0000000000..1c6ea6d97c --- /dev/null +++ b/tests/unit/api/auth/test_ticket_store.py @@ -0,0 +1,310 @@ +"""Tests for WsTicketStore.""" + +import math +import re +from unittest.mock import patch + +import pytest + +from synthorg.api.auth.models import AuthenticatedUser, AuthMethod +from synthorg.api.auth.ticket_store import TicketLimitExceededError, WsTicketStore +from synthorg.api.guards import HumanRole + + +def _make_user( + *, + user_id: str = "test-user-001", + username: str = "testadmin", + role: HumanRole = HumanRole.CEO, +) -> AuthenticatedUser: + return AuthenticatedUser( + user_id=user_id, + username=username, + role=role, + auth_method=AuthMethod.WS_TICKET, + ) + + +@pytest.mark.timeout(30) +@pytest.mark.unit +class TestWsTicketStoreCreate: + """Tests for ticket creation.""" + + def test_create_returns_url_safe_string(self) -> None: + store = WsTicketStore() + user = _make_user() + ticket = store.create(user) + + assert isinstance(ticket, str) + assert len(ticket) > 0 + # URL-safe base64 characters only + assert re.fullmatch(r"[A-Za-z0-9_-]+", ticket) + + def test_create_returns_unique_tickets(self) -> None: + store = WsTicketStore() + # Use different user IDs to avoid per-user ticket cap + tickets = {store.create(_make_user(user_id=f"user-{i}")) for i in range(100)} + assert len(tickets) == 100 + + def test_ttl_seconds_property(self) -> None: + store = WsTicketStore(ttl_seconds=60.0) + assert store.ttl_seconds == 60.0 + + @pytest.mark.parametrize( + ("ttl", "match"), + [ + (0.0, "positive"), + (-5.0, "positive"), + (math.nan, "finite positive"), + (math.inf, "finite positive"), + (-math.inf, "finite positive"), + ], + ids=["zero", "negative", "nan", "inf", "neg_inf"], + ) + def test_invalid_ttl_rejected(self, ttl: float, match: str) -> None: + with pytest.raises(ValueError, match=match): + WsTicketStore(ttl_seconds=ttl) + + def test_per_user_ticket_cap(self) -> None: + """Creating more than _MAX_PENDING_PER_USER tickets raises.""" + store = WsTicketStore() + user = _make_user() + for _ in range(5): + store.create(user) + with pytest.raises(TicketLimitExceededError): + store.create(user) + + def test_per_user_ticket_cap_different_users(self) -> None: + """Different users have independent ticket caps.""" + store = WsTicketStore() + user_a = _make_user(user_id="user-a") + user_b = _make_user(user_id="user-b") + for _ in range(5): + store.create(user_a) + # user_b should still be able to create tickets + ticket = store.create(user_b) + assert isinstance(ticket, str) + + +@pytest.mark.timeout(30) +@pytest.mark.unit +class TestWsTicketStoreValidateAndConsume: + """Tests for ticket validation and consumption.""" + + def test_validate_and_consume_returns_user(self) -> None: + store = WsTicketStore() + user = _make_user() + ticket = store.create(user) + + result = store.validate_and_consume(ticket) + + assert result is not None + assert result.user_id == user.user_id + assert result.username == user.username + assert result.role == user.role + assert result.auth_method == AuthMethod.WS_TICKET + + def test_validate_and_consume_single_use(self) -> None: + store = WsTicketStore() + user = _make_user() + ticket = store.create(user) + + first = store.validate_and_consume(ticket) + second = store.validate_and_consume(ticket) + + assert first is not None + assert second is None + + def test_validate_and_consume_single_use_concurrent(self) -> None: + """Exactly one concurrent consumer wins the ticket.""" + import threading + from concurrent.futures import ThreadPoolExecutor + + store = WsTicketStore() + user = _make_user() + ticket = store.create(user) + + barrier = threading.Barrier(10) + + def consume() -> AuthenticatedUser | None: + barrier.wait() + return store.validate_and_consume(ticket) + + with ThreadPoolExecutor(max_workers=10) as pool: + results = list(pool.map(lambda _: consume(), range(10))) + + winners = [r for r in results if r is not None] + assert len(winners) == 1 + assert winners[0].user_id == user.user_id + + def test_validate_and_consume_expired(self) -> None: + store = WsTicketStore(ttl_seconds=10.0) + user = _make_user() + + base_time = 1000.0 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + ticket = store.create(user) + + # Advance past expiry + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 11.0, + ): + result = store.validate_and_consume(ticket) + + assert result is None + + def test_validate_and_consume_just_before_expiry(self) -> None: + store = WsTicketStore(ttl_seconds=10.0) + user = _make_user() + + base_time = 1000.0 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + ticket = store.create(user) + + # Just before expiry — should still work + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 9.9, + ): + result = store.validate_and_consume(ticket) + + assert result is not None + + def test_validate_and_consume_unknown_ticket(self) -> None: + store = WsTicketStore() + result = store.validate_and_consume("nonexistent-ticket") + assert result is None + + def test_validate_and_consume_empty_string(self) -> None: + store = WsTicketStore() + result = store.validate_and_consume("") + assert result is None + + def test_custom_ttl(self) -> None: + store = WsTicketStore(ttl_seconds=5.0) + user = _make_user() + + base_time = 1000.0 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + ticket = store.create(user) + + # Within custom TTL + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 4.0, + ): + result = store.validate_and_consume(ticket) + assert result is not None + + def test_custom_ttl_expired(self) -> None: + store = WsTicketStore(ttl_seconds=5.0) + user = _make_user() + + base_time = 1000.0 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + ticket = store.create(user) + + # Past custom TTL + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 6.0, + ): + result = store.validate_and_consume(ticket) + assert result is None + + +@pytest.mark.timeout(30) +@pytest.mark.unit +class TestWsTicketStoreCleanup: + """Tests for expired ticket cleanup.""" + + def test_cleanup_expired_removes_old_entries(self) -> None: + store = WsTicketStore(ttl_seconds=10.0) + user = _make_user() + + base_time = 1000.0 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + store.create(user) + store.create(user) + + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 11.0, + ): + removed = store.cleanup_expired() + + assert removed == 2 + + def test_cleanup_preserves_valid_entries(self) -> None: + store = WsTicketStore(ttl_seconds=10.0) + user = _make_user() + + base_time = 1000.0 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + ticket = store.create(user) + + # Still within TTL + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 5.0, + ): + removed = store.cleanup_expired() + + assert removed == 0 + # Ticket should still be valid + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 5.0, + ): + result = store.validate_and_consume(ticket) + assert result is not None + + def test_cleanup_mixed_expired_and_valid(self) -> None: + store = WsTicketStore(ttl_seconds=10.0) + user = _make_user() + + base_time = 1000.0 + # Create two tickets at different times + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", return_value=base_time + ): + store.create(user) # expires at 1010 + + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 8.0, + ): + valid_ticket = store.create(user) # expires at 1018 + + # At t=1012: first expired, second still valid + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 12.0, + ): + removed = store.cleanup_expired() + + assert removed == 1 + with patch( + "synthorg.api.auth.ticket_store.time.monotonic", + return_value=base_time + 12.0, + ): + result = store.validate_and_consume(valid_ticket) + assert result is not None + + def test_cleanup_empty_store(self) -> None: + store = WsTicketStore() + removed = store.cleanup_expired() + assert removed == 0 diff --git a/tests/unit/api/controllers/test_ws.py b/tests/unit/api/controllers/test_ws.py index 3b80ade842..8d8b956b22 100644 --- a/tests/unit/api/controllers/test_ws.py +++ b/tests/unit/api/controllers/test_ws.py @@ -1,10 +1,18 @@ -"""Tests for WebSocket handler message parsing.""" +"""Tests for WebSocket handler message parsing and ticket auth.""" import json +from typing import Any import pytest +from litestar.testing import TestClient -from synthorg.api.controllers.ws import _handle_message +from synthorg.api.auth.models import AuthMethod +from synthorg.api.controllers.ws import ( + _WS_CLOSE_AUTH_FAILED, + _WS_CLOSE_FORBIDDEN, + _handle_message, +) +from synthorg.api.guards import _READ_ROLES, HumanRole @pytest.mark.unit @@ -160,7 +168,7 @@ def test_message_size_limit_boundary(self) -> None: small_msg = json.dumps({"action": "subscribe", "channels": ["tasks"]}) result = _handle_message(small_msg, subscribed, filters) data = json.loads(result) - assert "error" not in data or data.get("action") == "subscribed" + assert data["action"] == "subscribed" # Message whose encoded bytes exceed 4096 should fail big_msg = json.dumps({"action": "subscribe", "data": "x" * 4096}) @@ -168,33 +176,133 @@ def test_message_size_limit_boundary(self) -> None: data = json.loads(result) assert data["error"] == "Message too large" - def test_non_dict_json_returns_error(self) -> None: + @pytest.mark.parametrize( + "value", + [[1, 2, 3], "hello", 42], + ids=["array", "string", "number"], + ) + def test_non_dict_json_returns_error(self, value: object) -> None: subscribed: set[str] = set() filters: dict[str, dict[str, str]] = {} - - # JSON array - result = _handle_message( - json.dumps([1, 2, 3]), - subscribed, - filters, - ) + result = _handle_message(json.dumps(value), subscribed, filters) data = json.loads(result) - assert "error" in data + assert data["error"] == "Expected JSON object" - # JSON string - result = _handle_message( - json.dumps("hello"), - subscribed, - filters, - ) - data = json.loads(result) - assert "error" in data - # JSON number - result = _handle_message( - json.dumps(42), - subscribed, - filters, +@pytest.mark.timeout(30) +@pytest.mark.unit +class TestWsTicketAuth: + """Tests for ticket-based WebSocket authentication logic. + + These tests validate the auth validation logic used by the WS + handler without opening actual WebSocket connections (which + require the channels plugin background task and hang in the + sync test client). + """ + + def test_ws_ticket_endpoint_returns_ticket( + self, + test_client: TestClient[Any], + ) -> None: + """POST /auth/ws-ticket returns a consumable ticket.""" + response = test_client.post("/api/v1/auth/ws-ticket") + assert response.status_code == 200 + data = response.json()["data"] + assert "ticket" in data + assert data["expires_in"] == 30 + + def test_ws_ticket_carries_ws_ticket_auth_method( + self, + test_client: TestClient[Any], + ) -> None: + """The ticket user has auth_method=WS_TICKET.""" + response = test_client.post("/api/v1/auth/ws-ticket") + ticket = response.json()["data"]["ticket"] + + app_state = test_client.app.state["app_state"] + user = app_state.ticket_store.validate_and_consume(ticket) + assert user is not None + assert user.auth_method == AuthMethod.WS_TICKET + + def test_ws_ticket_single_use_via_store( + self, + test_client: TestClient[Any], + ) -> None: + """Ticket is consumed on first validate_and_consume.""" + response = test_client.post("/api/v1/auth/ws-ticket") + ticket = response.json()["data"]["ticket"] + + app_state = test_client.app.state["app_state"] + first = app_state.ticket_store.validate_and_consume(ticket) + second = app_state.ticket_store.validate_and_consume(ticket) + assert first is not None + assert second is None + + def test_ws_ticket_user_has_correct_identity( + self, + test_client: TestClient[Any], + ) -> None: + """The ticket preserves the original user's identity.""" + response = test_client.post("/api/v1/auth/ws-ticket") + ticket = response.json()["data"]["ticket"] + + app_state = test_client.app.state["app_state"] + user = app_state.ticket_store.validate_and_consume(ticket) + assert user is not None + assert user.role == HumanRole.CEO + assert user.username == "test-ceo" + + def test_ws_endpoint_excluded_from_auth_middleware(self) -> None: + """The /ws path must be in the auto-derived auth exclude list.""" + from synthorg.api.app import _build_middleware + from synthorg.api.config import ApiConfig + + api_config = ApiConfig() + middleware = _build_middleware(api_config) + auth_cls = middleware[0] + + from litestar import Litestar + + dummy_app = Litestar(route_handlers=[]) + instance = auth_cls(dummy_app) # type: ignore[operator,call-arg] + # Litestar compiles exclude paths into a single regex pattern. + assert instance.exclude is not None # type: ignore[union-attr] + assert instance.exclude.match("/api/v1/ws"), ( # type: ignore[union-attr] + f"/api/v1/ws not matched by exclude: {instance.exclude.pattern}" # type: ignore[union-attr] ) - data = json.loads(result) - assert "error" in data + + def test_read_roles_includes_all_human_roles(self) -> None: + """The WS handler's _READ_ROLES should include all HumanRole values.""" + for role in HumanRole: + assert role in _READ_ROLES + + def test_ws_close_codes_in_application_range(self) -> None: + """WS close codes should be in the RFC 6455 application range.""" + assert 4000 <= _WS_CLOSE_AUTH_FAILED <= 4999 + assert 4000 <= _WS_CLOSE_FORBIDDEN <= 4999 + + def test_ws_rejects_missing_ticket( + self, + test_client: TestClient[Any], + ) -> None: + """WS connection without ?ticket= is rejected (close before accept).""" + from litestar.exceptions import WebSocketDisconnect + + with ( + pytest.raises(WebSocketDisconnect), + test_client.websocket_connect("/api/v1/ws"), + ): + pass + + def test_ws_rejects_invalid_ticket( + self, + test_client: TestClient[Any], + ) -> None: + """WS connection with a bogus ticket is rejected.""" + from litestar.exceptions import WebSocketDisconnect + + with ( + pytest.raises(WebSocketDisconnect), + test_client.websocket_connect("/api/v1/ws?ticket=bogus-ticket"), + ): + pass diff --git a/web/src/__tests__/composables/useWebSocketSubscription.test.ts b/web/src/__tests__/composables/useWebSocketSubscription.test.ts index 1ea1977a08..e6a3818c85 100644 --- a/web/src/__tests__/composables/useWebSocketSubscription.test.ts +++ b/web/src/__tests__/composables/useWebSocketSubscription.test.ts @@ -2,12 +2,20 @@ import { describe, it, expect, vi, beforeEach, type Mock } from 'vitest' import { setActivePinia, createPinia } from 'pinia' import { onMounted, onUnmounted } from 'vue' -// Mock Vue lifecycle hooks since we're not in a component context +// Mock Vue lifecycle hooks since we're not in a component context. +// onMounted callback is now async — we invoke it and store the promise +// so tests can await it when they need to verify post-connect behaviour. +let mountedPromise: Promise | undefined vi.mock('vue', async () => { const actual = await vi.importActual('vue') return { ...actual, - onMounted: vi.fn((cb: () => void) => cb()), + onMounted: vi.fn((cb: () => void | Promise) => { + const result = cb() + if (result instanceof Promise) { + mountedPromise = result // preserve rejections for test assertion + } + }), onUnmounted: vi.fn(), } }) @@ -26,12 +34,18 @@ describe('useWebSocketSubscription', () => { setActivePinia(createPinia()) wsStore = useWebSocketStore() authStore = useAuthStore() + mountedPromise = undefined vi.clearAllMocks() consoleSpy = vi.spyOn(console, 'error').mockImplementation(() => {}) // Re-establish lifecycle mocks after clearAllMocks: - // - onMounted: synchronously invokes callback so setup logic runs during test + // - onMounted: invokes callback (may be async); stores promise for awaiting // - onUnmounted: no-op recorder; getUnmountCallback() reads from mock.calls - ;(onMounted as Mock).mockImplementation((cb: () => void) => cb()) + ;(onMounted as Mock).mockImplementation((cb: () => void | Promise) => { + const result = cb() + if (result instanceof Promise) { + mountedPromise = result + } + }) ;(onUnmounted as Mock).mockImplementation(() => {}) }) @@ -55,16 +69,17 @@ describe('useWebSocketSubscription', () => { expect(result.setupError.value).toBeNull() }) - it('calls connect when auth token exists and not connected', () => { - const connectSpy = vi.spyOn(wsStore, 'connect') + it('calls connect when auth token exists and not connected', async () => { + const connectSpy = vi.spyOn(wsStore, 'connect').mockResolvedValue() authStore.$patch({ token: 'test-token' }) const handler: WsEventHandler = vi.fn() useWebSocketSubscription({ bindings: [{ channel: 'tasks', handler }], }) + await mountedPromise - expect(connectSpy).toHaveBeenCalledWith('test-token') + expect(connectSpy).toHaveBeenCalledWith() }) it('skips connect when already connected', () => { @@ -95,7 +110,8 @@ describe('useWebSocketSubscription', () => { expect(onSpy).not.toHaveBeenCalled() }) - it('subscribes to deduplicated channels from bindings', () => { + it('subscribes to deduplicated channels from bindings', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() const subscribeSpy = vi.spyOn(wsStore, 'subscribe') authStore.$patch({ token: 'test-token' }) const handler1: WsEventHandler = vi.fn() @@ -107,11 +123,13 @@ describe('useWebSocketSubscription', () => { { channel: 'budget', handler: handler2 }, ], }) + await mountedPromise expect(subscribeSpy).toHaveBeenCalledWith(['tasks', 'budget'], undefined) }) - it('forwards filters to subscribe', () => { + it('forwards filters to subscribe', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() const subscribeSpy = vi.spyOn(wsStore, 'subscribe') authStore.$patch({ token: 'test-token' }) const handler: WsEventHandler = vi.fn() @@ -121,11 +139,13 @@ describe('useWebSocketSubscription', () => { bindings: [{ channel: 'tasks', handler }], filters, }) + await mountedPromise expect(subscribeSpy).toHaveBeenCalledWith(['tasks'], filters) }) - it('calls onChannelEvent for each binding', () => { + it('calls onChannelEvent for each binding', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() const onSpy = vi.spyOn(wsStore, 'onChannelEvent') authStore.$patch({ token: 'test-token' }) const handler1: WsEventHandler = vi.fn() @@ -137,13 +157,15 @@ describe('useWebSocketSubscription', () => { { channel: 'budget', handler: handler2 }, ], }) + await mountedPromise expect(onSpy).toHaveBeenCalledTimes(2) expect(onSpy).toHaveBeenCalledWith('tasks', handler1) expect(onSpy).toHaveBeenCalledWith('budget', handler2) }) - it('deduplicates channels but wires both handlers for same channel', () => { + it('deduplicates channels but wires both handlers for same channel', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() const subscribeSpy = vi.spyOn(wsStore, 'subscribe') const onSpy = vi.spyOn(wsStore, 'onChannelEvent') authStore.$patch({ token: 'test-token' }) @@ -156,6 +178,7 @@ describe('useWebSocketSubscription', () => { { channel: 'tasks', handler: handler2 }, ], }) + await mountedPromise // Subscribe only lists channel once expect(subscribeSpy).toHaveBeenCalledWith(['tasks'], undefined) @@ -224,7 +247,8 @@ describe('useWebSocketSubscription', () => { expect(onSpy).not.toHaveBeenCalled() }) - it('sets setupError and logs when subscribe throws', () => { + it('sets setupError and logs when subscribe throws', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() vi.spyOn(wsStore, 'subscribe').mockImplementation(() => { throw new Error('subscribe failed') }) @@ -234,6 +258,7 @@ describe('useWebSocketSubscription', () => { const { setupError } = useWebSocketSubscription({ bindings: [{ channel: 'tasks', handler }], }) + await mountedPromise expect(setupError.value).toBe('WebSocket subscription failed.') expect(consoleSpy).toHaveBeenCalledWith( @@ -243,7 +268,8 @@ describe('useWebSocketSubscription', () => { ) }) - it('still wires handlers when subscribe throws', () => { + it('still wires handlers when subscribe throws', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() const onSpy = vi.spyOn(wsStore, 'onChannelEvent') vi.spyOn(wsStore, 'subscribe').mockImplementation(() => { throw new Error('subscribe failed') @@ -254,16 +280,19 @@ describe('useWebSocketSubscription', () => { useWebSocketSubscription({ bindings: [{ channel: 'tasks', handler }], }) + await mountedPromise expect(onSpy).toHaveBeenCalledWith('tasks', handler) }) - it('handles empty bindings array with token', () => { + it('handles empty bindings array with token', async () => { + vi.spyOn(wsStore, 'connect').mockResolvedValue() const subscribeSpy = vi.spyOn(wsStore, 'subscribe') const onSpy = vi.spyOn(wsStore, 'onChannelEvent') authStore.$patch({ token: 'test-token' }) const result = useWebSocketSubscription({ bindings: [] }) + await mountedPromise expect(subscribeSpy).toHaveBeenCalledWith([], undefined) expect(onSpy).not.toHaveBeenCalled() diff --git a/web/src/__tests__/stores/websocket.test.ts b/web/src/__tests__/stores/websocket.test.ts index 15cd907645..095d2bb957 100644 --- a/web/src/__tests__/stores/websocket.test.ts +++ b/web/src/__tests__/stores/websocket.test.ts @@ -3,6 +3,13 @@ import { setActivePinia, createPinia } from 'pinia' import { useWebSocketStore } from '@/stores/websocket' import type { WsEvent } from '@/api/types' +// Mock the auth API module — getWsTicket returns a one-time ticket +vi.mock('@/api/endpoints/auth', () => ({ + getWsTicket: vi.fn().mockResolvedValue({ ticket: 'test-ticket-abc', expires_in: 30 }), +})) + +import { getWsTicket } from '@/api/endpoints/auth' + // Track all created MockWebSocket instances let mockInstances: MockWebSocket[] = [] @@ -40,6 +47,7 @@ beforeEach(() => { mockInstances = [] // @ts-expect-error -- mock WebSocket for testing globalThis.WebSocket = MockWebSocket + vi.mocked(getWsTicket).mockResolvedValue({ ticket: 'test-ticket-abc', expires_in: 30 }) }) afterEach(() => { @@ -66,19 +74,32 @@ describe('useWebSocketStore', () => { it('connects and sets connected to true', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) expect(store.connected).toBe(true) }) + it('fetches a ticket before opening WebSocket', async () => { + const store = useWebSocketStore() + vi.mocked(getWsTicket).mockClear() + await store.connect() + await vi.advanceTimersByTimeAsync(0) + + expect(getWsTicket).toHaveBeenCalledTimes(1) + // URL should contain ticket, not JWT + const ws = mockInstances[0] + expect(ws.url).toContain('ticket=test-ticket-abc') + expect(ws.url).not.toContain('token=') + }) + it('does not create duplicate connections', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) expect(mockInstances).toHaveLength(1) - store.connect('test-token') // should be no-op + await store.connect() // should be no-op — already connected expect(mockInstances).toHaveLength(1) // no new WebSocket created }) @@ -89,14 +110,12 @@ describe('useWebSocketStore', () => { // No WebSocket exists, so no send should have been called expect(mockInstances).toHaveLength(0) - // Verify subscription is queued by connecting and checking send was called - store.connect('test-token') }) it('replays pending subscriptions on connect', async () => { const store = useWebSocketStore() store.subscribe(['tasks']) - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) expect(store.connected).toBe(true) @@ -118,7 +137,7 @@ describe('useWebSocketStore', () => { store.subscribe(['tasks', 'agents']) // Connect and verify only one subscribe message is sent (not three) - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const ws = mockInstances[0] @@ -130,7 +149,7 @@ describe('useWebSocketStore', () => { it('disconnect sets state correctly', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) expect(store.connected).toBe(true) @@ -141,7 +160,7 @@ describe('useWebSocketStore', () => { it('dispatches events to channel handlers via onmessage', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const handler = vi.fn() @@ -170,7 +189,7 @@ describe('useWebSocketStore', () => { it('wildcard handlers receive events from all channels', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const handler = vi.fn() @@ -201,7 +220,7 @@ describe('useWebSocketStore', () => { it('handles malformed JSON messages gracefully', async () => { const consoleSpy = vi.spyOn(console, 'error').mockImplementation(() => {}) const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const ws = mockInstances[0] @@ -216,7 +235,7 @@ describe('useWebSocketStore', () => { it('subscription ack updates subscribedChannels when array is valid', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const ws = mockInstances[0] @@ -237,7 +256,7 @@ describe('useWebSocketStore', () => { it('scheduleReconnect stops after max attempts', async () => { const consoleSpy = vi.spyOn(console, 'error').mockImplementation(() => {}) const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) expect(store.connected).toBe(true) @@ -282,7 +301,7 @@ describe('useWebSocketStore', () => { it('re-subscribes to active subscriptions on reconnect', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) // Subscribe while connected @@ -306,9 +325,28 @@ describe('useWebSocketStore', () => { ) }) + it('reconnect fetches a fresh ticket each time', async () => { + vi.mocked(getWsTicket).mockClear() + const store = useWebSocketStore() + await store.connect() + await vi.advanceTimersByTimeAsync(0) + expect(getWsTicket).toHaveBeenCalledTimes(1) + + // Simulate disconnect + const ws1 = mockInstances[0] + ws1.readyState = MockWebSocket.CLOSED + ws1.onclose?.() + + // Trigger reconnect + await vi.advanceTimersByTimeAsync(5_000) + + // Should have fetched a second ticket + expect(getWsTicket).toHaveBeenCalledTimes(2) + }) + it('unsubscribe removes channels from active subscriptions so reconnect does not re-subscribe', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) // Subscribe then unsubscribe @@ -334,7 +372,7 @@ describe('useWebSocketStore', () => { it('sanitizes error messages from server', async () => { const consoleSpy = vi.spyOn(console, 'error').mockImplementation(() => {}) const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const ws = mockInstances[0] @@ -352,7 +390,7 @@ describe('useWebSocketStore', () => { it('send failures queue subscriptions for replay', async () => { const store = useWebSocketStore() - store.connect('test-token') + await store.connect() await vi.advanceTimersByTimeAsync(0) const ws = mockInstances[0] @@ -365,4 +403,21 @@ describe('useWebSocketStore', () => { // Should not throw — caught internally and queued for replay expect(store.connected).toBe(true) }) + + it('connect fails gracefully when ticket exchange fails', async () => { + const consoleSpy = vi.spyOn(console, 'error').mockImplementation(() => {}) + vi.mocked(getWsTicket).mockRejectedValueOnce(new Error('401 Unauthorized')) + + const store = useWebSocketStore() + await store.connect() + + // Should not have created any WebSocket (ticket exchange failed) + expect(mockInstances).toHaveLength(0) + expect(store.connected).toBe(false) + expect(consoleSpy).toHaveBeenCalledWith( + 'WebSocket ticket exchange failed:', + expect.any(String), + ) + consoleSpy.mockRestore() + }) }) diff --git a/web/src/api/endpoints/auth.ts b/web/src/api/endpoints/auth.ts index 00195f4395..d3b786700a 100644 --- a/web/src/api/endpoints/auth.ts +++ b/web/src/api/endpoints/auth.ts @@ -6,6 +6,7 @@ import type { SetupRequest, TokenResponse, UserInfoResponse, + WsTicketResponse, } from '../types' export async function setup(data: SetupRequest): Promise { @@ -27,3 +28,8 @@ export async function getMe(): Promise { const response = await apiClient.get>('/auth/me') return unwrap(response) } + +export async function getWsTicket(): Promise { + const response = await apiClient.post>('/auth/ws-ticket') + return unwrap(response) +} diff --git a/web/src/api/types.ts b/web/src/api/types.ts index e81e12ee09..3a7a83eaa8 100644 --- a/web/src/api/types.ts +++ b/web/src/api/types.ts @@ -175,6 +175,11 @@ export interface TokenResponse { must_change_password: boolean } +export interface WsTicketResponse { + ticket: string + expires_in: number +} + export interface UserInfoResponse { id: string username: string diff --git a/web/src/composables/useWebSocketSubscription.ts b/web/src/composables/useWebSocketSubscription.ts index cad5fe0d0a..0549f9a692 100644 --- a/web/src/composables/useWebSocketSubscription.ts +++ b/web/src/composables/useWebSocketSubscription.ts @@ -52,20 +52,25 @@ export function useWebSocketSubscription( const setupError = ref(null) const uniqueChannels: WsChannel[] = [...new Set(options.bindings.map((b) => b.channel))] + let disposed = false - onMounted(() => { + onMounted(async () => { if (!authStore.token) return try { if (!wsStore.connected) { - wsStore.connect(authStore.token) + await wsStore.connect() } } catch (err) { + if (disposed) return setupError.value = 'WebSocket connection failed.' console.error('WebSocket connect failed:', sanitizeForLog(err), err) return } + // Component may have unmounted while awaiting connect + if (disposed) return + try { wsStore.subscribe(uniqueChannels, options.filters) } catch (err) { @@ -83,6 +88,7 @@ export function useWebSocketSubscription( }) onUnmounted(() => { + disposed = true try { wsStore.unsubscribe(uniqueChannels) } catch (err) { diff --git a/web/src/stores/websocket.ts b/web/src/stores/websocket.ts index 9b1f30ed90..c133dfb8db 100644 --- a/web/src/stores/websocket.ts +++ b/web/src/stores/websocket.ts @@ -1,6 +1,8 @@ import { defineStore } from 'pinia' import { ref } from 'vue' +import { AxiosError } from 'axios' import type { WsChannel, WsEvent, WsEventHandler, WsSubscriptionFilters } from '@/api/types' +import { getWsTicket } from '@/api/endpoints/auth' import { WS_RECONNECT_BASE_DELAY, WS_RECONNECT_MAX_DELAY, WS_MAX_RECONNECT_ATTEMPTS, WS_MAX_MESSAGE_SIZE } from '@/utils/constants' import { sanitizeForLog } from '@/utils/logging' @@ -25,7 +27,9 @@ export const useWebSocketStore = defineStore('websocket', () => { let reconnectAttempts = 0 let reconnectTimer: ReturnType | null = null let intentionalClose = false - let currentToken: string | null = null + let shouldBeConnected = false + let connectPromise: Promise | null = null + let connectGeneration = 0 const channelHandlers = new Map>() let pendingSubscriptions: { channels: WsChannel[]; filters?: Record }[] = [] // Track active subscriptions so reconnect can re-subscribe automatically @@ -37,15 +41,49 @@ export const useWebSocketStore = defineStore('websocket', () => { return `${protocol}//${host}/api/v1/ws` } - function connect(token: string) { + async function connect() { if (socket?.readyState === WebSocket.OPEN || socket?.readyState === WebSocket.CONNECTING) return - reconnectExhausted.value = false + // Deduplicate concurrent connect() calls — the ticket exchange is async, + // so two callers could pass the readyState guard before the first resolves. + if (connectPromise) return connectPromise + const generation = connectGeneration + connectPromise = doConnect(generation).finally(() => { + // Only clear if no disconnect() happened — disconnect increments + // connectGeneration, so stale finally callbacks are no-ops. + if (generation === connectGeneration) connectPromise = null + }) + return connectPromise + } - currentToken = token + async function doConnect(generation: number) { + reconnectExhausted.value = false + shouldBeConnected = true intentionalClose = false - // TODO(#343): Replace with one-time WS ticket endpoint for production security. - // Currently passes JWT as query param which is logged in server/proxy/browser. - const url = `${getWsUrl()}?token=${encodeURIComponent(token)}` + + // Fetch a one-time ticket from the backend (requires valid JWT in Authorization header). + // The ticket is short-lived (30s) and single-use — browsers cannot set custom headers + // on WebSocket upgrade requests, so a query parameter is the only viable transport. + let ticket: string + try { + const resp = await getWsTicket() + ticket = resp.ticket + } catch (err) { + console.error('WebSocket ticket exchange failed:', sanitizeForLog(err)) + // Don't reconnect on auth failure — the 401 interceptor handles redirect + const isAuthError = err instanceof AxiosError && err.response?.status === 401 + if (shouldBeConnected && !isAuthError) { + scheduleReconnect() + } + return + } + + // Guard against stale connect attempts — if disconnect() was called while + // we were awaiting the ticket, bail out instead of opening a new socket. + if (!shouldBeConnected || generation !== connectGeneration) { + return + } + + const url = `${getWsUrl()}?ticket=${encodeURIComponent(ticket)}` socket = new WebSocket(url) socket.onopen = () => { @@ -105,7 +143,7 @@ export const useWebSocketStore = defineStore('websocket', () => { socket.onclose = () => { connected.value = false socket = null - if (!intentionalClose && currentToken) { + if (!intentionalClose && shouldBeConnected) { scheduleReconnect() } } @@ -129,13 +167,19 @@ export const useWebSocketStore = defineStore('websocket', () => { ) reconnectAttempts++ reconnectTimer = setTimeout(() => { - if (currentToken) connect(currentToken) + if (shouldBeConnected) { + connect().catch((err) => { + console.error('WebSocket reconnect failed:', sanitizeForLog(err)) + }) + } }, delay) } function disconnect() { intentionalClose = true - currentToken = null + shouldBeConnected = false + connectGeneration++ + connectPromise = null reconnectAttempts = 0 if (reconnectTimer) { clearTimeout(reconnectTimer) diff --git a/web/src/views/MeetingLogsPage.vue b/web/src/views/MeetingLogsPage.vue index 42fc1ab39c..4c9a226d3d 100644 --- a/web/src/views/MeetingLogsPage.vue +++ b/web/src/views/MeetingLogsPage.vue @@ -1,5 +1,5 @@