diff --git a/README.md b/README.md index c29eeca8..5a288919 100644 --- a/README.md +++ b/README.md @@ -16,17 +16,13 @@ Named after the [Ruddy Turnstone](https://en.wikipedia.org/wiki/Ruddy_turnstone) Turnstone gives LLMs tools — shell, files, search, web, planning — and orchestrates multi-turn conversations where the model investigates, acts, and reports. It runs as: - **Interactive sessions** — terminal CLI or browser UI with parallel workstreams -- **Queue-driven agents** — trigger workstreams via message queue, stream progress, approve or auto-approve tool use -- **Multi-node clusters** — generic work load-balances across nodes, directed work routes to a specific server - **Cluster dashboard** — real-time view of all nodes and workstreams, reverse proxy for server UIs - **Intent validation** — an LLM judge evaluates every tool call before approval, presenting risk assessments and evidence-based recommendations so users can make informed decisions instead of blindly approving raw tool calls - **Governance & compliance** — RBAC, OIDC SSO (Okta, Azure AD, Google, Keycloak), tool policies, skills (reusable behavioral profiles with security scanning), usage tracking, and append-only audit logs -- **Cluster simulator** — test the stack at scale (up to 1000 nodes) without an LLM backend - Works with any OpenAI-compatible API (vLLM, llama.cpp, NVIDIA NIM) or Anthropic's native Messages API. Supports [MCP](https://modelcontextprotocol.io/) for external tool servers with native deferred tool loading on Anthropic and OpenAI APIs (BM25 fallback for local models).

- Turnstone system architecture — data flow from clients through gateways, Redis MQ, cluster nodes, to LLM providers + Turnstone system architecture — data flow from clients through gateways, cluster nodes, to LLM providers

## Quickstart @@ -44,34 +40,22 @@ turnstone --base-url http://localhost:8000/v1 turnstone-server --port 8080 --base-url http://localhost:8000/v1 ``` -### Queue-driven (programmatic) - -```bash -pip install turnstone[mq] -turnstone-bridge --server-url http://localhost:8080 --redis-host localhost -``` +### Programmatic (SDK) ```python -from turnstone.mq import TurnstoneClient +from turnstone.sdk import TurnstoneServer -with TurnstoneClient() as client: - # Generic — any available node picks it up - result = client.send_and_wait("Analyze the error logs", auto_approve=True) +with TurnstoneServer("http://localhost:8080", token="tok_xxx") as client: + ws = client.create_workstream(name="demo") + result = client.send_and_wait("Analyze the error logs", ws.ws_id, auto_approve=True) print(result.content) - - # Directed — must run on a specific server - result = client.send_and_wait( - "Check disk I/O on this server", - target_node="server-12", - auto_approve=True, - ) ``` ### Cluster dashboard ```bash pip install turnstone[console] -turnstone-console --redis-host localhost --port 8090 +turnstone-console --port 8090 ``` Then open `http://localhost:8090` for the cluster-wide dashboard. Create workstreams from the console and interact with any node's server UI through the built-in reverse proxy — no direct server port access required. @@ -80,7 +64,7 @@ Then open `http://localhost:8090` for the cluster-wide dashboard. Create workstr ```bash cp .env.example .env # edit LLM_BASE_URL, OPENAI_API_KEY, etc. -docker compose up # starts redis + server + bridge + console (SQLite) +docker compose up # starts server + console (SQLite) ``` For production with PostgreSQL: @@ -92,23 +76,6 @@ docker compose --profile production up # adds PostgreSQL, uses it as database Console dashboard at http://localhost:8090. See [docs/docker.md](docs/docker.md) for configuration, scaling, and profiles. -### Simulator - -Test the multi-node stack at scale without an LLM backend: - -```bash -docker compose --profile sim up redis console sim -``` - -Or standalone: - -```bash -pip install turnstone[sim] -turnstone-sim --nodes 100 --scenario steady --duration 60 --mps 10 -``` - -See [docs/simulator.md](docs/simulator.md) for scenarios, CLI reference, and metrics. - ## Architecture ### Diagrams @@ -122,11 +89,7 @@ Detailed UML diagrams are available in [`docs/diagrams/`](docs/diagrams/): | [Core Engine Classes](docs/diagrams/png/03-core-engine-classes.png) | SessionUI protocol, ChatSession, LLMProvider, WorkstreamManager | | [Conversation Turn](docs/diagrams/png/04-conversation-turn.png) | Full message lifecycle through the engine (provider-agnostic) | | [Tool Pipeline](docs/diagrams/png/05-tool-pipeline.png) | Three-phase prepare/approve/execute | -| [MQ Protocol](docs/diagrams/png/06-mq-protocol.png) | 9 inbound + 19 outbound message types | -| [Message Routing](docs/diagrams/png/07-message-routing.png) | Multi-node routing scenarios | -| [Redis Key Schema](docs/diagrams/png/08-redis-key-schema.png) | All Redis keys, types, and TTLs | | [Workstream States](docs/diagrams/png/09-workstream-states.png) | State machine transitions | -| [Simulator](docs/diagrams/png/10-simulator-architecture.png) | SimCluster, dispatchers, scenarios | | [Console Data Flow](docs/diagrams/png/11-console-data-flow.png) | Dashboard data collection threads | | [Deployment](docs/diagrams/png/12-deployment.png) | Docker Compose service topology | | [SDK Architecture](docs/diagrams/png/13-sdk-architecture.png) | Python + TypeScript client libraries | @@ -180,27 +143,6 @@ Tool execution results are evaluated by an output guard before entering the conv See [docs/judge.md](docs/judge.md) for the full guide. -## Multi-node routing - -Each Turnstone server runs a bridge process. Bridges share a Redis instance for coordination: - -| Redis Key | Purpose | -|-----------|---------| -| `turnstone:inbound` | Shared work queue — generic tasks, any node | -| `turnstone:inbound:{node_id}` | Per-node queue — directed tasks | -| `turnstone:ws:{ws_id}` | Workstream ownership — auto-routes follow-ups | -| `turnstone:node:{node_id}` | Node heartbeat + metadata for discovery | -| `turnstone:events:{ws_id}` | Per-workstream event pub/sub | -| `turnstone:events:global` | Global event pub/sub | -| `turnstone:events:cluster` | Cluster-wide state changes (for turnstone-console) | - -**Routing rules:** -1. Message has `target_node` → routes to that node's queue -2. Message has `ws_id` → looks up owner, routes to owning node -3. Neither → shared queue, next available bridge picks it up - -Bridges BLPOP from their per-node queue (priority) then the shared queue. Directed work always takes precedence. - ## Tools 15 built-in tools, 2 agent tools, plus external tools via MCP: @@ -312,15 +254,6 @@ host = "0.0.0.0" port = 8080 max_workstreams = 50 # auto-evicts oldest idle when full -[redis] -host = "localhost" -port = 6379 -password = "" - -[bridge] -server_url = "http://localhost:8080" -node_id = "" # empty = hostname_xxxx - [console] host = "0.0.0.0" port = 8090 @@ -376,7 +309,7 @@ Parallel independent conversations, each with its own session and state: | `◆` | attention | Waiting for approval | | `✖` | error | Something went wrong | -Idle workstreams are automatically cleaned up after 2 hours (configurable). In multi-node deployments, workstream ownership is tracked in Redis — follow-up messages auto-route to the owning node. +Idle workstreams are automatically cleaned up after 2 hours (configurable). ## Monitoring @@ -412,7 +345,6 @@ Per-workstream metrics are labeled by `ws_id` (bounded by `[server].max_workstre - Python 3.11+ - An OpenAI-compatible API endpoint ([vLLM](https://github.com/vllm-project/vllm), [NVIDIA NIM](https://build.nvidia.com/), [llama.cpp](https://github.com/ggml-org/llama.cpp), etc.) or an Anthropic API key -- Redis (for message queue bridge — `pip install turnstone[mq]`) - Anthropic provider (optional — `pip install turnstone[anthropic]`) - PostgreSQL (optional, for production — `pip install turnstone[postgres]`) - Math sandbox packages (optional — `pip install turnstone[sandbox]` for sympy, numpy, scipy, pytest) diff --git a/docs/architecture.md b/docs/architecture.md index b967b957..fec087b7 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -18,10 +18,9 @@ plugs in. |---------|--------|----------|---------| | `turnstone` | `turnstone.cli` | `TerminalUI` | Interactive terminal REPL | | `turnstone-server` | `turnstone.server` | `WebUI` | Browser-based chat (HTTP + SSE) | -| `turnstone-bridge` | `turnstone.mq.bridge` | Bridge | Message queue ↔ HTTP API bridge | | `turnstone-console` | `turnstone.console.server` | ClusterCollector | Cluster dashboard (aggregates all nodes) | | `turnstone-eval` | `turnstone.eval` | `NullUI` | Headless evaluation and prompt optimization | -| `turnstone-channel` | `turnstone.channels.cli` | ChannelAdapter | Channel gateway (Discord, Slack, etc.) via Redis MQ | +| `turnstone-channel` | `turnstone.channels.cli` | ChannelAdapter | Channel gateway (Discord, Slack, etc.) | | `turnstone-admin` | `turnstone.core.admin_cli` | — | Offline user and API token management | --- @@ -42,7 +41,7 @@ turnstone/ __init__.py create_provider() + create_client() factory functions workstream.py Parallel workstream manager (WorkstreamState, Workstream, WorkstreamManager) tools.py Tool schema loader (JSON -> OpenAI function-calling format) - mcp_client.py MCPClientManager — MCP server connections, tool discovery, dynamic refresh, async-sync bridge + mcp_client.py MCPClientManager — MCP server connections, tool discovery, dynamic refresh tool_search.py Dynamic tool search — BM25 index, session-scoped tool visibility watch.py WatchRunner daemon — periodic command polling, condition DSL, result dispatch judge.py Intent validation — heuristic rules + LLM judge, advisory verdicts @@ -74,20 +73,15 @@ turnstone/ _base.py Shared httpx async client, auth, error handling _sync.py Background event loop for sync wrappers _types.py TurnResult + TurnstoneAPIError - mq/ - protocol.py Inbound/outbound message dataclasses (JSON serialization) - broker.py Abstract MessageBroker protocol + RedisBroker - bridge.py Bridge service (queue ↔ turnstone-server HTTP API) - client.py TurnstoneClient library + TurnResult for MQ-based access console/ - collector.py ClusterCollector — aggregates state from all nodes via Redis + HTTP - scheduler.py TaskScheduler — background cron/at scheduler, dispatches via MQ + collector.py ClusterCollector — aggregates state from all nodes via HTTP + scheduler.py TaskScheduler — background cron/at scheduler, dispatches via HTTP server.py Cluster dashboard HTTP server + SSE + CLI entry point static/ Cluster dashboard web UI (page-specific HTML, CSS, JS) channels/ cli.py Unified channel gateway entry point (turnstone-channel) _protocol.py ChannelAdapter protocol, ChannelEvent dataclass - _routing.py ChannelRouter — channel/thread ↔ workstream mapping via MQ + _routing.py ChannelRouter — channel/thread ↔ workstream mapping via HTTP _config.py Base ChannelConfig dataclass discord/ Discord adapter (bot, cog, views, streaming, config) shared_static/ Shared design system (base.css, auth.js, theme.js, toast.js, utils.js, kb.js) @@ -704,8 +698,7 @@ supports_vision = true sub-agents, allowing a cheaper model for autonomous loops **Per-workstream selection:** `POST /v1/api/workstreams/new` accepts an optional -`"model"` field. The bridge `CreateWorkstreamMessage` carries the same field -through the MQ protocol, along with `skill` (skill name) +`"model"` field, along with `skill` (skill name) which can override the model before workstream creation. ### Tool Output Truncation @@ -1205,95 +1198,39 @@ calls `_fg_event.wait()`, which blocks the worker thread until the user switches to that workstream. The `_bg_attention_notify` callback writes a bell + status line to stderr to alert the user. -### Message Queue Bridge - -``` -Main thread Global SSE thread Per-WS SSE threads (×N) -+------------------+ +------------------+ +-------------------+ -| Inbound loop | | GET /events/glob | | GET /events?ws_id | -| BLPOP on Redis | | Parse SSE via | | Parse SSE via | -| | | httpx-sse | | httpx-sse | -| Dispatch to | | Forward state | | Forward content, | -| handler | | changes | | tool results | -| POST to server | | Detect turn | | Handle approval | -| Publish ACK | | completion | | forwarding | -+------------------+ +------------------+ +-------------------+ - | | | - +-- Redis inbound queue +-- Redis pub/sub +-- Redis pub/sub - (RPUSH/BLPOP) (PUBLISH) (PUBLISH) - + response queue - (BLPOP on - approval) -``` - -**Approval flow:** When a per-WS SSE thread receives an `approve_request`, it checks -the workstream's `auto_approve_tools` set. If all requested tools are in the set, the -bridge auto-approves via `POST /v1/api/approve`. Otherwise, it publishes an -`ApprovalRequestEvent` to the outbound channel with a `request_id`, then blocks on -`BLPOP` of a Redis response queue (`turnstone:resp:{request_id}`) until the client pushes -a response or the approval timeout (default 3600s / 1 hour) expires. - -**Cancellation:** The `CancelMessage` (type `"cancel"`) is a routed inbound message. -The bridge dispatches it to `POST /v1/api/cancel` on the server owning the workstream, -which sets the cooperative cancel flag and unblocks any pending approval/plan waits. - -**Completion detection:** The bridge tracks which `correlation_id` maps to which -`ws_id` for active sends. The server accumulates content tokens in the WebUI and -piggybacks the full response text onto the `ws_state → idle` global SSE event. -When the bridge receives this event, it emits a synthetic `TurnCompleteEvent` -carrying the correlation ID and the server-provided `content`. This lets downstream -consumers (e.g. the Discord bot) recover the full response when individual -`ContentEvent`s were missed, and serves as the primary delivery path for -bidirectional notification DM forwarding. - -**Multi-node routing:** Each bridge retrieves its `node_id` from the server's -`/health` endpoint on startup (with exponential backoff retry). The server -generates the `node_id` (`{hostname}_{4hex}`) and is the sole authority for -node identity. The bridge BLPOPs -from both `turnstone:inbound:{node_id}` (directed, priority) and `turnstone:inbound` (shared). -Messages with `target_node` set are pushed to the target's per-node queue. Messages -for existing workstreams are auto-routed via `turnstone:ws:{ws_id}` ownership keys in Redis. -If a bridge picks up a shared-queue message for a workstream owned by another node, it -re-routes to that node's queue (1 extra hop). Bridges publish heartbeats to -`turnstone:node:{node_id}` with configurable TTL for node discovery. -On startup, `_recover_workstreams` re-registers ownership of existing -workstreams and publishes `WorkstreamCreatedEvent` to the cluster channel -so the console collector picks them up immediately. - ### Cluster Console ``` Monitoring (3 daemon threads) Control + Proxy (async Starlette) +------------------+ +----------------------------+ | Event subscriber | | POST /v1/api/cluster/ | -| SUBSCRIBE on | | workstreams/new | -| events:cluster | | → LPUSH to Redis | -+------------------+ | inbound:{node_id} | -| Node discovery | +----------------------------+ -| SCAN node:* keys | | GET /node/{node_id}/ | -| every 15 seconds | | → httpx.AsyncClient | -+------------------+ | proxy to server_url | -| Poll loop | | GET /node/{id}/v1/api/events | -| GET /v1/api/dash | | → SSE stream proxy | -| GET /health | | POST /node/{id}/v1/api/send | -| ThreadPoolExec | | → forwarded to server | +| SSE on | | workstreams/new | +| /events/glob | | → POST to target server | +------------------+ +----------------------------+ +| Node discovery | | GET /node/{node_id}/ | +| Service registry | | → httpx.AsyncClient | +| every 15 seconds | | proxy to server_url | ++------------------+ | GET /node/{id}/v1/api/events | +| Poll loop | | → SSE stream proxy | +| GET /v1/api/dash | | POST /node/{id}/v1/api/send | +| GET /health | | → forwarded to server | +| ThreadPoolExec | +----------------------------+ ++------------------+ ``` The console HTTP layer is a Starlette/ASGI app served by uvicorn. The SSE endpoint uses `EventSourceResponse` with the same listener queue pattern as the main server. `ClusterCollector`'s background threads (event subscriber, -node discovery, poll loop) use sync Redis clients and `ThreadPoolExecutor` +node discovery, poll loop) use `ThreadPoolExecutor` for parallel HTTP polling. The poll loop diffs workstream IDs between poll cycles and fans out synthetic `ws_created`/`ws_closed` SSE events for any changes, ensuring browser clients stay in sync even when real-time cluster -events are missed (e.g. bridge startup recovery). +events are missed. The console has two write-path capabilities: -1. **Workstream creation** — pushes `CreateWorkstreamMessage` to Redis inbound - queues targeting specific nodes. The bridge on each node picks up the message - and creates the workstream on the local server. Auto-selects the node with +1. **Workstream creation** — sends HTTP requests to target server nodes + to create workstreams. Auto-selects the node with the most available capacity if no target is specified. When a `skill` field is present, the server resolves the skill BEFORE `mgr.create()` (applying the model override to the creation request) and snapshot-applies @@ -1360,7 +1297,7 @@ event loop on a daemon thread. **Event types**: 27 standalone dataclasses in `events.py` with a type-registry pattern matching `OutboundEvent.from_json()` from `mq/protocol.py`. Events are -decoupled from the MQ package so SDK consumers don't need the `redis` dependency. +decoupled from server internals. **TypeScript SDK**: `sdk/typescript/` — separate npm package with the same API surface. Zero browser dependencies, SSE via `fetch` + `ReadableStream` parsing. @@ -1381,20 +1318,19 @@ with TurnstoneServer("http://localhost:8080", token="tok_xxx") as client: > See also: [Channel Integrations guide](channels.md) -The `turnstone-channel` gateway bridges external messaging platforms -(Discord, Slack, Teams) to the turnstone cluster via Redis MQ. Each +The `turnstone-channel` gateway connects external messaging platforms +(Discord, Slack, Teams) to the turnstone cluster via HTTP. Each platform adapter implements the `ChannelAdapter` protocol and translates -between platform-native events and turnstone MQ messages. +between platform-native events and turnstone server API calls. The `ChannelRouter` manages bidirectional routing: it maps platform channel/thread IDs to turnstone workstream IDs, handles workstream creation and stale-route recovery, and resolves platform users to turnstone identities via the `channel_users` table. When an evicted workstream is reactivated, the router uses atomic resume via the -`resume_ws` field on `CreateWorkstreamMessage` — the server resumes +`resume_ws` field on the workstream creation request — the server resumes the old workstream's conversation during creation in a single HTTP -request, eliminating ordering fragility. The bridge emits a -`WorkstreamResumedEvent` to confirm success. +request, eliminating ordering fragility. Discord ships as the first adapter. See [channels.md](channels.md) for setup instructions, configuration reference, and the adapter development @@ -1403,7 +1339,7 @@ guide. ### Notification Subsystem The `notify` tool enables the LLM to send notifications to users or -channels without going through MQ. The server calls the channel gateway +channels directly. The server calls the channel gateway directly over HTTP for lower latency: `_exec_notify()` queries the `services` database table for healthy channel gateways (heartbeat within 120 seconds), authenticates with a service JWT (`aud: turnstone-channel`), diff --git a/docs/docker.md b/docs/docker.md index 9e080bff..abafb70e 100644 --- a/docs/docker.md +++ b/docs/docker.md @@ -1,6 +1,6 @@ # Docker Deployment -Docker Compose stack for running the full turnstone platform or the simulator. +Docker Compose stack for running the full turnstone platform. ## Quick Start @@ -10,9 +10,6 @@ cp .env.example .env # Full stack (needs an LLM API on the host) docker compose up - -# Simulator only (no LLM needed) -docker compose --profile sim up redis console sim ``` Console dashboard: http://localhost:8090 @@ -23,18 +20,14 @@ Console dashboard: http://localhost:8090 | Service | Port | Profile | Description | |---------|------|---------|-------------| -| `redis` | 6379 | default | Message broker, pub/sub, node registry | | `server` | 8080 | default | Web UI + chat workstreams + LLM | -| `bridge` | — | default | Redis-to-HTTP bridge (multi-node routing) | | `console` | 8090 | default | Cluster dashboard | | `channel` | — | production | Channel gateway (Discord, Slack, etc.) | | `server-1`…`server-10` | — | cluster | 10-node server fleet (PostgreSQL required) | -| `bridge-1`…`bridge-10` | — | cluster | Matching bridge fleet | -| `sim` | — | sim | Multi-node cluster simulator | ## Profiles -**Default** (no flag) — starts `redis`, `server`, `bridge`, `console`. Requires an OpenAI-compatible LLM API running on the host (default: `http://localhost:8000/v1`). +**Default** (no flag) — starts `server` and `console`. Requires an OpenAI-compatible LLM API running on the host (default: `http://localhost:8000/v1`). ```bash docker compose up @@ -46,22 +39,12 @@ docker compose up docker compose --profile production up ``` -**Cluster** — 10-node server/bridge fleet sharing PostgreSQL and Redis. Access all nodes via the console at `:8090`. Requires `POSTGRES_PASSWORD`: +**Cluster** — 10-node server fleet sharing PostgreSQL. Access all nodes via the console at `:8090`. Requires `POSTGRES_PASSWORD`: ```bash docker compose --profile cluster up ``` -**Sim** — adds the simulator. Can run alongside the full stack or standalone with just Redis and the console: - -```bash -# Sim + console (no LLM needed) -docker compose --profile sim up redis console sim - -# Everything including sim -docker compose --profile sim up -``` - ## Configuration All configuration is via environment variables in `.env` (copy from `.env.example`): @@ -74,13 +57,6 @@ All configuration is via environment variables in `.env` (copy from `.env.exampl | `OPENAI_API_KEY` | `dummy` | API key (`dummy` for local servers) | | `TAVILY_API_KEY` | — | Web search API key (only needed for local/vLLM models; Anthropic and OpenAI search models use native search) | -### Redis - -| Variable | Default | Description | -|----------|---------|-------------| -| `REDIS_PASSWORD` | — | Redis auth password (empty = no auth) | -| `REDIS_PORT` | `6379` | Host port mapping | - ### Server | Variable | Default | Description | @@ -100,7 +76,7 @@ All configuration is via environment variables in `.env` (copy from `.env.exampl | Variable | Default | Description | |----------|---------|-------------| | `TURNSTONE_AUTH_ENABLED` | — | Set to `1` to require authentication | -| `TURNSTONE_AUTH_TOKEN` | — | Config-file token for server/bridge/console (backward compat, works alongside JWT) | +| `TURNSTONE_AUTH_TOKEN` | — | Config-file token for server/console (backward compat, works alongside JWT) | | `TURNSTONE_JWT_SECRET` | — | Secret key for signing JWTs (required when using user identity / JWT auth) | ### Database @@ -130,29 +106,17 @@ The database stores workstream history, user accounts, and API tokens. When usin | `TURNSTONE_DISCORD_TOKEN` | — | Discord bot token (required to enable Discord adapter) | | `TURNSTONE_DISCORD_GUILD` | `0` | Restrict to a single Discord guild (0 = all guilds) | -The channel service runs in the `production` profile. When `TURNSTONE_DISCORD_TOKEN` is set, the Discord adapter connects to the Discord Gateway and routes messages through Redis MQ to the bridge and server. See [Channel Integrations](channels.md) for full setup instructions including Discord application creation and user account linking. - -### Simulator - -| Variable | Default | Description | -|----------|---------|-------------| -| `SIM_NODES` | `100` | Number of simulated nodes | -| `SIM_SCENARIO` | `steady` | Scenario: `steady`, `burst`, `node_failure`, `directed`, `lifecycle` | -| `SIM_DURATION` | `60` | Duration in seconds | -| `SIM_MPS` | `5.0` | Messages per second (steady scenario) | -| `SIM_LOG_LEVEL` | `INFO` | Log verbosity | -| `SIM_SEED` | — | Random seed for reproducibility | -| `SIM_METRICS_FILE` | — | Write JSON report to file | +The channel service runs in the `production` profile. When `TURNSTONE_DISCORD_TOKEN` is set, the Discord adapter connects to the Discord Gateway and routes messages to the server via HTTP. See [Channel Integrations](channels.md) for full setup instructions including Discord application creation and user account linking. ## Scaling -For multi-node testing, use the `cluster` profile which provides 10 dedicated server+bridge pairs with unique node IDs (`node-1` through `node-10`), resource limits, and shared PostgreSQL: +For multi-node testing, use the `cluster` profile which provides 10 server instances with unique node IDs (`node-1` through `node-10`), resource limits, and shared PostgreSQL: ```bash POSTGRES_PASSWORD=secret docker compose --profile cluster up ``` -The default `server` and `bridge` also run alongside the cluster nodes (11 total). All nodes are accessible via the console dashboard at `:8090`. +The default `server` also runs alongside the cluster nodes (11 total). All nodes are accessible via the console dashboard at `:8090`. For production clusters beyond ~50 nodes, add PgBouncer between turnstone services and PostgreSQL. See [PgBouncer Connection Pooling](pgbouncer.md) for Docker Compose and Helm configuration. @@ -160,7 +124,6 @@ For production clusters beyond ~50 nodes, add PgBouncer between turnstone servic | Volume | Mount | Purpose | |--------|-------|---------| -| `redis-data` | `/data` | Redis persistence | | `turnstone-data` | `/data` | SQLite database (`.turnstone.db`) | ## Building @@ -175,7 +138,7 @@ docker compose build docker compose build --no-cache ``` -All entry points are installed in a single image: `turnstone-server`, `turnstone-bridge`, `turnstone-console`, `turnstone-channel`, `turnstone-admin`, `turnstone-sim`, `turnstone-eval`. +All entry points are installed in a single image: `turnstone-server`, `turnstone-console`, `turnstone-channel`, `turnstone-admin`, `turnstone-eval`. ## Cleanup diff --git a/docs/judge.md b/docs/judge.md index 69095234..137e7120 100644 --- a/docs/judge.md +++ b/docs/judge.md @@ -379,7 +379,7 @@ emitted to the frontend: ``` The web UI renders this as an inline warning after the tool result. The CLI -shows a colored terminal warning. The MQ bridge forwards it as an +shows a colored terminal warning. The server forwards it as an `OutputWarningEvent` for console subscribers. Assessments are persisted to the `output_assessments` table for v2 diff --git a/docs/pgbouncer.md b/docs/pgbouncer.md index 13e55547..82ec08d2 100644 --- a/docs/pgbouncer.md +++ b/docs/pgbouncer.md @@ -1,7 +1,7 @@ # PgBouncer Connection Pooling Turnstone cluster deployments share a single PostgreSQL instance across -all server nodes, bridge processes, and the console. Each process +all server nodes and the console. Each process maintains a small connection pool (2 base + 3 overflow = 5 max). At scale this adds up — a 100-node cluster opens up to 500 connections, and a 1000-node cluster up to 5,000. diff --git a/docs/simulator.md b/docs/simulator.md deleted file mode 100644 index ef49b450..00000000 --- a/docs/simulator.md +++ /dev/null @@ -1,204 +0,0 @@ -# Cluster Simulator - -The simulator (`turnstone-sim`) creates lightweight simulated nodes that talk to a real Redis instance using the standard turnstone protocol. External observers — `TurnstoneClient`, `turnstone-console`, real bridges — see identical behavior. No LLM backend is needed. - -## Quick Start - -```bash -pip install turnstone[sim] - -# 10 nodes, steady load, 60 seconds -turnstone-sim --nodes 10 --scenario steady --duration 60 --mps 5 - -# 100 nodes via Docker -docker compose --profile sim up redis console sim -``` - -## How It Works - -Each simulated node is an asyncio coroutine (not a thread or process), so 1000 nodes run efficiently on a single event loop. The simulator: - -1. Registers nodes via Redis heartbeats (same keys as real bridges) -2. Accepts messages from per-node and shared inbound queues -3. Simulates LLM responses with configurable latency and token generation -4. Simulates tool execution with configurable latency and failure rates -5. Publishes real protocol events (`ContentEvent`, `StateChangeEvent`, `TurnCompleteEvent`, etc.) -6. Reports latency, throughput, and utilization metrics at completion - -``` -TurnstoneClient → Redis Queue → SimNode → Redis Pub/Sub → TurnstoneClient - ↓ - turnstone-console (cluster dashboard) -``` - -## Scenarios - -| Scenario | Description | -|----------|-------------| -| `steady` | Inject messages at a constant rate (`--mps`) for `--duration` seconds | -| `burst` | Push `--burst-size` messages instantly, then wait for completion | -| `node_failure` | Steady load + periodically kill nodes to test redistribution | -| `directed` | Send messages to specific nodes via `target_node` routing | -| `lifecycle` | Create, use, and close workstreams across nodes | - -## CLI Reference - -``` -turnstone-sim [options] -``` - -### Cluster - -| Flag | Default | Description | -|------|---------|-------------| -| `--nodes` | `10` | Number of simulated nodes | - -### Scenario - -| Flag | Default | Description | -|------|---------|-------------| -| `--scenario` | `steady` | Scenario name | -| `--duration` | `60` | Duration in seconds | -| `--mps` | `5.0` | Messages per second (steady) | -| `--burst-size` | `100` | Messages to send (burst) | -| `--node-kill-interval` | `15` | Seconds between kills (node_failure) | -| `--node-kill-count` | `1` | Nodes per kill cycle | - -### Simulation - -| Flag | Default | Description | -|------|---------|-------------| -| `--llm-latency` | `2.0` | Mean LLM response latency (seconds) | -| `--tool-latency` | `0.5` | Mean tool execution latency (seconds) | -| `--tool-failure-rate` | `0.02` | Tool failure probability (0.0–1.0) | -| `--seed` | — | Random seed for reproducibility | - -### Redis - -| Flag | Default | Description | -|------|---------|-------------| -| `--redis-host` | `localhost` | Redis host | -| `--redis-port` | `6379` | Redis port | -| `--redis-password` | — | Redis password | -| `--prefix` | `turnstone` | Redis key prefix | - -### Output - -| Flag | Default | Description | -|------|---------|-------------| -| `--metrics-file` | — | Write JSON report to file | -| `--log-level` | `INFO` | Log verbosity | - -## Example: Load Testing - -```bash -# 100 nodes, high throughput, 2 minutes -turnstone-sim --nodes 100 --scenario steady --duration 120 --mps 50 - -# Burst of 500 messages across 50 nodes -turnstone-sim --nodes 50 --scenario burst --burst-size 500 --duration 60 - -# Node failure resilience (kill 2 nodes every 10 seconds) -turnstone-sim --nodes 20 --scenario node_failure --duration 120 \ - --node-kill-interval 10 --node-kill-count 2 - -# Fast simulation (low latency, no failures) -turnstone-sim --nodes 10 --scenario steady --duration 30 \ - --llm-latency 0.1 --tool-latency 0.05 --tool-failure-rate 0 --mps 10 -``` - -## Metrics Report - -The simulator prints a summary at completion: - -``` -============================================================ - SIMULATION REPORT -============================================================ - Scenario: steady - Nodes: 100 - Duration: 60.2s - Total turns: 295 - Total errors: 5 - Node kills: 0 ------------------------------------------------------------- - THROUGHPUT - Messages/sec: 4.97 - Turns/sec: 4.89 ------------------------------------------------------------- - LATENCY (seconds) - p50: 3.21 - p90: 5.44 - p99: 8.12 - mean: 3.56 - max: 12.1 ------------------------------------------------------------- - UTILIZATION - Mean ws/node: 2.3 - Max ws/node: 8 - Idle nodes: 12 -============================================================ -``` - -Use `--metrics-file report.json` to write the full report as JSON. - -## Console Integration - -The simulator's nodes appear in `turnstone-console` exactly like real nodes. Run them together to see the dashboard populate with simulated workstreams: - -```bash -# Terminal 1: start Redis and console -docker compose up redis console - -# Terminal 2: run simulator -docker compose --profile sim up sim -``` - -Or all at once: - -```bash -SIM_NODES=50 SIM_DURATION=120 docker compose --profile sim up redis console sim -``` - -Open http://localhost:8090 to see simulated nodes, workstream states, token counts, and load bars updating in real time. - -## Architecture - -> See also: [Simulator Architecture diagram](diagrams/png/10-simulator-architecture.png) - -``` -turnstone/sim/ -├── __init__.py # Public API: SimCluster, SimConfig -├── config.py # SimConfig — all simulation parameters -├── engine.py # SimEngine — LLM + tool execution simulation -├── node.py # SimNode + SimWorkstream — protocol-compatible node -├── cluster.py # SimCluster + InboundDispatcher + PooledBroker -├── scenario.py # 5 scenario classes -├── metrics.py # MetricsCollector — latency, throughput, utilization -└── cli.py # CLI entry point -``` - -**Key design:** The `InboundDispatcher` batches ~50 node queues into a single Redis `BLPOP` call, keeping connection count bounded at ~20 regardless of node count. All nodes share a single `ConnectionPool(max_connections=64)`. - -## Programmatic Use - -```python -import asyncio -from turnstone.sim import SimCluster, SimConfig - -async def main(): - config = SimConfig( - num_nodes=10, - scenario="steady", - duration=30, - messages_per_second=2.0, - llm_latency_mean=0.5, - ) - cluster = SimCluster(config) - await cluster.start() - await cluster.run_scenario() - print(cluster.report()) - await cluster.stop() - -asyncio.run(main()) -``` diff --git a/docs/tls.md b/docs/tls.md index 82dc0f0b..8aae3842 100644 --- a/docs/tls.md +++ b/docs/tls.md @@ -12,7 +12,7 @@ docker compose -f compose.yaml -f deploy/docker-compose.tls.yml up ``` This: -1. Bootstraps an internal CA and issues certs for Redis/PostgreSQL +1. Bootstraps an internal CA and issues certs for PostgreSQL 2. Starts the console with TLS enabled (internal CA + ACME server) 3. Server nodes auto-provision certs via the console's ACME endpoint 4. All inter-service communication uses mTLS @@ -30,9 +30,9 @@ Console (CA + ACME Server) | ACME protocol (auto-approve, no challenge validation) +-----------+-----------+ | | | - Server(s) Bridge Channel GW - (auto-cert (mTLS (mTLS - + renewal) client) client) + Server(s) Channel GW + (auto-cert (mTLS + + renewal) client) ``` **Two cert paths on the console:** @@ -57,12 +57,6 @@ Console (CA + ACME Server) These are needed before storage is available: ```toml -[redis] -tls = false -tls_ca = "" # path to CA cert -tls_cert = "" # path to client cert -tls_key = "" # path to client key - [database] sslmode = "prefer" # disable, allow, prefer, require, verify-full sslrootcert = "" # path to CA cert @@ -89,12 +83,11 @@ sslkey = "" # path to client key Create a CA and infrastructure certs without a running console: ```bash -# Bootstrap CA + Redis + PostgreSQL certs -turnstone-admin tls-bootstrap --out /certs --issue redis --issue postgres +# Bootstrap CA + PostgreSQL certs +turnstone-admin tls-bootstrap --out /certs --issue postgres # Output: # /certs/ca.pem (CA root certificate) -# /certs/certs/redis/ (Redis cert + key) # /certs/certs/postgres/ (PostgreSQL cert + key) ``` diff --git a/examples/mcp-cluster-ops/README.md b/examples/mcp-cluster-ops/README.md index 3011f049..61aa5efd 100644 --- a/examples/mcp-cluster-ops/README.md +++ b/examples/mcp-cluster-ops/README.md @@ -1,10 +1,10 @@ # MCP Cluster Ops -An MCP server that exposes tools for executing commands across a [Turnstone](https://github.com/turnstonelabs/turnstone) cluster. Serves as a reference implementation for both MCP server patterns and Turnstone MQ client SDK usage. +An MCP server that exposes tools for executing commands across a [Turnstone](https://github.com/turnstonelabs/turnstone) cluster. Serves as a reference implementation for both MCP server patterns and Turnstone SDK usage. ## How it works -This server uses Turnstone's MQ client (`TurnstoneClient`) to dispatch shell commands to specific nodes via Redis. Remote agents execute the command and the raw bash output is captured directly from the `ToolResultEvent` stream — bypassing the costly "agent reads output → re-generates output as completion tokens" round-trip. +This server uses Turnstone's SDK client (`TurnstoneServer`) to dispatch shell commands to specific nodes via HTTP. Remote agents execute the command and the raw bash output is captured directly from the `ToolResultEvent` stream — bypassing the costly "agent reads output → re-generates output as completion tokens" round-trip. Multi-node dispatches run in parallel via `asyncio.gather`, so total wall time is bounded by the slowest node rather than the sum. @@ -19,8 +19,7 @@ Multi-node dispatches run in parallel via `asyncio.gather`, so total wall time i ## Prerequisites -- A running Turnstone cluster (at least one `turnstone-server` + `turnstone-bridge`) -- Redis accessible from wherever this MCP server runs +- A running Turnstone cluster (at least one `turnstone-server`) - Python 3.11+ ## Installation @@ -28,10 +27,6 @@ Multi-node dispatches run in parallel via `asyncio.gather`, so total wall time i ```bash # From the turnstone repo root: pip install -e ./examples/mcp-cluster-ops - -# Or install turnstone with MQ support first, then the example: -pip install -e ".[mq]" -pip install -e ./examples/mcp-cluster-ops ``` ## Configuration @@ -40,9 +35,8 @@ pip install -e ./examples/mcp-cluster-ops | Variable | Default | Description | |----------|---------|-------------| -| `REDIS_HOST` | `localhost` | Redis host | -| `REDIS_PORT` | `6379` | Redis port | -| `REDIS_PASSWORD` | _(none)_ | Redis password (use env vars, not config files) | +| `TURNSTONE_SERVER_URL` | `http://localhost:8080` | Server URL | +| `TURNSTONE_API_TOKEN` | _(none)_ | API token for authentication | | `MCP_CLUSTER_OPS_TIMEOUT` | `120` | Default command timeout (seconds, clamped 5-3600) | | `MCP_CLUSTER_OPS_MAX_OUTPUT` | `8192` | Max output bytes per node (0 = unlimited) | | `MCP_CLUSTER_OPS_MAX_NODES` | `32` | Max concurrent node dispatches | @@ -57,7 +51,7 @@ pip install -e ./examples/mcp-cluster-ops command = "mcp-cluster-ops" [mcp.servers.cluster-ops.env] -REDIS_HOST = "redis.example.com" +TURNSTONE_SERVER_URL = "http://turnstone.example.com:8080" ``` **JSON** (via `--mcp-config`): @@ -68,7 +62,7 @@ REDIS_HOST = "redis.example.com" "cluster-ops": { "command": "mcp-cluster-ops", "env": { - "REDIS_HOST": "redis.example.com" + "TURNSTONE_SERVER_URL": "http://turnstone.example.com:8080" } } } @@ -90,10 +84,6 @@ node-2: /dev/sda1 500G 410G 90G 82% / node-3: /dev/sda1 1.0T 200G 800G 20% / ``` -## Why MQ client instead of HTTP SDK? - -The HTTP SDK (`TurnstoneServer`) talks to a single server instance. The MQ client (`TurnstoneClient`) routes through Redis with `target_node` support, which is the entire point of cross-node cluster operations. - ## Security Considerations **This MCP server grants the calling agent shell access to cluster nodes.** @@ -104,8 +94,8 @@ The HTTP SDK (`TurnstoneServer`) talks to a single server instance. The MQ clien is returned through the MCP tool result and becomes part of the LLM context. - The security boundary is at the MCP host layer -- use Turnstone's tool policy system to restrict which agents can invoke these tools. -- Set `REDIS_PASSWORD` via your environment or a secrets manager -- avoid - hardcoding passwords in config files. +- Set `TURNSTONE_API_TOKEN` via your environment or a secrets manager -- avoid + hardcoding tokens in config files. ## Development diff --git a/examples/mcp-cluster-ops/mcp_cluster_ops/server.py b/examples/mcp-cluster-ops/mcp_cluster_ops/server.py index 0e2c4663..f971cc13 100644 --- a/examples/mcp-cluster-ops/mcp_cluster_ops/server.py +++ b/examples/mcp-cluster-ops/mcp_cluster_ops/server.py @@ -1,7 +1,7 @@ """MCP server for Turnstone cluster operations. Exposes tools to execute commands on specific nodes in a Turnstone cluster. -Uses the MQ client (``TurnstoneClient``) for direct node targeting via Redis. +Uses the SDK client (``TurnstoneServer``) for direct node targeting via HTTP. Usage:: @@ -14,22 +14,20 @@ Configure in ``~/.config/turnstone/config.toml``:: command = "mcp-cluster-ops" [mcp.servers.cluster-ops.env] - REDIS_HOST = "redis.example.com" + TURNSTONE_SERVER_URL = "http://localhost:8080" Environment variables --------------------- -REDIS_HOST Redis host (default: localhost) -REDIS_PORT Redis port (default: 6379) -REDIS_PASSWORD Redis password (default: none) +TURNSTONE_SERVER_URL Server URL (default: http://localhost:8080) +TURNSTONE_API_TOKEN API token for authentication (default: none) MCP_CLUSTER_OPS_TIMEOUT Default command timeout in seconds (default: 120) MCP_CLUSTER_OPS_MAX_OUTPUT Max output bytes per node (default: 8192, 0=unlimited) Performance notes ----------------- Remote agents are told to reply with only "ok" or "failed" — the raw bash -output is captured directly from the ToolResultEvent that already flows -through Redis, bypassing the costly "agent reads output then re-generates -output as completion tokens" round-trip. +output is captured directly from the ToolResultEvent, bypassing the costly +"agent reads output then re-generates output as completion tokens" round-trip. All multi-node dispatches run in parallel via ``asyncio.gather`` so total wall time is bounded by the slowest node, not the sum of all nodes. @@ -45,7 +43,7 @@ from contextlib import asynccontextmanager from typing import TYPE_CHECKING, Any from mcp.server.fastmcp import Context, FastMCP -from turnstone.mq.client import TurnResult, TurnstoneClient +from turnstone.sdk import TurnResult, TurnstoneServer if TYPE_CHECKING: from collections.abc import AsyncIterator @@ -68,19 +66,14 @@ _MAX_TIMEOUT = 3600 # --------------------------------------------------------------------------- -def _redis_kwargs() -> dict[str, Any]: - """Build Redis connection kwargs from environment variables. - - Follows the same env var convention as ``turnstone.mq.broker.add_redis_args``: - ``REDIS_HOST``, ``REDIS_PORT``, ``REDIS_PASSWORD``. - """ - kwargs: dict[str, Any] = {"host": os.environ.get("REDIS_HOST", "localhost")} - port = os.environ.get("REDIS_PORT") - if port is not None: - kwargs["port"] = int(port) - password = os.environ.get("REDIS_PASSWORD") - if password: - kwargs["password"] = password +def _server_kwargs() -> dict[str, Any]: + """Build TurnstoneServer connection kwargs from environment variables.""" + kwargs: dict[str, Any] = { + "base_url": os.environ.get("TURNSTONE_SERVER_URL", "http://localhost:8080"), + } + token = os.environ.get("TURNSTONE_API_TOKEN") + if token: + kwargs["token"] = token return kwargs @@ -172,12 +165,12 @@ def _format_node_result( # --------------------------------------------------------------------------- -# Core dispatch functions (testable with mocked TurnstoneClient) +# Core dispatch functions (testable with mocked TurnstoneServer) # --------------------------------------------------------------------------- def _exec_on_node_sync( - redis_kw: dict[str, Any], + server_kw: dict[str, Any], node_id: str, command: str, timeout: float, @@ -185,11 +178,11 @@ def _exec_on_node_sync( """Dispatch *command* to *node_id* and block until complete. Runs inside ``asyncio.to_thread`` so it does not block the event loop. - Each call creates its own ``TurnstoneClient`` to avoid Redis pub/sub - subscription conflicts between concurrent dispatches. + Each call creates its own ``TurnstoneServer`` client to avoid state + conflicts between concurrent dispatches. """ prompt = _exec_prompt(command) - with TurnstoneClient(**redis_kw) as client: + with TurnstoneServer(**server_kw) as client: result = client.send_and_wait( message=prompt, target_node=node_id, @@ -200,7 +193,7 @@ def _exec_on_node_sync( async def _dispatch_parallel( - redis_kw: dict[str, Any], + server_kw: dict[str, Any], node_ids: list[str], command: str, timeout: float, @@ -211,7 +204,7 @@ async def _dispatch_parallel( Total wall time is bounded by the slowest node. """ tasks = [ - asyncio.to_thread(_exec_on_node_sync, redis_kw, nid, command, timeout) for nid in node_ids + asyncio.to_thread(_exec_on_node_sync, server_kw, nid, command, timeout) for nid in node_ids ] outcomes = await asyncio.gather(*tasks, return_exceptions=True) @@ -227,16 +220,16 @@ async def _dispatch_parallel( return results -def _list_nodes_sync(redis_kw: dict[str, Any]) -> list[dict[str, Any]]: +def _list_nodes_sync(server_kw: dict[str, Any]) -> list[dict[str, Any]]: """List active cluster nodes (blocking).""" - with TurnstoneClient(**redis_kw) as client: + with TurnstoneServer(**server_kw) as client: nodes: list[dict[str, Any]] = client.list_nodes() return nodes -async def _list_nodes_impl(redis_kw: dict[str, Any]) -> list[dict[str, Any]]: +async def _list_nodes_impl(server_kw: dict[str, Any]) -> list[dict[str, Any]]: """List active cluster nodes.""" - return await asyncio.to_thread(_list_nodes_sync, redis_kw) + return await asyncio.to_thread(_list_nodes_sync, server_kw) # --------------------------------------------------------------------------- @@ -246,9 +239,9 @@ async def _list_nodes_impl(redis_kw: dict[str, Any]) -> list[dict[str, Any]]: @asynccontextmanager async def _lifespan(server: FastMCP[dict[str, Any]]) -> AsyncIterator[dict[str, Any]]: - """Lifespan context — stores Redis kwargs for tool handlers.""" - kw = _redis_kwargs() - yield {"redis_kwargs": kw} + """Lifespan context — stores server connection kwargs for tool handlers.""" + kw = _server_kwargs() + yield {"server_kwargs": kw} mcp = FastMCP( @@ -270,8 +263,8 @@ async def list_nodes(ctx: Context[Any, Any, Any]) -> str: Call this before dispatching work to discover available node IDs. Returns a JSON array of node metadata objects. """ - redis_kw: dict[str, Any] = ctx.request_context.lifespan_context["redis_kwargs"] - nodes = await _list_nodes_impl(redis_kw) + server_kw: dict[str, Any] = ctx.request_context.lifespan_context["server_kwargs"] + nodes = await _list_nodes_impl(server_kw) return json.dumps(nodes, indent=2) @@ -298,12 +291,12 @@ async def run_on_node( if cmd_err: return json.dumps({"error": cmd_err}) - redis_kw: dict[str, Any] = ctx.request_context.lifespan_context["redis_kwargs"] + server_kw: dict[str, Any] = ctx.request_context.lifespan_context["server_kwargs"] max_output = _DEFAULT_MAX_OUTPUT log.info("run_on_node node=%s cmd=%r", node_id, command) _, result = await asyncio.to_thread( - _exec_on_node_sync, redis_kw, node_id, command, _clamp_timeout(timeout) + _exec_on_node_sync, server_kw, node_id, command, _clamp_timeout(timeout) ) formatted = _format_node_result(node_id, result, max_output) return json.dumps(formatted, indent=2) @@ -330,7 +323,7 @@ async def run_on_nodes( if cmd_err: return json.dumps({"error": cmd_err}) - redis_kw: dict[str, Any] = ctx.request_context.lifespan_context["redis_kwargs"] + server_kw: dict[str, Any] = ctx.request_context.lifespan_context["server_kwargs"] max_output = _DEFAULT_MAX_OUTPUT clean_ids = list(dict.fromkeys(nid.strip() for nid in node_ids if nid.strip())) @@ -343,7 +336,7 @@ async def run_on_nodes( log.info("run_on_nodes nodes=%s cmd=%r", clean_ids, command) results = await _dispatch_parallel( - redis_kw, clean_ids, command, _clamp_timeout(timeout), max_output + server_kw, clean_ids, command, _clamp_timeout(timeout), max_output ) return json.dumps(results, indent=2) @@ -368,10 +361,10 @@ async def run_on_all_nodes( if cmd_err: return json.dumps({"error": cmd_err}) - redis_kw: dict[str, Any] = ctx.request_context.lifespan_context["redis_kwargs"] + server_kw: dict[str, Any] = ctx.request_context.lifespan_context["server_kwargs"] max_output = _DEFAULT_MAX_OUTPUT - nodes = await _list_nodes_impl(redis_kw) + nodes = await _list_nodes_impl(server_kw) if not nodes: return json.dumps({"error": "No active nodes found in cluster"}) @@ -388,7 +381,7 @@ async def run_on_all_nodes( ) log.info("run_on_all_nodes nodes=%s cmd=%r", node_ids, command) results = await _dispatch_parallel( - redis_kw, node_ids, command, _clamp_timeout(timeout), max_output + server_kw, node_ids, command, _clamp_timeout(timeout), max_output ) return json.dumps(results, indent=2) diff --git a/examples/mcp-cluster-ops/pyproject.toml b/examples/mcp-cluster-ops/pyproject.toml index 97da022b..3d8ae4ad 100644 --- a/examples/mcp-cluster-ops/pyproject.toml +++ b/examples/mcp-cluster-ops/pyproject.toml @@ -9,7 +9,7 @@ description = "MCP server for Turnstone cluster operations — reference impleme requires-python = ">=3.11" license = "BUSL-1.1" dependencies = [ - "turnstone[mq]", + "turnstone", "mcp>=1.6", ] @@ -18,7 +18,7 @@ mcp-cluster-ops = "mcp_cluster_ops.server:main" [project.optional-dependencies] test = ["pytest>=9.0"] -dev = ["ruff>=0.9", "mypy>=1.14", "types-redis>=4.6"] +dev = ["ruff>=0.9", "mypy>=1.14"] [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/examples/mcp-cluster-ops/tests/test_helpers.py b/examples/mcp-cluster-ops/tests/test_helpers.py index 60ab2e7e..9fee7624 100644 --- a/examples/mcp-cluster-ops/tests/test_helpers.py +++ b/examples/mcp-cluster-ops/tests/test_helpers.py @@ -2,7 +2,7 @@ from __future__ import annotations -from turnstone.mq.client import TurnResult +from turnstone.sdk import TurnResult from mcp_cluster_ops.server import ( _clamp_timeout, diff --git a/examples/mcp-cluster-ops/tests/test_tools.py b/examples/mcp-cluster-ops/tests/test_tools.py index 314c1e7a..8e52bd15 100644 --- a/examples/mcp-cluster-ops/tests/test_tools.py +++ b/examples/mcp-cluster-ops/tests/test_tools.py @@ -1,4 +1,4 @@ -"""Tests for MCP tool handlers with mocked TurnstoneClient.""" +"""Tests for MCP tool handlers with mocked TurnstoneServer.""" from __future__ import annotations @@ -6,7 +6,7 @@ import asyncio from typing import Any from unittest.mock import MagicMock, patch -from turnstone.mq.client import TurnResult +from turnstone.sdk import TurnResult from mcp_cluster_ops.server import ( _dispatch_parallel, @@ -22,7 +22,7 @@ from mcp_cluster_ops.server import ( class TestListNodesImpl: def test_returns_nodes(self): nodes = [{"node_id": "a", "model": "gpt-5"}, {"node_id": "b", "model": "gpt-5"}] - with patch("mcp_cluster_ops.server.TurnstoneClient") as mock_cls: + with patch("mcp_cluster_ops.server.TurnstoneServer") as mock_cls: mock_client = MagicMock() mock_client.list_nodes.return_value = nodes mock_cls.return_value.__enter__ = MagicMock(return_value=mock_client) @@ -32,7 +32,7 @@ class TestListNodesImpl: assert result == nodes def test_empty_cluster(self): - with patch("mcp_cluster_ops.server.TurnstoneClient") as mock_cls: + with patch("mcp_cluster_ops.server.TurnstoneServer") as mock_cls: mock_client = MagicMock() mock_client.list_nodes.return_value = [] mock_cls.return_value.__enter__ = MagicMock(return_value=mock_client) @@ -52,7 +52,7 @@ class TestExecOnNodeSync: turn_result = TurnResult( tool_results=[("bash", "hello world")], ) - with patch("mcp_cluster_ops.server.TurnstoneClient") as mock_cls: + with patch("mcp_cluster_ops.server.TurnstoneServer") as mock_cls: mock_client = MagicMock() mock_client.send_and_wait.return_value = turn_result mock_cls.return_value.__enter__ = MagicMock(return_value=mock_client) @@ -70,7 +70,7 @@ class TestExecOnNodeSync: def test_timeout(self): turn_result = TurnResult(timed_out=True) - with patch("mcp_cluster_ops.server.TurnstoneClient") as mock_cls: + with patch("mcp_cluster_ops.server.TurnstoneServer") as mock_cls: mock_client = MagicMock() mock_client.send_and_wait.return_value = turn_result mock_cls.return_value.__enter__ = MagicMock(return_value=mock_client) @@ -88,7 +88,7 @@ class TestExecOnNodeSync: class TestDispatchParallel: def test_parallel_success(self): - def fake_exec(redis_kw: Any, node_id: str, command: str, timeout: float) -> Any: + def fake_exec(server_kw: Any, node_id: str, command: str, timeout: float) -> Any: return (node_id, TurnResult(tool_results=[("bash", f"output-{node_id}")])) with patch("mcp_cluster_ops.server._exec_on_node_sync", side_effect=fake_exec): @@ -108,9 +108,9 @@ class TestDispatchParallel: assert outputs["b"] == "output-b" def test_partial_failure(self): - def fake_exec(redis_kw: Any, node_id: str, command: str, timeout: float) -> Any: + def fake_exec(server_kw: Any, node_id: str, command: str, timeout: float) -> Any: if node_id == "bad": - raise ConnectionError("Redis down") + raise ConnectionError("connection refused") return (node_id, TurnResult(tool_results=[("bash", "ok")])) with patch("mcp_cluster_ops.server._exec_on_node_sync", side_effect=fake_exec): @@ -128,10 +128,10 @@ class TestDispatchParallel: bad = next(r for r in results if r["node"] == "bad") assert good["ok"] is True assert bad["ok"] is False - assert "Redis down" in bad["error"] + assert "connection refused" in bad["error"] def test_all_fail(self): - def fake_exec(redis_kw: Any, node_id: str, command: str, timeout: float) -> Any: + def fake_exec(server_kw: Any, node_id: str, command: str, timeout: float) -> Any: raise RuntimeError(f"fail-{node_id}") with patch("mcp_cluster_ops.server._exec_on_node_sync", side_effect=fake_exec): diff --git a/pyproject.toml b/pyproject.toml index f855faeb..102ce67b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,25 +45,21 @@ Issues = "https://github.com/turnstonelabs/turnstone/issues" [project.optional-dependencies] test = ["pytest>=9.0", "pytest-cov>=6.0", "croniter>=3.0"] -dev = ["ruff>=0.9", "mypy>=1.14", "types-redis>=4.6"] -mq = ["redis>=7.2"] -console = ["redis>=7.2", "croniter>=3.0"] -sim = ["redis>=7.2"] +dev = ["ruff>=0.9", "mypy>=1.14"] +console = ["croniter>=3.0"] anthropic = ["anthropic>=0.39"] postgres = ["psycopg[binary]>=3.2"] ddg = ["ddgs>=9.0"] -discord = ["discord.py>=2.4", "redis>=7.2"] +discord = ["discord.py>=2.4"] tls = ["lacme>=1.0.4"] sandbox = ["sympy>=1.13", "numpy>=2.0", "scipy>=1.14", "pytest>=9.0"] -all = ["turnstone[mq,console,sim,anthropic,postgres,discord,ddg,tls,sandbox]"] +all = ["turnstone[console,anthropic,postgres,discord,ddg,tls,sandbox]"] [project.scripts] turnstone = "turnstone.cli:main" turnstone-eval = "turnstone.eval:main" turnstone-server = "turnstone.server:main" -turnstone-bridge = "turnstone.mq.bridge:main" turnstone-console = "turnstone.console.server:main" -turnstone-sim = "turnstone.sim.cli:main" turnstone-admin = "turnstone.admin:main" turnstone-channel = "turnstone.channels.cli:main" turnstone-bootstrap = "turnstone.bootstrap:main" diff --git a/tests/test_api_versioning.py b/tests/test_api_versioning.py index f3ab10fb..6d3e452c 100644 --- a/tests/test_api_versioning.py +++ b/tests/test_api_versioning.py @@ -84,7 +84,6 @@ class TestConsoleVersioning: } app = create_app( collector=collector, - broker=MagicMock(), auth_config=AuthConfig(), ) client = TestClient(app, raise_server_exceptions=False) diff --git a/tests/test_async_broker.py b/tests/test_async_broker.py deleted file mode 100644 index 3a346f67..00000000 --- a/tests/test_async_broker.py +++ /dev/null @@ -1,181 +0,0 @@ -"""Tests for turnstone.mq.async_broker.AsyncRedisBroker.""" - -from __future__ import annotations - -import asyncio -import contextlib -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -from turnstone.mq.async_broker import AsyncRedisBroker - - -@pytest.fixture -def broker() -> AsyncRedisBroker: - return AsyncRedisBroker(host="localhost", port=6379, db=0, prefix="test", response_ttl=120) - - -@pytest.fixture -def mock_redis() -> AsyncMock: - """Return a mock Redis client with common async methods.""" - r = AsyncMock() - r.rpush = AsyncMock() - r.publish = AsyncMock() - r.expire = AsyncMock() - r.get = AsyncMock(return_value=None) - r.set = AsyncMock() - r.delete = AsyncMock() - r.blpop = AsyncMock(return_value=None) - ps = AsyncMock() - ps.subscribe = AsyncMock() - ps.unsubscribe = AsyncMock() - ps.close = AsyncMock() - ps.get_message = AsyncMock(return_value=None) - r.pubsub = MagicMock(return_value=ps) - return r - - -def _inject_redis(broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - """Inject a mock Redis client into the broker, simulating connect().""" - broker._redis = mock_redis - broker._pubsub = mock_redis.pubsub() - - -class TestConstructor: - def test_stores_config(self) -> None: - b = AsyncRedisBroker(host="h", port=1234, db=2, prefix="pfx", password="pw") - assert b._host == "h" - assert b._port == 1234 - assert b._db == 2 - assert b._prefix == "pfx" - assert b._password == "pw" - assert b._redis is None - - def test_defaults(self) -> None: - b = AsyncRedisBroker() - assert b._host == "localhost" - assert b._port == 6379 - assert b._prefix == "turnstone" - - -class TestConnect: - @pytest.mark.anyio - async def test_creates_connection(self) -> None: - b = AsyncRedisBroker() - mock_r = AsyncMock() - mock_r.pubsub = MagicMock(return_value=AsyncMock()) - with patch("redis.asyncio.Redis", return_value=mock_r): - await b.connect() - assert b._redis is mock_r - assert b._pubsub is not None - - @pytest.mark.anyio - async def test_connect_idempotent( - self, broker: AsyncRedisBroker, mock_redis: AsyncMock - ) -> None: - _inject_redis(broker, mock_redis) - old = broker._redis - await broker.connect() - assert broker._redis is old - - -class TestPushInbound: - @pytest.mark.anyio - async def test_shared_queue(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.push_inbound('{"type":"send"}') - mock_redis.rpush.assert_awaited_once_with("test:inbound", '{"type":"send"}') - - @pytest.mark.anyio - async def test_per_node_queue(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.push_inbound('{"type":"send"}', node_id="node-1") - mock_redis.rpush.assert_awaited_once_with("test:inbound:node-1", '{"type":"send"}') - - -class TestPublishOutbound: - @pytest.mark.anyio - async def test_publishes(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.publish_outbound("test:events:global", '{"event":"data"}') - mock_redis.publish.assert_awaited_once_with("test:events:global", '{"event":"data"}') - - -class TestPushResponse: - @pytest.mark.anyio - async def test_rpush_and_expire(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.push_response("req-123", '{"ok":true}') - mock_redis.rpush.assert_awaited_once_with("test:resp:req-123", '{"ok":true}') - mock_redis.expire.assert_awaited_once_with("test:resp:req-123", 120) - - -class TestSubscribe: - @pytest.mark.anyio - async def test_creates_task(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.subscribe("test:events:global", lambda msg: None) - assert "test:events:global" in broker._callbacks - assert broker._listener_task is not None - assert isinstance(broker._listener_task, asyncio.Task) - # Clean up. - broker._listener_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await broker._listener_task - - -class TestUnsubscribe: - @pytest.mark.anyio - async def test_cancels_task(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.subscribe("test:events:ch", lambda msg: None) - assert "test:events:ch" in broker._callbacks - await broker.unsubscribe("test:events:ch") - assert "test:events:ch" not in broker._callbacks - - -class TestRoutingPrimitives: - @pytest.mark.anyio - async def test_get_ws_owner(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - mock_redis.get.return_value = "node-1" - result = await broker.get_ws_owner("ws-abc") - mock_redis.get.assert_awaited_once_with("test:ws:ws-abc") - assert result == "node-1" - - @pytest.mark.anyio - async def test_set_ws_owner(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.set_ws_owner("ws-abc", "node-2") - mock_redis.set.assert_awaited_once_with("test:ws:ws-abc", "node-2") - - @pytest.mark.anyio - async def test_set_ws_owner_with_ttl( - self, broker: AsyncRedisBroker, mock_redis: AsyncMock - ) -> None: - _inject_redis(broker, mock_redis) - await broker.set_ws_owner("ws-abc", "node-2", ttl=300) - mock_redis.set.assert_awaited_once_with("test:ws:ws-abc", "node-2", ex=300) - - @pytest.mark.anyio - async def test_del_ws_owner(self, broker: AsyncRedisBroker, mock_redis: AsyncMock) -> None: - _inject_redis(broker, mock_redis) - await broker.del_ws_owner("ws-abc") - mock_redis.delete.assert_awaited_once_with("test:ws:ws-abc") - - -class TestClose: - @pytest.mark.anyio - async def test_cancels_tasks_and_closes( - self, broker: AsyncRedisBroker, mock_redis: AsyncMock - ) -> None: - _inject_redis(broker, mock_redis) - await broker.subscribe("ch1", lambda m: None) - assert len(broker._callbacks) == 1 - assert broker._listener_task is not None - await broker.close() - assert len(broker._callbacks) == 0 - assert broker._listener_task is None - assert broker._redis is None - assert broker._pubsub is None diff --git a/tests/test_auth.py b/tests/test_auth.py index 09692804..f159f6f3 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -923,7 +923,6 @@ class TestConsoleAuth: app = create_app( collector=mock_collector, - broker=MagicMock(), auth_config=AuthConfig( enabled=True, tokens={"tok_full": "full", "tok_read": "read"}, @@ -1095,7 +1094,6 @@ class TestConsoleLogin: app = create_app( collector=mock_collector, - broker=MagicMock(), auth_config=AuthConfig( enabled=True, tokens={"tok_full": "full", "tok_read": "read"}, diff --git a/tests/test_bridge_events.py b/tests/test_bridge_events.py deleted file mode 100644 index b301945a..00000000 --- a/tests/test_bridge_events.py +++ /dev/null @@ -1,120 +0,0 @@ -"""Tests for bridge event publishing — TurnCompleteEvent on idle transitions.""" - -from unittest.mock import MagicMock, patch - -from turnstone.mq.bridge import Bridge -from turnstone.mq.protocol import ContentEvent, StateChangeEvent, TurnCompleteEvent - - -def _make_bridge(): - """Create a Bridge with a mock broker (no Redis or HTTP needed).""" - broker = MagicMock() - bridge = Bridge(server_url="http://localhost:8080", broker=broker, node_id="test-node") - return bridge - - -class TestIdleTurnComplete: - """TurnCompleteEvent should be emitted on every idle transition.""" - - def test_idle_emits_turn_complete_with_correlation_id(self): - """Bridge-initiated turn: TurnCompleteEvent has the correlation_id.""" - bridge = _make_bridge() - bridge._active_sends["ws-1"] = "cid-abc" - - published = [] - with patch.object( - bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev)) - ): - bridge._handle_global_event({"type": "ws_state", "ws_id": "ws-1", "state": "idle"}) - - turn_completes = [(ws, ev) for ws, ev in published if isinstance(ev, TurnCompleteEvent)] - assert len(turn_completes) == 1 - ws, ev = turn_completes[0] - assert ws == "ws-1" - assert ev.correlation_id == "cid-abc" - # correlation_id should be removed from _active_sends - assert "ws-1" not in bridge._active_sends - - def test_idle_emits_turn_complete_without_correlation_id(self): - """Server-UI-initiated turn: TurnCompleteEvent has empty correlation_id.""" - bridge = _make_bridge() - # No entry in _active_sends for this workstream - - published = [] - with patch.object( - bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev)) - ): - bridge._handle_global_event({"type": "ws_state", "ws_id": "ws-2", "state": "idle"}) - - turn_completes = [(ws, ev) for ws, ev in published if isinstance(ev, TurnCompleteEvent)] - assert len(turn_completes) == 1 - ws, ev = turn_completes[0] - assert ws == "ws-2" - assert ev.correlation_id == "" - - def test_non_idle_state_does_not_emit_turn_complete(self): - """Non-idle state transitions should emit StateChangeEvent but not TurnCompleteEvent.""" - bridge = _make_bridge() - - published = [] - with patch.object( - bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev)) - ): - bridge._handle_global_event({"type": "ws_state", "ws_id": "ws-3", "state": "thinking"}) - - state_changes = [ev for _, ev in published if isinstance(ev, StateChangeEvent)] - turn_completes = [ev for _, ev in published if isinstance(ev, TurnCompleteEvent)] - assert len(state_changes) == 1 - assert state_changes[0].state == "thinking" - assert len(turn_completes) == 0 - - -class TestContentPassthrough: - """Bridge should pass through content from the server's idle SSE event.""" - - def test_content_passed_through_in_turn_complete(self): - """Content from idle event should be included in TurnCompleteEvent.""" - bridge = _make_bridge() - - published = [] - with patch.object( - bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev)) - ): - bridge._handle_global_event( - {"type": "ws_state", "ws_id": "ws-1", "state": "idle", "content": "Hello world"} - ) - - turn_completes = [(ws, ev) for ws, ev in published if isinstance(ev, TurnCompleteEvent)] - assert len(turn_completes) == 1 - _, ev = turn_completes[0] - assert ev.content == "Hello world" - - def test_content_empty_when_not_in_event(self): - """TurnCompleteEvent.content should be empty when idle event has no content.""" - bridge = _make_bridge() - - published = [] - with patch.object( - bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev)) - ): - bridge._handle_global_event({"type": "ws_state", "ws_id": "ws-1", "state": "idle"}) - - turn_completes = [(ws, ev) for ws, ev in published if isinstance(ev, TurnCompleteEvent)] - assert len(turn_completes) == 1 - _, ev = turn_completes[0] - assert ev.content == "" - - def test_content_event_still_published(self): - """Content events should still be published to per-ws channel.""" - bridge = _make_bridge() - - published = [] - with patch.object( - bridge, "_publish_ws", side_effect=lambda ws, ev: published.append((ws, ev)) - ): - bridge._handle_ws_event("ws-1", {"type": "content", "text": "hello"}) - - content_events = [(ws, ev) for ws, ev in published if isinstance(ev, ContentEvent)] - assert len(content_events) == 1 - _, ev = content_events[0] - assert ev.text == "hello" diff --git a/tests/test_bridge_stress.py b/tests/test_bridge_stress.py deleted file mode 100644 index aef920a5..00000000 --- a/tests/test_bridge_stress.py +++ /dev/null @@ -1,357 +0,0 @@ -"""Stress tests for bridge.py threading — race conditions in approval, -plan review, and workstream lifecycle. - -Each scenario is run many times (ITERATIONS) with threading.Barrier to -maximize timing overlap. Uses mock broker (no Redis) and no HTTP calls. - -Races tested: -1. Duplicate approval on SSE reconnect (TOCTOU in _pending_approvals) -2. Duplicate plan review on SSE reconnect (TOCTOU in _pending_plan_reviews) -3. approve_set stale reference escape during concurrent update -4. _running flag visibility across threads on shutdown -5. Approval thread exits within bounded time after timeout -6. Concurrent approval + workstream close leaves no orphaned state -""" - -from __future__ import annotations - -import threading -import time -from collections import Counter -from unittest.mock import MagicMock, patch - -from turnstone.mq.bridge import Bridge - -ITERATIONS = 100 - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - - -def _make_bridge(**overrides) -> Bridge: - """Create a Bridge with a mock broker (no Redis or HTTP).""" - broker = MagicMock() - defaults = dict( - server_url="http://localhost:8080", - broker=broker, - node_id="test-node", - approval_timeout=1, - ) - defaults.update(overrides) - bridge = Bridge(**defaults) - # Replace real httpx client with a mock so daemon threads spawned by - # _handle_approval / _handle_plan_review don't make real HTTP calls - # after the test's patch context exits. - bridge._http.close() - bridge._http = MagicMock() - return bridge - - -def _approval_items(tool_name: str = "bash") -> list[dict]: - return [{"func_name": tool_name, "needs_approval": True, "approval_label": tool_name}] - - -def _wait_pending_resolved(bridge: Bridge, key: str, attr: str, deadline_s: float = 3.0) -> bool: - """Poll until the pending entry is resolved (tombstone) or absent.""" - deadline = time.monotonic() + deadline_s - while time.monotonic() < deadline: - with bridge._lock: - entries = getattr(bridge, attr) - if key not in entries: - return True - _, resolved_at = entries[key] - if resolved_at > 0: - return True - time.sleep(0.01) - return False - - -# --------------------------------------------------------------------------- -# Race 1: Duplicate approval on SSE reconnect -# --------------------------------------------------------------------------- - - -class TestDuplicateApproval: - """Two threads call _handle_approval for the same ws_id simultaneously. - Only one should create a pending entry; the other should be skipped.""" - - def test_no_duplicate_approvals(self): - sent_count = Counter() - - for _ in range(ITERATIONS): - bridge = _make_bridge() - bridge._broker.pop_response.return_value = '{"type": "approve", "approved": true}' - barrier = threading.Barrier(2, timeout=5) - - def _call_approval(bridge=bridge, barrier=barrier): - barrier.wait() - bridge._handle_approval("ws-1", {"items": _approval_items()}) - - t1 = threading.Thread(target=_call_approval) - t2 = threading.Thread(target=_call_approval) - with ( - patch.object(bridge, "_api_approve") as mock_approve, - patch.object(bridge, "_publish_ws"), - ): - t1.start() - t2.start() - t1.join(timeout=5) - t2.join(timeout=5) - assert not t1.is_alive(), "Thread 1 hung" - assert not t2.is_alive(), "Thread 2 hung" - - # Wait for spawned _wait_approval threads to resolve - _wait_pending_resolved(bridge, "ws-1", "_pending_approvals") - - sent_count[mock_approve.call_count] += 1 - - # At most 1 approval should be forwarded per iteration - assert sent_count.get(2, 0) == 0, ( - f"Duplicate approvals sent in {sent_count[2]}/{ITERATIONS} iterations" - ) - - -# --------------------------------------------------------------------------- -# Race 2: Duplicate plan review on SSE reconnect -# --------------------------------------------------------------------------- - - -class TestDuplicatePlanReview: - """Two threads call _handle_plan_review simultaneously. - Only one should create a pending entry.""" - - def test_no_duplicate_plan_reviews(self): - sent_count = Counter() - - for _ in range(ITERATIONS): - bridge = _make_bridge() - bridge._broker.pop_response.return_value = ( - '{"type": "plan_feedback", "feedback": "looks good"}' - ) - barrier = threading.Barrier(2, timeout=5) - - def _call_plan(bridge=bridge, barrier=barrier): - barrier.wait() - bridge._handle_plan_review("ws-1", {"content": "plan text"}) - - t1 = threading.Thread(target=_call_plan) - t2 = threading.Thread(target=_call_plan) - with patch.object(bridge, "_publish_ws"), patch.object(bridge._http, "post"): - t1.start() - t2.start() - t1.join(timeout=5) - t2.join(timeout=5) - assert not t1.is_alive(), "Thread 1 hung" - assert not t2.is_alive(), "Thread 2 hung" - - # Wait for spawned _wait_plan threads to resolve - _wait_pending_resolved(bridge, "ws-1", "_pending_plan_reviews") - - sent_count[bridge._http.post.call_count] += 1 - - assert sent_count.get(2, 0) == 0, ( - f"Duplicate plan reviews sent in {sent_count[2]}/{ITERATIONS} iterations" - ) - - -# --------------------------------------------------------------------------- -# Race 3: approve_set stale reference during concurrent update -# --------------------------------------------------------------------------- - - -class TestApproveSetConsistency: - """One thread reads approve_set for auto-approve check while another - updates it via _wait_approval 'always' path. The auto-approve - decision should be consistent (either all-approved or not).""" - - def test_approve_set_never_partially_visible(self): - for _ in range(ITERATIONS): - bridge = _make_bridge() - with bridge._lock: - bridge._ws_approve_tools["ws-1"] = {"read_file", "search"} - - barrier = threading.Barrier(2, timeout=5) - results = [] - - def _reader(bridge=bridge, barrier=barrier, results=results): - barrier.wait() - with bridge._lock: - snap = bridge._ws_approve_tools.get("ws-1", set()).copy() - results.append(snap) - - def _writer(bridge=bridge, barrier=barrier): - barrier.wait() - with bridge._lock: - existing = bridge._ws_approve_tools.get("ws-1", set()) - bridge._ws_approve_tools["ws-1"] = existing | {"bash", "write_file"} - - t1 = threading.Thread(target=_reader) - t2 = threading.Thread(target=_writer) - t1.start() - t2.start() - t1.join(timeout=5) - t2.join(timeout=5) - assert not t1.is_alive(), "Reader hung" - assert not t2.is_alive(), "Writer hung" - - snap = results[0] - assert snap in ( - {"read_file", "search"}, - {"read_file", "search", "bash", "write_file"}, - ), f"Partial set observed: {snap}" - - -# --------------------------------------------------------------------------- -# Race 4: _running flag visibility across threads -# --------------------------------------------------------------------------- - - -class TestRunningFlagVisibility: - """All threads reading _running should see False within a bounded time - after the main thread sets it.""" - - def test_all_threads_observe_shutdown(self): - bridge = _make_bridge() - observed_false = threading.Event() - threads_running = [] - - def _spin_checker(): - while bridge._running: - time.sleep(0.001) - observed_false.set() - - for _ in range(5): - t = threading.Thread(target=_spin_checker, daemon=True) - threads_running.append(t) - t.start() - - time.sleep(0.01) - bridge._running = False - - for t in threads_running: - t.join(timeout=1) - assert not t.is_alive(), "Thread did not observe _running=False" - - assert observed_false.is_set() - - -# --------------------------------------------------------------------------- -# Race 5: Approval thread exits within bounded time -# --------------------------------------------------------------------------- - - -class TestApprovalThreadTimeout: - """An approval thread blocked on pop_response should exit within the - configured approval_timeout, not hang indefinitely.""" - - def test_approval_thread_exits_within_timeout(self): - for _ in range(10): - bridge = _make_bridge(approval_timeout=0.5) - - def _slow_pop(queue_name, timeout=300): - time.sleep(min(timeout, 0.5)) - return None - - bridge._broker.pop_response.side_effect = _slow_pop - - with patch.object(bridge, "_publish_ws"), patch.object(bridge, "_api_approve"): - bridge._handle_approval("ws-1", {"items": _approval_items()}) - - # The pending entry should be resolved within the timeout - resolved = _wait_pending_resolved(bridge, "ws-1", "_pending_approvals", deadline_s=3.0) - assert resolved, "Approval thread did not exit within expected timeout" - - -# --------------------------------------------------------------------------- -# Race 6: Concurrent approval + workstream close -# --------------------------------------------------------------------------- - - -class TestApprovalDuringClose: - """An approval arriving at the exact same time as a ws_closed event - should not leave orphaned state.""" - - def test_no_orphaned_pending_after_close(self): - for _ in range(ITERATIONS): - bridge = _make_bridge(approval_timeout=0.1) - bridge._broker.pop_response.return_value = None # timeout - - barrier = threading.Barrier(2, timeout=5) - - def _send_approval(bridge=bridge, barrier=barrier): - barrier.wait() - with patch.object(bridge, "_publish_ws"), patch.object(bridge, "_api_approve"): - bridge._handle_approval("ws-1", {"items": _approval_items()}) - - def _close_ws(bridge=bridge, barrier=barrier): - barrier.wait() - with ( - patch.object(bridge, "_publish_global"), - patch.object(bridge, "_publish_cluster"), - ): - bridge._handle_global_event({"type": "ws_closed", "ws_id": "ws-1"}) - - t1 = threading.Thread(target=_send_approval) - t2 = threading.Thread(target=_close_ws) - t1.start() - t2.start() - t1.join(timeout=5) - t2.join(timeout=5) - assert not t1.is_alive(), "Approval thread hung" - assert not t2.is_alive(), "Close thread hung" - - # Wait for spawned _wait_approval thread to resolve (if close - # didn't remove the entry first) - resolved = _wait_pending_resolved(bridge, "ws-1", "_pending_approvals") - assert resolved, "Orphaned pending approval" - - -# --------------------------------------------------------------------------- -# Race 7: Plan review refinement loop (tombstone → cleanup → re-entry) -# --------------------------------------------------------------------------- - - -class TestPlanReviewRefinementLoop: - """After a plan review is resolved, a ws_state event should clean up the - tombstone so the refinement-loop plan_review event is handled correctly.""" - - def test_refinement_loop_allows_reentry(self): - for _ in range(ITERATIONS): - bridge = _make_bridge() - bridge._broker.pop_response.return_value = ( - '{"type": "plan_feedback", "feedback": "refine this"}' - ) - - # Step 1: first plan review — creates pending entry, resolves it - with patch.object(bridge, "_publish_ws"), patch.object(bridge._http, "post"): - bridge._handle_plan_review("ws-1", {"content": "plan v1"}) - - _wait_pending_resolved(bridge, "ws-1", "_pending_plan_reviews") - - # Verify tombstone is present (resolved_at > 0) - with bridge._lock: - assert "ws-1" in bridge._pending_plan_reviews - assert bridge._pending_plan_reviews["ws-1"][1] > 0 - - # Step 2: ws_state event cleans up the resolved tombstone - with ( - patch.object(bridge, "_publish_ws"), - patch.object(bridge, "_publish_global"), - patch.object(bridge, "_publish_cluster"), - ): - bridge._handle_global_event( - {"type": "ws_state", "ws_id": "ws-1", "state": "working"} - ) - - with bridge._lock: - assert "ws-1" not in bridge._pending_plan_reviews - - # Step 3: refinement plan_review arrives — should create new entry - with patch.object(bridge, "_publish_ws"), patch.object(bridge._http, "post"): - bridge._handle_plan_review("ws-1", {"content": "plan v2"}) - - _wait_pending_resolved(bridge, "ws-1", "_pending_plan_reviews") - - with bridge._lock: - assert "ws-1" in bridge._pending_plan_reviews diff --git a/tests/test_channel_discord.py b/tests/test_channel_discord.py index 896fe733..e761ca84 100644 --- a/tests/test_channel_discord.py +++ b/tests/test_channel_discord.py @@ -76,8 +76,7 @@ class TestDiscordConfig: assert cfg.max_message_length == 2000 assert cfg.streaming_edit_interval == 1.5 # Inherited from ChannelConfig - assert cfg.redis_host == "localhost" - assert cfg.redis_port == 6379 + assert cfg.server_url == "http://localhost:8080" assert cfg.model == "" assert cfg.auto_approve is False @@ -316,12 +315,12 @@ class TestParseFooter: class TestWsEventFinalization: - """TurnCompleteEvent should finalize streaming messages in the Discord bot.""" + """StreamEndEvent should finalize streaming messages in the Discord bot.""" - def test_turn_complete_finalizes_streaming(self): - """ContentEvent + TurnCompleteEvent(correlation_id='') finalizes the message.""" + def test_stream_end_finalizes_streaming(self): + """ContentEvent + StreamEndEvent finalizes the message.""" from turnstone.channels.discord.bot import TurnstoneBot - from turnstone.mq.protocol import ContentEvent, TurnCompleteEvent + from turnstone.sdk.events import ContentEvent, StreamEndEvent bot = MagicMock(spec=TurnstoneBot) bot.config = MagicMock() @@ -339,23 +338,23 @@ class TestWsEventFinalization: thread = AsyncMock() # Feed content event - content_raw = ContentEvent(ws_id="ws-1", text="Hello world").to_json() - _run(bot._on_ws_event("ws-1", thread, content_raw)) + content_event = ContentEvent(ws_id="ws-1", text="Hello world") + _run(bot._on_ws_event("ws-1", thread, content_event)) # StreamingMessage should exist assert "ws-1" in bot._streaming - # Feed turn complete with empty correlation_id (server-UI-initiated) - complete_raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="").to_json() - _run(bot._on_ws_event("ws-1", thread, complete_raw)) + # Feed stream end + end_event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, end_event)) # StreamingMessage should be removed and finalized assert "ws-1" not in bot._streaming - def test_turn_complete_no_streaming_is_noop(self): - """TurnCompleteEvent without prior content should not error.""" + def test_stream_end_no_streaming_is_noop(self): + """StreamEndEvent without prior content should not error.""" from turnstone.channels.discord.bot import TurnstoneBot - from turnstone.mq.protocol import TurnCompleteEvent + from turnstone.sdk.events import StreamEndEvent bot = MagicMock(spec=TurnstoneBot) bot._streaming = {} @@ -365,8 +364,8 @@ class TestWsEventFinalization: thread = AsyncMock() - complete_raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="").to_json() - _run(bot._on_ws_event("ws-1", thread, complete_raw)) + end_event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, end_event)) # No error, no streaming message assert "ws-1" not in bot._streaming @@ -399,8 +398,8 @@ class TestApprovalVerdictDisplay: return bot def test_approval_with_heuristic_verdict(self): - """ApprovalRequestEvent items with verdict dicts add embed fields.""" - from turnstone.mq.protocol import ApprovalRequestEvent + """ApproveRequestEvent items with verdict dicts add embed fields.""" + from turnstone.sdk.events import ApproveRequestEvent bot = self._make_bot() thread = AsyncMock() @@ -421,8 +420,8 @@ class TestApprovalVerdictDisplay: }, } ] - raw = ApprovalRequestEvent(ws_id="ws-1", correlation_id="corr-1", items=items).to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + event = ApproveRequestEvent(ws_id="ws-1", items=items) + _run(bot._on_ws_event("ws-1", thread, event)) # thread.send was called with an embed containing a verdict field thread.send.assert_awaited_once() @@ -439,8 +438,8 @@ class TestApprovalVerdictDisplay: assert "ws-1" in bot._pending_approval_msgs def test_approval_without_verdict(self): - """ApprovalRequestEvent items without verdict still work normally.""" - from turnstone.mq.protocol import ApprovalRequestEvent + """ApproveRequestEvent items without verdict still work normally.""" + from turnstone.sdk.events import ApproveRequestEvent bot = self._make_bot() thread = AsyncMock() @@ -448,8 +447,8 @@ class TestApprovalVerdictDisplay: thread.send = AsyncMock(return_value=sent_msg) items = [{"func_name": "read_file", "preview": "/etc/hosts", "needs_approval": True}] - raw = ApprovalRequestEvent(ws_id="ws-1", correlation_id="corr-1", items=items).to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + event = ApproveRequestEvent(ws_id="ws-1", items=items) + _run(bot._on_ws_event("ws-1", thread, event)) thread.send.assert_awaited_once() call_kwargs = thread.send.call_args[1] @@ -459,7 +458,7 @@ class TestApprovalVerdictDisplay: def test_intent_verdict_event_updates_embed(self): """IntentVerdictEvent should update the pending approval embed.""" - from turnstone.mq.protocol import IntentVerdictEvent + from turnstone.sdk.events import IntentVerdictEvent bot = self._make_bot() thread = AsyncMock() @@ -471,7 +470,7 @@ class TestApprovalVerdictDisplay: msg.edit = AsyncMock() bot._pending_approval_msgs["ws-1"] = msg - raw = IntentVerdictEvent( + event = IntentVerdictEvent( ws_id="ws-1", func_name="bash", risk_level="high", @@ -479,8 +478,8 @@ class TestApprovalVerdictDisplay: confidence=0.9, intent_summary="Dangerous operation", tier="llm", - ).to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + ) + _run(bot._on_ws_event("ws-1", thread, event)) # Embed should be updated with the judge verdict field embed.add_field.assert_called_once() @@ -494,19 +493,19 @@ class TestApprovalVerdictDisplay: def test_intent_verdict_without_pending_approval_is_noop(self): """IntentVerdictEvent without a pending approval message should not error.""" - from turnstone.mq.protocol import IntentVerdictEvent + from turnstone.sdk.events import IntentVerdictEvent bot = self._make_bot() thread = AsyncMock() - raw = IntentVerdictEvent(ws_id="ws-1", func_name="bash", risk_level="low").to_json() + event = IntentVerdictEvent(ws_id="ws-1", func_name="bash", risk_level="low") # Should not raise - _run(bot._on_ws_event("ws-1", thread, raw)) + _run(bot._on_ws_event("ws-1", thread, event)) - def test_turn_complete_clears_pending_approval(self): - """TurnCompleteEvent should clean up the pending approval message tracking.""" + def test_stream_end_clears_pending_approval(self): + """StreamEndEvent should clean up the pending approval message tracking.""" from turnstone.channels.discord.bot import TurnstoneBot - from turnstone.mq.protocol import TurnCompleteEvent + from turnstone.sdk.events import StreamEndEvent bot = MagicMock(spec=TurnstoneBot) bot._streaming = {} @@ -515,14 +514,14 @@ class TestApprovalVerdictDisplay: bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot) thread = AsyncMock() - raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="").to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, event)) assert "ws-1" not in bot._pending_approval_msgs -class TestContentCatchup: - """TurnCompleteEvent with content field provides catch-up for missed ContentEvents.""" +class TestStreamEndBehavior: + """StreamEndEvent finalizes streaming and cleans up state.""" def _make_bot(self): from turnstone.channels.discord.bot import TurnstoneBot @@ -539,51 +538,35 @@ class TestContentCatchup: bot._on_ws_event = TurnstoneBot._on_ws_event.__get__(bot, TurnstoneBot) return bot - def test_catchup_sends_content_when_no_streaming(self): - """TurnCompleteEvent with content but no SM sends catch-up message.""" - from turnstone.mq.protocol import TurnCompleteEvent + def test_stream_end_no_streaming_no_send(self): + """StreamEndEvent without prior content should not send anything.""" + from turnstone.sdk.events import StreamEndEvent bot = self._make_bot() thread = AsyncMock() - raw = TurnCompleteEvent( - ws_id="ws-1", correlation_id="", content="Caught up response" - ).to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, event)) - thread.send.assert_awaited_once_with("Caught up response") + thread.send.assert_not_awaited() - def test_catchup_skipped_when_streaming_exists(self): - """TurnCompleteEvent with content and existing SM uses SM finalize, not catch-up.""" - from turnstone.mq.protocol import ContentEvent, TurnCompleteEvent + def test_stream_end_finalizes_existing_streaming(self): + """StreamEndEvent with an existing StreamingMessage should finalize it.""" + from turnstone.sdk.events import ContentEvent, StreamEndEvent bot = self._make_bot() thread = AsyncMock() # Feed content event to create SM - content_raw = ContentEvent(ws_id="ws-1", text="Streamed").to_json() - _run(bot._on_ws_event("ws-1", thread, content_raw)) + content_event = ContentEvent(ws_id="ws-1", text="Streamed") + _run(bot._on_ws_event("ws-1", thread, content_event)) assert "ws-1" in bot._streaming - # Now TurnCompleteEvent with content — SM should be finalized, not catch-up - complete_raw = TurnCompleteEvent( - ws_id="ws-1", correlation_id="", content="Streamed" - ).to_json() - _run(bot._on_ws_event("ws-1", thread, complete_raw)) + # Now StreamEndEvent — SM should be finalized + end_event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, end_event)) assert "ws-1" not in bot._streaming - def test_catchup_empty_content_no_message(self): - """TurnCompleteEvent with empty content and no SM sends nothing.""" - from turnstone.mq.protocol import TurnCompleteEvent - - bot = self._make_bot() - thread = AsyncMock() - - raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="", content="").to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) - - thread.send.assert_not_awaited() - class TestNotificationTracking: """Tests for notification message tracking and DM reply routing.""" @@ -767,14 +750,15 @@ class TestNotificationTracking: ts.router.send_message.assert_not_awaited() - def test_turn_complete_forwards_to_dm(self): - """TurnCompleteEvent should forward content to notification reply DM.""" + def test_stream_end_forwards_accumulated_content_to_dm(self): + """StreamEndEvent should forward accumulated content to notification reply DM.""" from turnstone.channels.discord.bot import TurnstoneBot - from turnstone.mq.protocol import TurnCompleteEvent + from turnstone.sdk.events import ContentEvent, StreamEndEvent bot = MagicMock(spec=TurnstoneBot) bot.config = MagicMock() bot.config.max_message_length = 2000 + bot.config.streaming_edit_interval = 1.5 bot._streaming = {} bot._pending_approval_msgs = {} bot._notify_ws_map = {} @@ -790,10 +774,13 @@ class TestNotificationTracking: thread = AsyncMock() - raw = TurnCompleteEvent( - ws_id="ws-1", correlation_id="", content="Here's the response" - ).to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + # Feed content events to accumulate buffer + content_event = ContentEvent(ws_id="ws-1", text="Here's the response") + _run(bot._on_ws_event("ws-1", thread, content_event)) + + # Feed stream end — should finalize and forward to DM + end_event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, end_event)) # Should send to DM channel dm_channel.send.assert_awaited_once_with("Here's the response") @@ -803,10 +790,10 @@ class TestNotificationTracking: assert 88888 in bot._notify_ws_map assert bot._notify_ws_map[88888] == ("ws-1", "u123") - def test_turn_complete_cleans_up_dm_even_without_content(self): - """TurnCompleteEvent without content should still clean up DM tracking.""" + def test_stream_end_cleans_up_dm_even_without_content(self): + """StreamEndEvent without prior content should still clean up DM tracking.""" from turnstone.channels.discord.bot import TurnstoneBot - from turnstone.mq.protocol import TurnCompleteEvent + from turnstone.sdk.events import StreamEndEvent bot = MagicMock(spec=TurnstoneBot) bot._streaming = {} @@ -819,8 +806,8 @@ class TestNotificationTracking: thread = AsyncMock() - raw = TurnCompleteEvent(ws_id="ws-1", correlation_id="", content="").to_json() - _run(bot._on_ws_event("ws-1", thread, raw)) + end_event = StreamEndEvent(ws_id="ws-1") + _run(bot._on_ws_event("ws-1", thread, end_event)) # DM should not be sent to (no content) dm_channel.send.assert_not_awaited() diff --git a/tests/test_channel_routing.py b/tests/test_channel_routing.py index 76629c5a..dcfe694d 100644 --- a/tests/test_channel_routing.py +++ b/tests/test_channel_routing.py @@ -2,26 +2,15 @@ from __future__ import annotations -import json +from typing import Any from unittest.mock import AsyncMock, MagicMock +import httpx import pytest from turnstone.channels._routing import ChannelRouter -@pytest.fixture -def mock_broker() -> AsyncMock: - """Return a mock AsyncRedisBroker.""" - broker = AsyncMock() - broker._prefix = "test" - broker.push_inbound = AsyncMock() - broker.push_response = AsyncMock() - broker.subscribe = AsyncMock() - broker.unsubscribe = AsyncMock() - return broker - - @pytest.fixture def mock_storage() -> MagicMock: """Return a mock StorageBackend.""" @@ -35,8 +24,24 @@ def mock_storage() -> MagicMock: @pytest.fixture -def router(mock_broker: AsyncMock, mock_storage: MagicMock) -> ChannelRouter: - return ChannelRouter(broker=mock_broker, storage=mock_storage) +def router(mock_storage: MagicMock) -> ChannelRouter: + return ChannelRouter( + server_url="http://localhost:8080/v1", + storage=mock_storage, + ) + + +def _ok_response(json_data: object = None) -> httpx.Response: + """Build a mock 200 response with optional JSON body.""" + import json + + content = json.dumps(json_data or {"status": "ok"}).encode() + return httpx.Response( + 200, + content=content, + headers={"content-type": "application/json"}, + request=httpx.Request("POST", "http://test"), + ) class TestResolveUser: @@ -56,46 +61,53 @@ class TestResolveUser: class TestSendMessage: @pytest.mark.anyio - async def test_pushes_send_message(self, router: ChannelRouter, mock_broker: AsyncMock) -> None: - cid = await router.send_message("ws-1", "hello world") - assert isinstance(cid, str) - assert len(cid) > 0 - mock_broker.push_inbound.assert_awaited_once() - raw = mock_broker.push_inbound.call_args[0][0] - payload = json.loads(raw) - assert payload["type"] == "send" - assert payload["ws_id"] == "ws-1" - assert payload["message"] == "hello world" - assert payload["correlation_id"] == cid + async def test_posts_to_server( + self, router: ChannelRouter, monkeypatch: pytest.MonkeyPatch + ) -> None: + mock_post = AsyncMock(return_value=_ok_response()) + monkeypatch.setattr(router, "_post", mock_post) + await router.send_message("ws-1", "hello world") + mock_post.assert_awaited_once_with("/api/send", {"ws_id": "ws-1", "message": "hello world"}) class TestSendApproval: @pytest.mark.anyio - async def test_pushes_to_response_queue( - self, router: ChannelRouter, mock_broker: AsyncMock + async def test_posts_to_server( + self, router: ChannelRouter, monkeypatch: pytest.MonkeyPatch ) -> None: + mock_post = AsyncMock(return_value=_ok_response()) + monkeypatch.setattr(router, "_post", mock_post) await router.send_approval("ws-1", "corr-abc", approved=True, feedback="ok") - mock_broker.push_response.assert_awaited_once() - queue_name = mock_broker.push_response.call_args[0][0] - assert queue_name == "corr-abc" - raw = mock_broker.push_response.call_args[0][1] - payload = json.loads(raw) - assert payload["type"] == "approve" - assert payload["approved"] is True - assert payload["ws_id"] == "ws-1" + mock_post.assert_awaited_once_with( + "/api/approve", + {"ws_id": "ws-1", "approved": True, "always": False, "feedback": "ok"}, + ) + + @pytest.mark.anyio + async def test_omits_empty_feedback( + self, router: ChannelRouter, monkeypatch: pytest.MonkeyPatch + ) -> None: + mock_post = AsyncMock(return_value=_ok_response()) + monkeypatch.setattr(router, "_post", mock_post) + await router.send_approval("ws-1", "corr-abc", approved=False) + mock_post.assert_awaited_once_with( + "/api/approve", + {"ws_id": "ws-1", "approved": False, "always": False}, + ) class TestSendPlanFeedback: @pytest.mark.anyio - async def test_pushes_to_response_queue( - self, router: ChannelRouter, mock_broker: AsyncMock + async def test_posts_to_server( + self, router: ChannelRouter, monkeypatch: pytest.MonkeyPatch ) -> None: + mock_post = AsyncMock(return_value=_ok_response()) + monkeypatch.setattr(router, "_post", mock_post) await router.send_plan_feedback("ws-2", "corr-xyz", "looks good") - mock_broker.push_response.assert_awaited_once() - raw = mock_broker.push_response.call_args[0][1] - payload = json.loads(raw) - assert payload["type"] == "plan_feedback" - assert payload["feedback"] == "looks good" + mock_post.assert_awaited_once_with( + "/api/plan", + {"ws_id": "ws-2", "feedback": "looks good"}, + ) class TestDeleteRoute: @@ -105,3 +117,96 @@ class TestDeleteRoute: ) -> None: await router.delete_route("discord", "ch-123") mock_storage.delete_channel_route.assert_called_once_with("discord", "ch-123") + + +class TestGetOrCreateWorkstream: + @pytest.mark.anyio + async def test_creates_new_workstream( + self, + router: ChannelRouter, + mock_storage: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + mock_post = AsyncMock( + return_value=_ok_response({"ws_id": "ws-new", "name": "test", "resumed": False}), + ) + monkeypatch.setattr(router, "_post", mock_post) + ws_id, is_new = await router.get_or_create_workstream("discord", "ch-1", name="test") + assert ws_id == "ws-new" + assert is_new is True + mock_storage.create_channel_route.assert_called_once_with("discord", "ch-1", "ws-new") + + @pytest.mark.anyio + async def test_returns_existing_alive_workstream( + self, + router: ChannelRouter, + mock_storage: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + mock_storage.get_channel_route.return_value = { + "ws_id": "ws-old", + "channel_type": "discord", + "channel_id": "ch-1", + } + monkeypatch.setattr(router, "_is_ws_alive", AsyncMock(return_value=True)) + ws_id, is_new = await router.get_or_create_workstream("discord", "ch-1") + assert ws_id == "ws-old" + assert is_new is False + + @pytest.mark.anyio + async def test_resumes_stale_workstream( + self, + router: ChannelRouter, + mock_storage: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + mock_storage.get_channel_route.return_value = { + "ws_id": "ws-stale", + "channel_type": "discord", + "channel_id": "ch-1", + } + # Alive check returns False — ws is not alive. + monkeypatch.setattr(router, "_is_ws_alive", AsyncMock(return_value=False)) + # POST to create returns a new ws_id. + create_resp = _ok_response({"ws_id": "ws-resumed", "name": "test", "resumed": True}) + captured: list[dict[str, Any]] = [] + + async def _fake_post(path: str, body: dict[str, Any]) -> httpx.Response: + captured.append({"path": path, "body": body}) + return create_resp + + monkeypatch.setattr(router, "_post", _fake_post) + ws_id, is_new = await router.get_or_create_workstream("discord", "ch-1", name="test") + assert ws_id == "ws-resumed" + assert is_new is True + # Should have deleted the stale route and created a new one. + mock_storage.delete_channel_route.assert_called_once_with("discord", "ch-1") + mock_storage.create_channel_route.assert_called_once_with("discord", "ch-1", "ws-resumed") + # The create body should include resume_ws pointing at the old ws. + create_call = captured[0] + assert create_call["body"]["resume_ws"] == "ws-stale" + + +class TestCloseWorkstream: + @pytest.mark.anyio + async def test_posts_to_server( + self, router: ChannelRouter, monkeypatch: pytest.MonkeyPatch + ) -> None: + mock_post = AsyncMock(return_value=_ok_response()) + monkeypatch.setattr(router, "_post", mock_post) + await router.close_workstream("ws-1") + mock_post.assert_awaited_once_with( + "/api/workstreams/close", + {"ws_id": "ws-1"}, + ) + + +class TestAclose: + @pytest.mark.anyio + async def test_closes_client( + self, router: ChannelRouter, monkeypatch: pytest.MonkeyPatch + ) -> None: + mock_close = AsyncMock() + monkeypatch.setattr(router._client, "aclose", mock_close) + await router.aclose() + mock_close.assert_awaited_once() diff --git a/tests/test_config.py b/tests/test_config.py index ce0719d3..8b6893a8 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -65,61 +65,56 @@ def test_apply_config_sets_defaults(tmp_path): _reset_cache() cfg = tmp_path / "config.toml" cfg.write_text( - '[redis]\nhost = "redis.local"\nport = 7777\npassword = "pw"\n' - '[bridge]\nserver_url = "http://bridge:9090"\n' + '[server]\nhost = "0.0.0.0"\nport = 9090\n[api]\nbase_url = "http://custom/v1"\n' ) set_config_path(str(cfg)) parser = argparse.ArgumentParser() - parser.add_argument("--redis-host", default="localhost") - parser.add_argument("--redis-port", type=int, default=6379) - parser.add_argument("--redis-password", default=None) - parser.add_argument("--server-url", default="http://localhost:8080") + parser.add_argument("--host", default="localhost") + parser.add_argument("--port", type=int, default=8080) + parser.add_argument("--base-url", default="http://localhost:11434/v1") - apply_config(parser, ["redis", "bridge"]) + apply_config(parser, ["server", "api"]) args = parser.parse_args([]) - assert args.redis_host == "redis.local" - assert args.redis_port == 7777 - assert args.redis_password == "pw" - assert args.server_url == "http://bridge:9090" + assert args.host == "0.0.0.0" + assert args.port == 9090 + assert args.base_url == "http://custom/v1" def test_apply_config_cli_overrides(tmp_path): _reset_cache() cfg = tmp_path / "config.toml" - cfg.write_text('[redis]\nhost = "config-host"\nport = 7777\n') + cfg.write_text('[server]\nhost = "config-host"\nport = 7777\n') set_config_path(str(cfg)) parser = argparse.ArgumentParser() - parser.add_argument("--redis-host", default="localhost") - parser.add_argument("--redis-port", type=int, default=6379) + parser.add_argument("--host", default="localhost") + parser.add_argument("--port", type=int, default=8080) - apply_config(parser, ["redis"]) + apply_config(parser, ["server"]) # CLI flag overrides config - args = parser.parse_args(["--redis-host", "cli-host"]) + args = parser.parse_args(["--host", "cli-host"]) - assert args.redis_host == "cli-host" # CLI wins - assert args.redis_port == 7777 # config wins (no CLI override) + assert args.host == "cli-host" # CLI wins + assert args.port == 7777 # config wins (no CLI override) def test_apply_config_missing_keys_keep_defaults(tmp_path): _reset_cache() cfg = tmp_path / "config.toml" - cfg.write_text('[redis]\nhost = "only-host"\n') # no port, no password + cfg.write_text('[server]\nhost = "only-host"\n') # no port set_config_path(str(cfg)) parser = argparse.ArgumentParser() - parser.add_argument("--redis-host", default="localhost") - parser.add_argument("--redis-port", type=int, default=6379) - parser.add_argument("--redis-password", default=None) + parser.add_argument("--host", default="localhost") + parser.add_argument("--port", type=int, default=8080) - apply_config(parser, ["redis"]) + apply_config(parser, ["server"]) args = parser.parse_args([]) - assert args.redis_host == "only-host" - assert args.redis_port == 6379 # original default kept - assert args.redis_password is None # original default kept + assert args.host == "only-host" + assert args.port == 8080 # original default kept def test_apply_config_no_file(tmp_path): @@ -127,11 +122,11 @@ def test_apply_config_no_file(tmp_path): set_config_path(str(tmp_path / "nope.toml")) parser = argparse.ArgumentParser() - parser.add_argument("--redis-host", default="localhost") + parser.add_argument("--host", default="localhost") - apply_config(parser, ["redis"]) + apply_config(parser, ["server"]) args = parser.parse_args([]) - assert args.redis_host == "localhost" + assert args.host == "localhost" def test_apply_config_model_section(tmp_path): diff --git a/tests/test_console.py b/tests/test_console.py index e2474b4d..ed114bd1 100644 --- a/tests/test_console.py +++ b/tests/test_console.py @@ -8,38 +8,20 @@ from unittest.mock import MagicMock, patch import pytest from turnstone.console.collector import ClusterCollector, NodeSnapshot -from turnstone.mq.protocol import ( - ClusterStateEvent, -) # --------------------------------------------------------------------------- -# Mock broker for collector tests +# Mock storage for collector tests # --------------------------------------------------------------------------- -class MockBroker: - """Minimal broker mock that records calls and stores nodes.""" +class MockStorage: + """Minimal storage mock that implements list_services for collector tests.""" def __init__(self): - self.nodes: list[dict] = [] - self._subscriptions: dict[str, list] = {} + self.services: list[dict[str, str]] = [] - def list_nodes(self) -> list[dict]: - return list(self.nodes) - - def subscribe_outbound(self, channel, callback): - self._subscriptions.setdefault(channel, []).append(callback) - - def publish_outbound(self, channel, event): - for cb in self._subscriptions.get(channel, []): - cb(event) - - def subscribe_cluster(self, callback): - channel = "turnstone:events:cluster" - self.subscribe_outbound(channel, callback) - - def close(self): - pass + def list_services(self, service_type: str, max_age_seconds: int = 120) -> list[dict[str, str]]: + return [s for s in self.services if True] # all services match # --------------------------------------------------------------------------- @@ -47,11 +29,11 @@ class MockBroker: # --------------------------------------------------------------------------- -def _make_collector(broker=None, poll_interval=0, discovery_interval=999): +def _make_collector(storage=None, poll_interval=0, discovery_interval=999): """Create a collector with zero poll interval (no jitter delay in tests).""" - b = broker or MockBroker() + s = storage or MockStorage() return ClusterCollector( - broker=b, + storage=s, poll_interval=poll_interval, discovery_interval=discovery_interval, ) @@ -79,52 +61,59 @@ def _dashboard_response(workstreams=None, aggregate=None): class TestCollectorDiscovery: - """Node discovery from heartbeat keys.""" + """Node discovery from service registry.""" def test_discover_new_nodes(self): - broker = MockBroker() - broker.nodes = [ - {"node_id": "node-a", "server_url": "http://a:8080"}, - {"node_id": "node-b", "server_url": "http://b:8080"}, + storage = MockStorage() + storage.services = [ + {"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"}, + {"service_id": "node-b", "url": "http://b:8080", "metadata": "{}"}, ] - c = _make_collector(broker) + c = _make_collector(storage) c._discover_nodes() overview = c.get_overview() assert overview["nodes"] == 2 def test_discover_removes_lost_nodes(self): - broker = MockBroker() - broker.nodes = [{"node_id": "node-a", "server_url": "http://a:8080"}] - c = _make_collector(broker) + storage = MockStorage() + storage.services = [ + {"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"}, + ] + c = _make_collector(storage) c._discover_nodes() assert c.get_overview()["nodes"] == 1 # Node disappears - broker.nodes = [] + storage.services = [] c._discover_nodes() assert c.get_overview()["nodes"] == 0 def test_discover_updates_server_url(self): - broker = MockBroker() - broker.nodes = [{"node_id": "node-a", "server_url": "http://a:8080"}] - c = _make_collector(broker) + storage = MockStorage() + storage.services = [ + {"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"}, + ] + c = _make_collector(storage) c._discover_nodes() - broker.nodes = [{"node_id": "node-a", "server_url": "http://a:9090"}] + storage.services = [ + {"service_id": "node-a", "url": "http://a:9090", "metadata": "{}"}, + ] c._discover_nodes() detail = c.get_node_detail("node-a") assert detail["server_url"] == "http://a:9090" def test_discover_emits_node_joined_event(self): - broker = MockBroker() - c = _make_collector(broker) - _events = [] + storage = MockStorage() + c = _make_collector(storage) q = queue.Queue() c.register_listener(q) - broker.nodes = [{"node_id": "node-a", "server_url": "http://a:8080"}] + storage.services = [ + {"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"}, + ] c._discover_nodes() event = q.get_nowait() @@ -132,21 +121,41 @@ class TestCollectorDiscovery: assert event["node_id"] == "node-a" def test_discover_emits_node_lost_event(self): - broker = MockBroker() - broker.nodes = [{"node_id": "node-a", "server_url": "http://a:8080"}] - c = _make_collector(broker) + storage = MockStorage() + storage.services = [ + {"service_id": "node-a", "url": "http://a:8080", "metadata": "{}"}, + ] + c = _make_collector(storage) c._discover_nodes() q = queue.Queue() c.register_listener(q) - broker.nodes = [] + storage.services = [] c._discover_nodes() event = q.get_nowait() assert event["type"] == "node_lost" assert event["node_id"] == "node-a" + def test_discover_parses_metadata(self): + storage = MockStorage() + storage.services = [ + { + "service_id": "node-a", + "url": "http://a:8080", + "metadata": '{"max_ws": 20, "started": 1234567890.0}', + }, + ] + c = _make_collector(storage) + c._discover_nodes() + + detail = c.get_node_detail("node-a") + assert detail is not None + # Verify metadata was parsed into the NodeSnapshot + assert c._nodes["node-a"].max_ws == 20 + assert c._nodes["node-a"].started == 1234567890.0 + class TestCollectorPolling: """Polling /v1/api/dashboard from nodes.""" @@ -316,92 +325,8 @@ class TestCollectorPolling: assert "ws1" in c._nodes["node-a"].workstreams -class TestCollectorEvents: - """Real-time event handling from cluster channel.""" - - def test_cluster_state_event_updates_workstream(self): - c = _make_collector() - c._nodes["node-a"] = NodeSnapshot( - node_id="node-a", - server_url="http://a:8080", - workstreams={"ws1": {"id": "ws1", "name": "test", "state": "idle", "node": "node-a"}}, - ) - - event = ClusterStateEvent( - ws_id="ws1", - state="running", - node_id="node-a", - tokens=5000, - context_ratio=0.25, - activity="bash: echo hi", - activity_state="tool", - ) - c._on_cluster_event(event.to_json()) - - ws = c._nodes["node-a"].workstreams["ws1"] - assert ws["state"] == "running" - assert ws["tokens"] == 5000 - assert ws["activity"] == "bash: echo hi" - - def test_ws_created_event_adds_workstream(self): - c = _make_collector() - c._nodes["node-a"] = NodeSnapshot(node_id="node-a", server_url="http://a:8080") - - event_json = json.dumps( - { - "type": "ws_created", - "ws_id": "ws-new", - "name": "new-task", - "node_id": "node-a", - "correlation_id": "abc", - } - ) - c._on_cluster_event(event_json) - - assert "ws-new" in c._nodes["node-a"].workstreams - assert c._nodes["node-a"].workstreams["ws-new"]["name"] == "new-task" - assert c._nodes["node-a"].workstreams["ws-new"]["server_url"] == "http://a:8080" - - def test_ws_closed_event_removes_workstream(self): - c = _make_collector() - c._nodes["node-a"] = NodeSnapshot( - node_id="node-a", - workstreams={"ws1": {"id": "ws1", "state": "idle"}}, - ) - - event_json = json.dumps({"type": "ws_closed", "ws_id": "ws1"}) - c._on_cluster_event(event_json) - - assert "ws1" not in c._nodes["node-a"].workstreams - - def test_ws_rename_event_updates_name(self): - c = _make_collector() - c._nodes["node-a"] = NodeSnapshot( - node_id="node-a", - workstreams={"ws1": {"id": "ws1", "name": "old-name", "state": "idle"}}, - ) - - event_json = json.dumps({"type": "ws_rename", "ws_id": "ws1", "name": "new-name"}) - c._on_cluster_event(event_json) - - assert c._nodes["node-a"].workstreams["ws1"]["name"] == "new-name" - - def test_event_fans_out_to_listeners(self): - c = _make_collector() - c._nodes["node-a"] = NodeSnapshot( - node_id="node-a", - workstreams={"ws1": {"id": "ws1", "state": "idle", "node": "node-a"}}, - ) - - q = queue.Queue() - c.register_listener(q) - - event = ClusterStateEvent(ws_id="ws1", state="running", node_id="node-a") - c._on_cluster_event(event.to_json()) - - fan_event = q.get_nowait() - assert fan_event["type"] == "cluster_state" - assert fan_event["ws_id"] == "ws1" +class TestCollectorFanout: + """SSE fan-out to registered listeners.""" def test_unregister_listener_stops_fanout(self): c = _make_collector() @@ -414,12 +339,6 @@ class TestCollectorEvents: c._fanout({"type": "test"}) assert q.empty() - def test_invalid_json_event_ignored(self): - c = _make_collector() - # Should not raise - c._on_cluster_event("not valid json {{{") - c._on_cluster_event("") - class TestCollectorQueries: """Query methods: get_overview, get_nodes, get_workstreams, get_node_detail.""" @@ -598,50 +517,6 @@ class TestCollectorQueries: assert snap["overview"]["version_drift"] == overview["version_drift"] -# --------------------------------------------------------------------------- -# ClusterStateEvent protocol tests -# --------------------------------------------------------------------------- - - -class TestClusterStateEventProtocol: - """Ensure ClusterStateEvent round-trips through JSON correctly.""" - - def test_round_trip(self): - event = ClusterStateEvent( - ws_id="ws1", - state="running", - node_id="node-a", - tokens=5000, - context_ratio=0.25, - activity="bash: ls", - activity_state="tool", - ) - raw = event.to_json() - data = json.loads(raw) - assert data["type"] == "cluster_state" - assert data["ws_id"] == "ws1" - assert data["node_id"] == "node-a" - assert data["tokens"] == 5000 - assert data["context_ratio"] == 0.25 - - def test_from_json(self): - from turnstone.mq.protocol import OutboundEvent - - raw = json.dumps( - { - "type": "cluster_state", - "ws_id": "ws1", - "state": "running", - "node_id": "node-a", - "tokens": 5000, - } - ) - event = OutboundEvent.from_json(raw) - assert isinstance(event, ClusterStateEvent) - assert event.node_id == "node-a" - assert event.tokens == 5000 - - # --------------------------------------------------------------------------- # Console HTTP server tests # --------------------------------------------------------------------------- @@ -726,7 +601,6 @@ class TestConsoleHTTPEndpoints: app = create_app( collector=mock_collector, - broker=MagicMock(), auth_config=AuthConfig(), ) client = TestClient(app, raise_server_exceptions=False) @@ -927,7 +801,7 @@ class TestCollectorVersionInfo: class TestConsoleWorkstreamCreation: - """Tests for POST /v1/api/cluster/workstreams/new.""" + """Tests for POST /v1/api/cluster/workstreams/new (HTTP dispatch).""" @pytest.fixture() def mock_collector(self): @@ -958,25 +832,40 @@ class TestConsoleWorkstreamCreation: return collector @pytest.fixture() - def client_and_broker(self, mock_collector): + def client_and_mock(self, mock_collector): + """Returns (TestClient, mock_proxy_post) where mock_proxy_post is the + patched proxy_client.post that captures outgoing HTTP calls.""" + import httpx from starlette.testclient import TestClient from turnstone.console.server import _load_static, create_app from turnstone.core.auth import AuthConfig _load_static() - mock_broker = MagicMock() app = create_app( collector=mock_collector, - broker=mock_broker, auth_config=AuthConfig(), ) + + # Set up a mock proxy_client (lifespan doesn't run in TestClient) + async def _mock_post(*args, **kwargs): + return httpx.Response( + 200, + json={"ws_id": "ws_new_123", "name": "test"}, + request=httpx.Request("POST", args[0] if args else "http://test"), + ) + + mock_post = MagicMock(side_effect=_mock_post) + mock_proxy = MagicMock(spec=httpx.AsyncClient) + mock_proxy.post = mock_post + app.state.proxy_client = mock_proxy + client = TestClient(app, raise_server_exceptions=False) - yield client, mock_broker + yield client, mock_post client.close() - def test_create_with_explicit_node(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_with_explicit_node(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "node-a", "name": "test-ws"}, @@ -986,47 +875,46 @@ class TestConsoleWorkstreamCreation: assert data["status"] == "ok" assert data["target_node"] == "node-a" assert "correlation_id" in data - broker.push_inbound.assert_called_once() - # Verify the pushed message - msg_json = broker.push_inbound.call_args[0][0] - msg = json.loads(msg_json) - assert msg["type"] == "create_workstream" - assert msg["target_node"] == "node-a" - assert msg["name"] == "test-ws" + mock_post.assert_called_once() + # Verify the HTTP call was to the right node + call_args = mock_post.call_args + assert "http://a:8080/v1/api/workstreams/new" in call_args[0] + body = call_args[1]["json"] + assert body["name"] == "test-ws" - def test_create_with_model(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_with_model(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "node-a", "model": "gpt-5"}, ) assert resp.status_code == 200 - msg_json = broker.push_inbound.call_args[0][0] - msg = json.loads(msg_json) - assert msg["model"] == "gpt-5" + body = mock_post.call_args[1]["json"] + assert body["model"] == "gpt-5" - def test_create_with_initial_message_directed(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_with_initial_message_directed(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "node-a", "initial_message": "Do the thing"}, ) assert resp.status_code == 200 - msg = json.loads(broker.push_inbound.call_args[0][0]) - assert msg["initial_message"] == "Do the thing" + body = mock_post.call_args[1]["json"] + assert body["initial_message"] == "Do the thing" - def test_create_with_initial_message_pool(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_with_initial_message_pool(self, client_and_mock, mock_collector): + """Pool mode picks the best node and dispatches via HTTP.""" + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "pool", "initial_message": "Pool task"}, ) assert resp.status_code == 200 - msg = json.loads(broker.push_inbound.call_args[0][0]) - assert msg["initial_message"] == "Pool task" + body = mock_post.call_args[1]["json"] + assert body["initial_message"] == "Pool task" - def test_create_auto_selects_best_node(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_auto_selects_best_node(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"name": "auto-test"}, @@ -1036,15 +924,15 @@ class TestConsoleWorkstreamCreation: # node-b has more headroom (10-3=7 vs 10-8=2) assert data["target_node"] == "node-b" - def test_create_no_reachable_nodes(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_no_reachable_nodes(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock mock_collector.get_nodes.return_value = ([], 0) resp = client.post("/v1/api/cluster/workstreams/new", json={}) assert resp.status_code == 503 assert "No reachable nodes" in resp.json()["error"] - def test_create_unknown_node(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_unknown_node(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock mock_collector.get_node_detail.return_value = None resp = client.post( "/v1/api/cluster/workstreams/new", @@ -1052,8 +940,8 @@ class TestConsoleWorkstreamCreation: ) assert resp.status_code == 404 - def test_create_invalid_json(self, client_and_broker): - client, broker = client_and_broker + def test_create_invalid_json(self, client_and_mock): + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", content=b"not json", @@ -1061,19 +949,19 @@ class TestConsoleWorkstreamCreation: ) assert resp.status_code == 400 - def test_create_pushes_to_directed_queue(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_dispatches_to_correct_node_url(self, client_and_mock, mock_collector): + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "node-a"}, ) assert resp.status_code == 200 - # Verify push_inbound called with node_id kwarg - call_kwargs = broker.push_inbound.call_args - assert call_kwargs[1]["node_id"] == "node-a" + call_args = mock_post.call_args + assert "http://a:8080/v1/api/workstreams/new" in call_args[0] - def test_create_pool_pushes_to_shared_queue(self, client_and_broker, mock_collector): - client, broker = client_and_broker + def test_create_pool_picks_best_node(self, client_and_mock, mock_collector): + """Pool mode dispatches to the best available node.""" + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "pool", "name": "pool-task"}, @@ -1081,60 +969,40 @@ class TestConsoleWorkstreamCreation: assert resp.status_code == 200 data = resp.json() assert data["status"] == "ok" - assert data["target_node"] == "pool" - broker.push_inbound.assert_called_once() - # Shared queue: no node_id kwarg (or empty) - call_args = broker.push_inbound.call_args - assert call_args[1].get("node_id", "") == "" - # Message should have no target_node - msg = json.loads(call_args[0][0]) - assert msg["type"] == "create_workstream" - assert msg["target_node"] == "" - assert msg["name"] == "pool-task" + # Pool picks best node (node-b has most headroom) + assert data["target_node"] == "node-b" - def test_create_pool_skips_node_validation(self, client_and_broker, mock_collector): - """Pool mode doesn't need a valid node_id — it goes to the shared queue.""" - client, broker = client_and_broker - mock_collector.get_node_detail.return_value = None # would 404 for directed + def test_create_pool_no_nodes_returns_503(self, client_and_mock, mock_collector): + """Pool mode with no reachable nodes returns 503.""" + client, mock_post = client_and_mock + mock_collector.get_nodes.return_value = ([], 0) resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "pool"}, ) - assert resp.status_code == 200 - assert resp.json()["target_node"] == "pool" + assert resp.status_code == 503 - def test_create_with_resume_ws_directed(self, client_and_broker, mock_collector): + def test_create_with_resume_ws_directed(self, client_and_mock, mock_collector): """resume_ws is forwarded in directed dispatch.""" - client, broker = client_and_broker + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"node_id": "node-a", "resume_ws": "old-ws-id-123"}, ) assert resp.status_code == 200 - msg = json.loads(broker.push_inbound.call_args[0][0]) - assert msg["resume_ws"] == "old-ws-id-123" + body = mock_post.call_args[1]["json"] + assert body["resume_ws"] == "old-ws-id-123" - def test_create_with_resume_ws_pool(self, client_and_broker, mock_collector): - """resume_ws is forwarded in pool dispatch.""" - client, broker = client_and_broker - resp = client.post( - "/v1/api/cluster/workstreams/new", - json={"node_id": "pool", "resume_ws": "old-ws-id-456"}, - ) - assert resp.status_code == 200 - msg = json.loads(broker.push_inbound.call_args[0][0]) - assert msg["resume_ws"] == "old-ws-id-456" - - def test_create_with_resume_ws_auto(self, client_and_broker, mock_collector): + def test_create_with_resume_ws_auto(self, client_and_mock, mock_collector): """resume_ws is forwarded in auto-select dispatch.""" - client, broker = client_and_broker + client, mock_post = client_and_mock resp = client.post( "/v1/api/cluster/workstreams/new", json={"resume_ws": "old-ws-id-789"}, ) assert resp.status_code == 200 - msg = json.loads(broker.push_inbound.call_args[0][0]) - assert msg["resume_ws"] == "old-ws-id-789" + body = mock_post.call_args[1]["json"] + assert body["resume_ws"] == "old-ws-id-789" # --------------------------------------------------------------------------- @@ -1174,7 +1042,6 @@ class TestConsoleProxy: _load_static() app = create_app( collector=mock_collector, - broker=MagicMock(), auth_config=AuthConfig(), ) client = TestClient(app, raise_server_exceptions=False) @@ -1346,7 +1213,6 @@ class TestConsoleVersionEndpoints: _load_static() app = create_app( collector=mock_collector, - broker=MagicMock(), auth_config=AuthConfig(), ) client = TestClient(app, raise_server_exceptions=False) @@ -1396,7 +1262,6 @@ class TestSharedStatic: } app = create_app( collector=collector, - broker=MagicMock(), auth_config=AuthConfig(), ) client = TestClient(app, raise_server_exceptions=False) @@ -1481,7 +1346,6 @@ class TestProxySharedStatic: def test_proxy_shim_injected_in_html(self): """Verify shim is injected as inline script in proxied HTML.""" - import json from turnstone.console.server import _CONSOLE_BANNER_TEMPLATE, _JS_PROXY_SHIM @@ -1516,7 +1380,6 @@ class TestProxySharedStatic: collector.get_node_detail.return_value = None app = create_app( collector=collector, - broker=MagicMock(), auth_config=AuthConfig(), ) client = TestClient(app, raise_server_exceptions=False) diff --git a/tests/test_model_registry.py b/tests/test_model_registry.py index fd445225..6f3d1790 100644 --- a/tests/test_model_registry.py +++ b/tests/test_model_registry.py @@ -840,32 +840,27 @@ class TestWorkstreamModelParam: # --------------------------------------------------------------------------- -# Protocol +# CreateWorkstreamRequest model field # --------------------------------------------------------------------------- -class TestProtocolModel: - def test_create_workstream_message_has_model(self) -> None: - from turnstone.mq.protocol import CreateWorkstreamMessage +class TestCreateWorkstreamRequestModel: + def test_request_has_model(self) -> None: + from turnstone.api.server_schemas import CreateWorkstreamRequest - msg = CreateWorkstreamMessage(name="test", model="openai") - assert msg.model == "openai" + req = CreateWorkstreamRequest(name="test", model="openai") + assert req.model == "openai" - def test_create_workstream_message_default(self) -> None: - from turnstone.mq.protocol import CreateWorkstreamMessage + def test_request_model_default(self) -> None: + from turnstone.api.server_schemas import CreateWorkstreamRequest - msg = CreateWorkstreamMessage(name="test") - assert msg.model == "" + req = CreateWorkstreamRequest(name="test") + assert req.model == "" - def test_round_trip(self) -> None: - from turnstone.mq.protocol import CreateWorkstreamMessage, InboundMessage - - msg = CreateWorkstreamMessage(name="ws1", model="local") - raw = msg.to_json() - restored = InboundMessage.from_json(raw) - assert isinstance(restored, CreateWorkstreamMessage) - assert restored.model == "local" - assert restored.name == "ws1" + def test_json_payload_carries_model(self) -> None: + body = {"name": "ws1", "model": "local"} + assert body["model"] == "local" + assert body["name"] == "ws1" # --------------------------------------------------------------------------- diff --git a/tests/test_protocol.py b/tests/test_protocol.py deleted file mode 100644 index 929cbbe1..00000000 --- a/tests/test_protocol.py +++ /dev/null @@ -1,239 +0,0 @@ -"""Tests for turnstone.mq.protocol message serialization.""" - -import json - -import pytest - -from turnstone.mq.protocol import ( - AckEvent, - ApprovalRequestEvent, - ApproveMessage, - CancelMessage, - CloseWorkstreamMessage, - CommandMessage, - ContentEvent, - CreateWorkstreamMessage, - ErrorEvent, - HealthMessage, - HealthResponseEvent, - InboundMessage, - InfoEvent, - ListNodesMessage, - ListWorkstreamsMessage, - NodeListEvent, - OutboundEvent, - PlanFeedbackMessage, - PlanReviewEvent, - ReasoningEvent, - SendMessage, - StateChangeEvent, - StatusEvent, - StreamEndEvent, - ToolInfoEvent, - ToolResultEvent, - TurnCompleteEvent, - WorkstreamClosedEvent, - WorkstreamCreatedEvent, - WorkstreamListEvent, - WorkstreamRenameEvent, -) - -# --------------------------------------------------------------------------- -# Inbound message round-trip tests -# --------------------------------------------------------------------------- - -INBOUND_TYPES = [ - ( - SendMessage, - { - "message": "hello", - "ws_id": "abc", - "auto_approve": True, - "auto_approve_tools": ["bash"], - }, - ), - ( - ApproveMessage, - {"ws_id": "abc", "request_id": "r1", "approved": True, "feedback": "ok"}, - ), - ( - PlanFeedbackMessage, - {"ws_id": "abc", "request_id": "r2", "feedback": "looks good"}, - ), - (CommandMessage, {"ws_id": "abc", "command": "/clear"}), - ( - CreateWorkstreamMessage, - {"name": "test-ws", "auto_approve": False, "auto_approve_tools": ["read_file"]}, - ), - (CloseWorkstreamMessage, {"ws_id": "abc"}), - (ListWorkstreamsMessage, {}), - (HealthMessage, {}), - (ListNodesMessage, {}), - (CancelMessage, {"ws_id": "abc"}), -] - - -@pytest.mark.parametrize("cls,kwargs", INBOUND_TYPES) -def test_inbound_round_trip(cls, kwargs): - msg = cls(**kwargs) - raw = msg.to_json() - parsed = json.loads(raw) - - # type field matches - assert parsed["type"] == msg.type - - # correlation_id auto-generated - assert len(msg.correlation_id) == 12 - assert parsed["correlation_id"] == msg.correlation_id - - # timestamp present - assert msg.timestamp > 0 - - # Deserialize back - restored = InboundMessage.from_json(raw) - assert type(restored) is cls - assert restored.type == msg.type - assert restored.correlation_id == msg.correlation_id - - # Check custom fields - for k, v in kwargs.items(): - assert getattr(restored, k) == v - - -def test_inbound_unknown_type(): - with pytest.raises(ValueError, match="Unknown inbound"): - InboundMessage.from_json('{"type": "nonexistent"}') - - -def test_inbound_extra_fields_ignored(): - raw = json.dumps({"type": "send", "message": "hi", "extra_field": 42}) - msg = InboundMessage.from_json(raw) - assert isinstance(msg, SendMessage) - assert msg.message == "hi" - assert not hasattr(msg, "extra_field") - - -# --------------------------------------------------------------------------- -# Outbound event round-trip tests -# --------------------------------------------------------------------------- - -OUTBOUND_TYPES = [ - (AckEvent, {"status": "ok", "detail": "done"}), - (ContentEvent, {"text": "hello world"}), - (ReasoningEvent, {"text": "thinking..."}), - (ToolInfoEvent, {"items": [{"name": "bash", "preview": "ls"}]}), - (ApprovalRequestEvent, {"items": [{"name": "bash", "needs_approval": True}]}), - (ToolResultEvent, {"call_id": "call_123", "name": "bash", "output": "file.txt"}), - (PlanReviewEvent, {"content": "# Plan\n\nStep 1: ..."}), - (StatusEvent, {"prompt_tokens": 100, "completion_tokens": 50, "pct": 0.42}), - (StateChangeEvent, {"state": "thinking"}), - (TurnCompleteEvent, {}), - (StreamEndEvent, {}), - (WorkstreamCreatedEvent, {"name": "test-ws"}), - (WorkstreamClosedEvent, {}), - (WorkstreamListEvent, {"workstreams": [{"id": "abc", "name": "ws"}]}), - (WorkstreamRenameEvent, {"name": "renamed"}), - (HealthResponseEvent, {"data": {"status": "ok"}}), - (ErrorEvent, {"message": "something broke"}), - (InfoEvent, {"message": "heads up"}), - ( - NodeListEvent, - {"nodes": [{"node_id": "server-12", "server_url": "http://x:8080"}]}, - ), -] - - -@pytest.mark.parametrize("cls,kwargs", OUTBOUND_TYPES) -def test_outbound_round_trip(cls, kwargs): - event = cls(ws_id="ws1", correlation_id="c1", **kwargs) - raw = event.to_json() - parsed = json.loads(raw) - - assert parsed["type"] == event.type - assert parsed["ws_id"] == "ws1" - assert parsed["correlation_id"] == "c1" - - restored = OutboundEvent.from_json(raw) - assert type(restored) is cls - assert restored.ws_id == "ws1" - assert restored.correlation_id == "c1" - - for k, v in kwargs.items(): - assert getattr(restored, k) == v - - -def test_outbound_unknown_type_falls_back(): - raw = json.dumps({"type": "future_event", "ws_id": "x"}) - event = OutboundEvent.from_json(raw) - assert isinstance(event, OutboundEvent) - assert event.ws_id == "x" - - -def test_send_message_defaults(): - msg = SendMessage(message="hello") - assert msg.ws_id == "" - assert msg.auto_approve is False - assert msg.auto_approve_tools == [] - assert msg.name == "" - assert msg.target_node == "" - assert len(msg.correlation_id) == 12 - - -def test_create_workstream_with_tools(): - msg = CreateWorkstreamMessage( - name="ci-runner", - auto_approve=False, - auto_approve_tools=["bash", "read_file", "search"], - ) - raw = msg.to_json() - restored = InboundMessage.from_json(raw) - assert restored.auto_approve_tools == ["bash", "read_file", "search"] - assert restored.name == "ci-runner" - - -def test_send_message_target_node(): - msg = SendMessage(message="check disk", target_node="server-12") - raw = msg.to_json() - restored = InboundMessage.from_json(raw) - assert isinstance(restored, SendMessage) - assert restored.target_node == "server-12" - assert restored.message == "check disk" - - -def test_create_workstream_target_node(): - msg = CreateWorkstreamMessage(name="debug-ws", target_node="gpu-node-3") - raw = msg.to_json() - restored = InboundMessage.from_json(raw) - assert isinstance(restored, CreateWorkstreamMessage) - assert restored.target_node == "gpu-node-3" - assert restored.name == "debug-ws" - - -def test_create_workstream_skill_field(): - msg = CreateWorkstreamMessage(name="ws", skill="code-review") - assert msg.skill == "code-review" - raw = msg.to_json() - restored = InboundMessage.from_json(raw) - assert isinstance(restored, CreateWorkstreamMessage) - assert restored.skill == "code-review" - - -def test_create_workstream_skill_default_empty(): - msg = CreateWorkstreamMessage(name="ws") - assert msg.skill == "" - - -def test_list_nodes_round_trip(): - msg = ListNodesMessage() - raw = msg.to_json() - restored = InboundMessage.from_json(raw) - assert isinstance(restored, ListNodesMessage) - - -def test_node_list_event_round_trip(): - nodes = [{"node_id": "a", "server_url": "http://a:8080"}] - event = NodeListEvent(nodes=nodes, correlation_id="c1") - raw = event.to_json() - restored = OutboundEvent.from_json(raw) - assert isinstance(restored, NodeListEvent) - assert restored.nodes == nodes diff --git a/tests/test_resume_flow.py b/tests/test_resume_flow.py index bf22c5cc..52a0b66f 100644 --- a/tests/test_resume_flow.py +++ b/tests/test_resume_flow.py @@ -1,97 +1,39 @@ -"""Tests for the atomic workstream resumption flow. +"""Tests for the workstream resume request schema. -Covers CreateWorkstreamMessage resume_ws field, WorkstreamResumedEvent, -WorkstreamCreatedEvent resumed fields, and server endpoint handling. +Verifies that the create-workstream JSON payload carries the resume_ws field +correctly, matching the server's ``CreateWorkstreamRequest`` schema. """ from __future__ import annotations -import json - -from turnstone.mq.protocol import ( - CreateWorkstreamMessage, - WorkstreamCreatedEvent, - WorkstreamResumedEvent, -) - # --------------------------------------------------------------------------- -# Protocol tests +# CreateWorkstreamRequest resume_ws field # --------------------------------------------------------------------------- -class TestCreateWorkstreamMessageResumeField: +class TestCreateWorkstreamResumeField: def test_resume_ws_defaults_empty(self) -> None: - msg = CreateWorkstreamMessage(name="test") - assert msg.resume_ws == "" + body: dict[str, str] = {"name": "test"} + assert body.get("resume_ws", "") == "" def test_resume_ws_set(self) -> None: - msg = CreateWorkstreamMessage(name="test", resume_ws="ws-abc") - assert msg.resume_ws == "ws-abc" + body = {"name": "test", "resume_ws": "ws-abc"} + assert body["resume_ws"] == "ws-abc" - def test_resume_ws_serializes(self) -> None: - msg = CreateWorkstreamMessage(resume_ws="ws-xyz") - data = json.loads(msg.to_json()) - assert data["resume_ws"] == "ws-xyz" + def test_resume_ws_present_in_payload(self) -> None: + body = {"name": "test", "resume_ws": "ws-xyz"} + assert "resume_ws" in body + assert body["resume_ws"] == "ws-xyz" - def test_resume_ws_deserializes(self) -> None: - msg = CreateWorkstreamMessage(resume_ws="ws-123") - raw = msg.to_json() - from turnstone.mq.protocol import InboundMessage + def test_pydantic_schema_has_resume_ws(self) -> None: + """CreateWorkstreamRequest schema includes resume_ws.""" + from turnstone.api.server_schemas import CreateWorkstreamRequest - restored = InboundMessage.from_json(raw) - assert getattr(restored, "resume_ws", "") == "ws-123" + req = CreateWorkstreamRequest(name="test", resume_ws="ws-123") + assert req.resume_ws == "ws-123" + def test_pydantic_schema_default_empty(self) -> None: + from turnstone.api.server_schemas import CreateWorkstreamRequest -class TestWorkstreamCreatedEventResumeFields: - def test_default_not_resumed(self) -> None: - event = WorkstreamCreatedEvent(ws_id="ws-1", name="test") - assert event.resumed is False - assert event.message_count == 0 - - def test_resumed_fields(self) -> None: - event = WorkstreamCreatedEvent(ws_id="ws-1", name="test", resumed=True, message_count=42) - assert event.resumed is True - assert event.message_count == 42 - - def test_serializes_resumed_fields(self) -> None: - event = WorkstreamCreatedEvent(ws_id="ws-1", resumed=True, message_count=10) - data = json.loads(event.to_json()) - assert data["resumed"] is True - assert data["message_count"] == 10 - - def test_deserializes_resumed_fields(self) -> None: - event = WorkstreamCreatedEvent(ws_id="ws-1", resumed=True, message_count=5) - from turnstone.mq.protocol import OutboundEvent - - restored = OutboundEvent.from_json(event.to_json()) - assert isinstance(restored, WorkstreamCreatedEvent) - assert restored.resumed is True - assert restored.message_count == 5 - - -class TestWorkstreamResumedEvent: - def test_defaults(self) -> None: - event = WorkstreamResumedEvent(ws_id="ws-1") - assert event.type == "ws_resumed" - assert event.message_count == 0 - assert event.name == "" - - def test_with_values(self) -> None: - event = WorkstreamResumedEvent(ws_id="ws-1", message_count=25, name="My Chat") - assert event.message_count == 25 - assert event.name == "My Chat" - - def test_round_trip(self) -> None: - event = WorkstreamResumedEvent(ws_id="ws-1", message_count=10, name="Chat") - from turnstone.mq.protocol import OutboundEvent - - restored = OutboundEvent.from_json(event.to_json()) - assert isinstance(restored, WorkstreamResumedEvent) - assert restored.message_count == 10 - assert restored.name == "Chat" - - def test_registered_in_outbound_registry(self) -> None: - from turnstone.mq.protocol import _OUTBOUND_REGISTRY - - assert "ws_resumed" in _OUTBOUND_REGISTRY - assert _OUTBOUND_REGISTRY["ws_resumed"] is WorkstreamResumedEvent + req = CreateWorkstreamRequest(name="test") + assert req.resume_ws == "" diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index 455d1b45..391dae97 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -2,21 +2,45 @@ from __future__ import annotations -from unittest.mock import MagicMock +import json +from unittest.mock import MagicMock, patch import pytest from turnstone.console.scheduler import TaskScheduler +def _wire_lock_storage(storage: MagicMock, initial: dict[str, str] | None = None) -> None: + """Configure *storage* mock so upsert/get track scheduler_lock state. + + The scheduler's ``_try_acquire_lock`` now writes then reads back to + verify ownership. The mock must reflect what was most recently + upserted so the read-back succeeds. + """ + state: dict[str, dict[str, str] | None] = {"scheduler_lock": initial} + + def _get(key: str, **_kw: object) -> dict[str, str] | None: + return state.get(key) + + def _upsert(key: str, value: str, **_kw: object) -> None: + state[key] = {"value": value} + + def _delete(key: str, **_kw: object) -> None: + state.pop(key, None) + + storage.get_system_setting.side_effect = _get + storage.upsert_system_setting.side_effect = _upsert + storage.delete_system_setting.side_effect = _delete + + @pytest.fixture def mocks(): - """Broker, collector, and storage mocks for scheduler tests.""" - broker = MagicMock() - broker._redis = MagicMock() + """Collector and storage mocks for scheduler tests.""" collector = MagicMock() storage = MagicMock() - return broker, collector, storage + # Default: no existing lock + _wire_lock_storage(storage, initial=None) + return collector, storage def _make_task(**overrides): @@ -58,70 +82,93 @@ class TestSchedulerTick: """Tests for _tick() lock acquisition and dispatch logic.""" def test_tick_acquires_lock(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks storage.list_due_tasks.return_value = [] - scheduler = TaskScheduler(broker, collector, storage) + scheduler = TaskScheduler(collector, storage) scheduler._tick() - broker._redis.set.assert_called_once() + storage.get_system_setting.assert_called() + storage.upsert_system_setting.assert_called() storage.list_due_tasks.assert_called_once() - # Lock released via Lua eval (conditional delete) - broker._redis.eval.assert_called_once() def test_tick_skips_when_locked(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = None # lock held by another console + collector, storage = mocks + # Another instance holds the lock (recent timestamp) + from datetime import UTC, datetime - scheduler = TaskScheduler(broker, collector, storage) + now_str = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S") + _wire_lock_storage( + storage, + initial={"value": json.dumps({"owner": "other-instance", "acquired": now_str})}, + ) + + scheduler = TaskScheduler(collector, storage) scheduler._tick() storage.list_due_tasks.assert_not_called() + def test_tick_takes_expired_lock(self, mocks): + """An expired lock from another instance should be taken over.""" + collector, storage = mocks + _wire_lock_storage( + storage, + initial={ + "value": json.dumps({"owner": "other-instance", "acquired": "2020-01-01T00:00:00"}) + }, + ) + storage.list_due_tasks.return_value = [] + + scheduler = TaskScheduler(collector, storage) + scheduler._tick() + + storage.list_due_tasks.assert_called_once() + def test_dispatch_auto_mode(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="auto") storage.list_due_tasks.return_value = [task] collector.get_nodes.return_value = ([_make_node()], 1) + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - broker.push_inbound.assert_called_once() - _, kwargs = broker.push_inbound.call_args - assert ( - kwargs.get("node_id") == "node-001" - or broker.push_inbound.call_args[1].get("node_id") == "node-001" - ) + mock_post.assert_called_once() + url = mock_post.call_args[0][0] + assert "http://node-001:8080/v1/api/workstreams/new" in url storage.record_task_run.assert_called_once() run_kwargs = storage.record_task_run.call_args[1] assert run_kwargs["node_id"] == "node-001" assert run_kwargs["status"] == "dispatched" def test_dispatch_pool_mode(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="pool") storage.list_due_tasks.return_value = [task] + collector.get_nodes.return_value = ([_make_node("node-001")], 1) + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - broker.push_inbound.assert_called_once() - # Pool dispatch calls push_inbound without node_id kwarg - args, kwargs = broker.push_inbound.call_args - assert kwargs.get("node_id") is None or "node_id" not in kwargs + mock_post.assert_called_once() storage.record_task_run.assert_called_once() - run_kwargs = storage.record_task_run.call_args[1] - assert run_kwargs["node_id"] == "pool" def test_dispatch_all_mode(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="all") storage.list_due_tasks.return_value = [task] @@ -129,37 +176,53 @@ class TestSchedulerTick: [_make_node("node-001"), _make_node("node-002")], 2, ) + collector.get_node_detail.side_effect = lambda nid: { + "server_url": f"http://{nid}:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - assert broker.push_inbound.call_count == 2 + assert mock_post.call_count == 2 assert storage.record_task_run.call_count == 2 def test_dispatch_specific_node(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="node-001") storage.list_due_tasks.return_value = [task] + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - broker.push_inbound.assert_called_once() - _, kwargs = broker.push_inbound.call_args - assert kwargs["node_id"] == "node-001" + mock_post.assert_called_once() + url = mock_post.call_args[0][0] + assert "node-001" in url def test_at_task_disables_after_dispatch(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(schedule_type="at", cron_expr="", at_time="2099-01-01T00:00:00") storage.list_due_tasks.return_value = [task] collector.get_nodes.return_value = ([_make_node()], 1) + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() # At-task should be disabled after dispatch update_calls = storage.update_scheduled_task.call_args_list @@ -170,15 +233,20 @@ class TestSchedulerTick: assert kwargs["next_run"] == "" def test_cron_task_updates_next_run(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(schedule_type="cron", cron_expr="0 9 * * *") storage.list_due_tasks.return_value = [task] collector.get_nodes.return_value = ([_make_node()], 1) + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() update_calls = storage.update_scheduled_task.call_args_list assert len(update_calls) == 1 @@ -187,8 +255,7 @@ class TestSchedulerTick: assert "enabled" not in kwargs # cron tasks stay enabled def test_no_reachable_nodes_records_failure(self, mocks): - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="auto") storage.list_due_tasks.return_value = [task] @@ -198,10 +265,9 @@ class TestSchedulerTick: 1, ) - scheduler = TaskScheduler(broker, collector, storage) + scheduler = TaskScheduler(collector, storage) scheduler._tick() - broker.push_inbound.assert_not_called() storage.record_task_run.assert_called_once() run_kwargs = storage.record_task_run.call_args[1] assert run_kwargs["status"] == "failed" @@ -209,14 +275,13 @@ class TestSchedulerTick: def test_failure_does_not_advance_schedule(self, mocks): """When dispatch fails, last_run/next_run should not be updated.""" - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="auto") storage.list_due_tasks.return_value = [task] collector.get_nodes.return_value = ([], 0) # no nodes at all - scheduler = TaskScheduler(broker, collector, storage) + scheduler = TaskScheduler(collector, storage) scheduler._tick() # update_scheduled_task should NOT be called (no last_run/next_run advance) @@ -224,49 +289,84 @@ class TestSchedulerTick: def test_fan_out_capped(self, mocks): """Fan-out 'all' mode should respect max_fan_out limit.""" - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="all") storage.list_due_tasks.return_value = [task] # 10 reachable nodes but max_fan_out=3 nodes = [_make_node(f"node-{i:03d}") for i in range(10)] collector.get_nodes.return_value = (nodes, 10) + collector.get_node_detail.side_effect = lambda nid: { + "server_url": f"http://{nid}:8080", + } - scheduler = TaskScheduler(broker, collector, storage, max_fan_out=3) - scheduler._tick() + scheduler = TaskScheduler(collector, storage, max_fan_out=3) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - assert broker.push_inbound.call_count == 3 + assert mock_post.call_count == 3 assert storage.record_task_run.call_count == 3 def test_specific_node_target(self, mocks): """Non-enum target_mode is treated as a specific node_id.""" - broker, collector, storage = mocks - broker._redis.set.return_value = True + collector, storage = mocks task = _make_task(target_mode="node-custom-123") storage.list_due_tasks.return_value = [task] + collector.get_node_detail.return_value = { + "server_url": "http://node-custom-123:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - broker.push_inbound.assert_called_once() - call_kwargs = broker.push_inbound.call_args - assert call_kwargs[1]["node_id"] == "node-custom-123" + mock_post.assert_called_once() + url = mock_post.call_args[0][0] + assert "node-custom-123" in url - def test_user_id_in_dispatched_message(self, mocks): - """Dispatched message should include created_by as user_id.""" - import json + def test_user_id_in_dispatched_body(self, mocks): + """Dispatched HTTP body should include created_by as user_id.""" + collector, storage = mocks - broker, collector, storage = mocks - broker._redis.set.return_value = True - - task = _make_task(target_mode="pool", created_by="u_scheduler_admin") + task = _make_task(target_mode="auto", created_by="u_scheduler_admin") storage.list_due_tasks.return_value = [task] + collector.get_nodes.return_value = ([_make_node()], 1) + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } - scheduler = TaskScheduler(broker, collector, storage) - scheduler._tick() + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.return_value = MagicMock(status_code=200) + mock_post.return_value.raise_for_status = MagicMock() + scheduler._tick() - msg_json = broker.push_inbound.call_args[0][0] - msg_data = json.loads(msg_json) - assert msg_data["user_id"] == "u_scheduler_admin" + body = mock_post.call_args[1]["json"] + assert body["user_id"] == "u_scheduler_admin" + + def test_http_failure_records_failure(self, mocks): + """HTTP errors during dispatch should record a failure.""" + import httpx + + collector, storage = mocks + + task = _make_task(target_mode="auto") + storage.list_due_tasks.return_value = [task] + collector.get_nodes.return_value = ([_make_node()], 1) + collector.get_node_detail.return_value = { + "server_url": "http://node-001:8080", + } + + scheduler = TaskScheduler(collector, storage) + with patch.object(scheduler._http_client, "post") as mock_post: + mock_post.side_effect = httpx.ConnectError("connection refused") + scheduler._tick() + + storage.record_task_run.assert_called_once() + run_kwargs = storage.record_task_run.call_args[1] + assert run_kwargs["status"] == "failed" diff --git a/tests/test_sim.py b/tests/test_sim.py deleted file mode 100644 index c715a089..00000000 --- a/tests/test_sim.py +++ /dev/null @@ -1,351 +0,0 @@ -"""Tests for the turnstone cluster simulator.""" - -from __future__ import annotations - -import asyncio -import random -from unittest.mock import MagicMock - -import pytest - -from turnstone.mq.protocol import ( - OutboundEvent, - SendMessage, - StateChangeEvent, -) -from turnstone.sim.config import SimConfig -from turnstone.sim.engine import SimEngine, ToolSimulationError -from turnstone.sim.metrics import MetricsCollector -from turnstone.sim.node import SimNode, SimWorkstream - - -def _run(coro): - """Run an async coroutine synchronously.""" - return asyncio.run(coro) - - -# --------------------------------------------------------------------------- -# SimConfig -# --------------------------------------------------------------------------- - - -class TestSimConfig: - def test_defaults(self): - cfg = SimConfig() - assert cfg.num_nodes == 10 - assert cfg.scenario == "steady" - assert cfg.llm_latency_mean == 2.0 - assert cfg.tool_failure_rate == 0.02 - - def test_frozen(self): - cfg = SimConfig() - with pytest.raises(AttributeError): - cfg.num_nodes = 5 # type: ignore[misc] - - def test_custom_values(self): - cfg = SimConfig(num_nodes=100, scenario="burst", seed=42) - assert cfg.num_nodes == 100 - assert cfg.scenario == "burst" - assert cfg.seed == 42 - - -# --------------------------------------------------------------------------- -# SimEngine -# --------------------------------------------------------------------------- - - -class TestSimEngine: - @pytest.fixture - def fast_config(self): - return SimConfig( - llm_latency_mean=0.01, - llm_latency_stddev=0.001, - llm_tokens_mean=20, - llm_tokens_stddev=5, - tool_latency_mean=0.01, - tool_latency_stddev=0.001, - tool_failure_rate=0.0, - seed=42, - ) - - @pytest.fixture - def engine(self, fast_config): - return SimEngine(fast_config) - - def test_llm_response_returns_content(self, engine): - async def _test(): - content, tool_calls = await engine.simulate_llm_response(True) - assert isinstance(content, str) - assert len(content) > 0 - assert isinstance(tool_calls, list) - - _run(_test()) - - def test_llm_response_reproducible_with_seed(self, fast_config): - async def _test(): - e1 = SimEngine(fast_config, rng=random.Random(123)) - e2 = SimEngine(fast_config, rng=random.Random(123)) - c1, t1 = await e1.simulate_llm_response(True) - c2, t2 = await e2.simulate_llm_response(True) - assert c1 == c2 - assert len(t1) == len(t2) - - _run(_test()) - - def test_tool_execution_success(self, engine): - async def _test(): - result = await engine.simulate_tool_execution("bash") - assert "bash" in result - assert "completed" in result - - _run(_test()) - - def test_tool_execution_failure(self, fast_config): - cfg = SimConfig( - llm_latency_mean=0.01, - tool_latency_mean=0.01, - tool_latency_stddev=0.001, - tool_failure_rate=1.0, # always fail - seed=42, - ) - engine = SimEngine(cfg) - - async def _test(): - with pytest.raises(ToolSimulationError, match="Simulated bash failure"): - await engine.simulate_tool_execution("bash") - - _run(_test()) - - def test_generate_content(self, engine): - content = engine._generate_content(10) - words = content.split() - assert len(words) == 10 - - -# --------------------------------------------------------------------------- -# MetricsCollector -# --------------------------------------------------------------------------- - - -class TestMetricsCollector: - def test_record_and_summary(self): - m = MetricsCollector() - m.record_inject() - m.record_turn("ws1", "node-0", 1.5) - m.record_turn("ws2", "node-0", 2.5) - m.record_turn("ws3", "node-1", 3.0) - m.record_error("node-0", "test error") - - report = m.summary() - assert report["total_turns"] == 3 - assert report["total_errors"] == 1 - assert report["latency"]["p50"] == 2.5 - assert report["latency"]["max"] == 3.0 - assert report["turns_per_node"]["node-0"] == 2 - assert report["turns_per_node"]["node-1"] == 1 - - def test_empty_summary(self): - m = MetricsCollector() - report = m.summary() - assert report["total_turns"] == 0 - assert report["latency"]["p50"] == 0 - - def test_node_kill_tracking(self): - m = MetricsCollector() - m.record_node_kill("node-0") - m.record_node_kill("node-1") - report = m.summary() - assert report["node_kills"] == 2 - - def test_utilization_snapshot(self): - m = MetricsCollector() - m.snapshot_utilization({"node-0": 3, "node-1": 5, "node-2": 0}) - report = m.summary() - assert report["utilization"]["mean_ws_per_node"] == pytest.approx(8 / 3) - assert report["utilization"]["max_ws_per_node"] == 5 - assert report["utilization"]["nodes_with_zero_ws"] == 1 - - -# --------------------------------------------------------------------------- -# SimNode — message dispatch -# --------------------------------------------------------------------------- - - -class TestSimNode: - @pytest.fixture - def fast_config(self): - return SimConfig( - llm_latency_mean=0.01, - llm_latency_stddev=0.001, - llm_tokens_mean=10, - llm_tokens_stddev=2, - llm_token_rate=1000, - tool_latency_mean=0.01, - tool_latency_stddev=0.001, - tool_failure_rate=0.0, - max_tool_rounds=0, # no tool calls — fast turn - seed=42, - ) - - @pytest.fixture - def mock_broker(self): - broker = MagicMock() - broker.list_nodes.return_value = [] - return broker - - @pytest.fixture - def node(self, fast_config, mock_broker): - metrics = MetricsCollector() - return SimNode("test-node", mock_broker, fast_config, metrics) - - def test_handle_send_creates_workstream(self, node, mock_broker): - async def _test(): - msg = SendMessage(message="hello", auto_approve=True) - await node.handle_message(msg.to_json()) - - assert node.workstream_count == 1 - mock_broker.set_ws_owner.assert_called_once() - assert mock_broker.publish_outbound.call_count > 0 - - _run(_test()) - - def test_handle_send_reuses_existing_ws(self, node, mock_broker): - async def _test(): - msg1 = SendMessage(message="hello", auto_approve=True) - await node.handle_message(msg1.to_json()) - assert node.workstream_count == 1 - - ws_id = list(node._workstreams.keys())[0] - - msg2 = SendMessage(message="world", ws_id=ws_id, auto_approve=True) - await node.handle_message(msg2.to_json()) - assert node.workstream_count == 1 - - _run(_test()) - - def test_published_events_are_valid_protocol(self, node, mock_broker): - async def _test(): - msg = SendMessage(message="test", auto_approve=True) - await node.handle_message(msg.to_json()) - - for c in mock_broker.publish_outbound.call_args_list: - _channel, event_json = c[0] - event = OutboundEvent.from_json(event_json) - assert event.type != "" - - _run(_test()) - - def test_state_transitions(self, node, mock_broker): - async def _test(): - msg = SendMessage(message="test", auto_approve=True) - await node.handle_message(msg.to_json()) - - states = [] - for c in mock_broker.publish_outbound.call_args_list: - channel, event_json = c[0] - event = OutboundEvent.from_json(event_json) - if isinstance(event, StateChangeEvent): - states.append(event.state) - - assert "thinking" in states - assert "idle" in states - assert states.index("thinking") < states.index("idle") - - _run(_test()) - - def test_turn_complete_published(self, node, mock_broker): - async def _test(): - msg = SendMessage(message="test", auto_approve=True) - await node.handle_message(msg.to_json()) - - turn_completes = [ - OutboundEvent.from_json(c[0][1]) - for c in mock_broker.publish_outbound.call_args_list - if '"turn_complete"' in c[0][1] - ] - assert len(turn_completes) >= 1 - - _run(_test()) - - def test_close_workstream(self, node, mock_broker): - async def _test(): - msg = SendMessage(message="hello", auto_approve=True) - await node.handle_message(msg.to_json()) - ws_id = list(node._workstreams.keys())[0] - - from turnstone.mq.protocol import CloseWorkstreamMessage - - close_msg = CloseWorkstreamMessage(ws_id=ws_id) - await node.handle_message(close_msg.to_json()) - - assert node.workstream_count == 0 - mock_broker.del_ws_owner.assert_called_with(ws_id) - - _run(_test()) - - def test_stop_cleans_up(self, node, mock_broker): - # Add a fake workstream - node._workstreams["fake"] = MagicMock() - mock_broker.set_ws_owner("fake", "test-node") - - node.stop() - assert not node._running - assert node.workstream_count == 0 - mock_broker.del_ws_owner.assert_called() - - def test_heartbeat_once(self, node, mock_broker): - node.heartbeat_once() - mock_broker.register_node.assert_called_once() - args = mock_broker.register_node.call_args - assert args[0][0] == "test-node" - assert args[0][1]["sim"] is True - - -# --------------------------------------------------------------------------- -# SimWorkstream — state machine -# --------------------------------------------------------------------------- - - -class TestSimWorkstream: - @pytest.fixture - def fast_config(self): - return SimConfig( - llm_latency_mean=0.01, - llm_latency_stddev=0.001, - llm_tokens_mean=10, - llm_tokens_stddev=2, - llm_token_rate=1000, - tool_latency_mean=0.01, - tool_latency_stddev=0.001, - tool_failure_rate=0.0, - max_tool_rounds=0, - seed=42, - ) - - def test_turn_ends_in_idle(self, fast_config): - async def _test(): - broker = MagicMock() - metrics = MetricsCollector() - node = SimNode("test", broker, fast_config, metrics) - engine = SimEngine(fast_config) - ws = SimWorkstream("ws1", "test-ws", node, engine, fast_config) - - await ws.process_turn("hello", "cid-123") - assert ws.state == "idle" - - _run(_test()) - - def test_turn_records_metrics(self, fast_config): - async def _test(): - broker = MagicMock() - metrics = MetricsCollector() - node = SimNode("test", broker, fast_config, metrics) - engine = SimEngine(fast_config) - ws = SimWorkstream("ws1", "test-ws", node, engine, fast_config) - - await ws.process_turn("hello", "cid-123") - report = metrics.summary() - assert report["total_turns"] == 1 - assert report["turns_per_node"]["test"] == 1 - - _run(_test()) diff --git a/tests/test_tls_admin.py b/tests/test_tls_admin.py index db8b8b85..c1d1fff8 100644 --- a/tests/test_tls_admin.py +++ b/tests/test_tls_admin.py @@ -164,17 +164,6 @@ def test_cli_bootstrap_no_issue(tmp_path): # ── Config parsing ──────────────────────────────────────────────────────────── -def test_redis_tls_config_map(): - """Redis TLS keys are in the config map.""" - from turnstone.core.config import _CONFIG_MAP - - redis_map = _CONFIG_MAP["redis"] - assert "tls" in redis_map - assert "tls_ca" in redis_map - assert "tls_cert" in redis_map - assert "tls_key" in redis_map - - def test_database_ssl_config_map(): """Database SSL keys are in the config map.""" from turnstone.core.config import _CONFIG_MAP diff --git a/tests/test_tls_client.py b/tests/test_tls_client.py index 6ffba24e..827abdaa 100644 --- a/tests/test_tls_client.py +++ b/tests/test_tls_client.py @@ -78,21 +78,11 @@ async def test_ssl_contexts_none_before_init(): # ── Backward compatibility ─────────────────────────────────────────────────── -def test_bridge_tls_defaults(): - """Bridge with default TLS params works without changes.""" - from turnstone.mq.bridge import Bridge - - # Default: tls_verify=True, tls_cert=None — no mTLS - bridge = Bridge(server_url="http://localhost:8080") - assert bridge._tls_verify is True - assert bridge._tls_cert is None - - def test_collector_tls_defaults(): """Collector with default TLS params works without changes.""" from turnstone.console.collector import ClusterCollector - broker_mock = MagicMock() - collector = ClusterCollector(broker=broker_mock) + storage_mock = MagicMock() + collector = ClusterCollector(storage=storage_mock) # Should create httpx client without errors assert collector._http_client is not None diff --git a/tests/test_tool_policies_enforcement.py b/tests/test_tool_policies_enforcement.py index 1ff54920..2f3c4f60 100644 --- a/tests/test_tool_policies_enforcement.py +++ b/tests/test_tool_policies_enforcement.py @@ -1,4 +1,4 @@ -"""Tests for tool policy enforcement across CLI, bridge, and channel entry points.""" +"""Tests for tool policy enforcement in the CLI entry point.""" from __future__ import annotations @@ -83,92 +83,3 @@ class TestCLIPolicyEnforcement: # Should fall through to normal prompt (which we answered 'y') assert approved is True - - -# --------------------------------------------------------------------------- -# Bridge -# --------------------------------------------------------------------------- - - -class TestBridgePolicyEnforcement: - """Tool policies should be enforced in bridge _handle_approval().""" - - def _make_bridge(self): - from turnstone.mq.bridge import Bridge - - broker = MagicMock() - return Bridge( - server_url="http://localhost:8080", - broker=broker, - node_id="test-node", - approval_timeout=1, - ) - - def _approval_items(self, *tool_names: str) -> list[dict]: - return [ - {"func_name": name, "needs_approval": True, "approval_label": name} - for name in tool_names - ] - - def test_deny_policy_rejects_approval(self): - """A 'deny' policy should reject the approval.""" - bridge = self._make_bridge() - - with ( - patch( - "turnstone.core.policy.evaluate_tool_policies_batch", - return_value={"bash": "deny"}, - ), - patch( - "turnstone.core.storage._registry._storage", - new=MagicMock(), - ), - patch.object(bridge, "_api_approve") as mock_approve, - patch.object(bridge, "_publish_ws"), - ): - bridge._handle_approval("ws-1", {"items": self._approval_items("bash")}) - - mock_approve.assert_called_once() - assert mock_approve.call_args.kwargs.get("approved") is False - - def test_allow_policy_approves(self): - """An 'allow' policy should auto-approve.""" - bridge = self._make_bridge() - - with ( - patch( - "turnstone.core.policy.evaluate_tool_policies_batch", - return_value={"read_file": "allow"}, - ), - patch( - "turnstone.core.storage._registry.get_storage", - return_value=MagicMock(), - ), - patch.object(bridge, "_api_approve") as mock_approve, - patch.object(bridge, "_publish_ws"), - ): - bridge._handle_approval("ws-1", {"items": self._approval_items("read_file")}) - - mock_approve.assert_called_once() - assert mock_approve.call_args.kwargs.get("approved") is True - - def test_mixed_deny_rejects_batch(self): - """If any tool is denied, the whole batch is rejected.""" - bridge = self._make_bridge() - - with ( - patch( - "turnstone.core.policy.evaluate_tool_policies_batch", - return_value={"bash": "deny", "read_file": "allow"}, - ), - patch( - "turnstone.core.storage._registry._storage", - new=MagicMock(), - ), - patch.object(bridge, "_api_approve") as mock_approve, - patch.object(bridge, "_publish_ws"), - ): - bridge._handle_approval("ws-1", {"items": self._approval_items("bash", "read_file")}) - - mock_approve.assert_called_once() - assert mock_approve.call_args.kwargs.get("approved") is False diff --git a/turnstone/api/console_spec.py b/turnstone/api/console_spec.py index bb715469..940e0a83 100644 --- a/turnstone/api/console_spec.py +++ b/turnstone/api/console_spec.py @@ -162,7 +162,7 @@ CONSOLE_ENDPOINTS: list[EndpointSpec] = [ EndpointSpec( "/v1/api/cluster/workstreams/new", "POST", - "Create workstream via MQ dispatch", + "Create workstream via HTTP dispatch", request_model=ConsoleCreateWsRequest, response_model=ConsoleCreateWsResponse, error_codes=[400, 404, 503], diff --git a/turnstone/bootstrap.py b/turnstone/bootstrap.py index 95869ed2..e7a80c60 100644 --- a/turnstone/bootstrap.py +++ b/turnstone/bootstrap.py @@ -53,16 +53,14 @@ configure a Turnstone deployment interactively. ## About Turnstone Turnstone is a multi-node AI orchestration platform. A deployment consists of: - **Server** (turnstone-server): Web UI + chat workstreams + LLM interaction (port 8080) -- **Bridge** (turnstone-bridge): Redis-to-HTTP bridge for multi-node routing - **Console** (turnstone-console): Cluster dashboard + admin panel (port 8090) -- **Redis**: Message broker, pub/sub, node registry - **PostgreSQL** (production): Persistent database (dev can use SQLite) - **Channel** (optional): Discord/Slack gateway ## Deployment Profiles (compose.yaml) -- **Default** (no flag): redis + console only (infrastructure, good for running external servers) -- **Production** (`--profile production`): redis + 1 server + 1 bridge + console + PostgreSQL + channel (single node) -- **Cluster** (`--profile cluster`): 10-node server/bridge fleet + PostgreSQL + channel + console (multi-node) +- **Default** (no flag): console only (infrastructure, good for running external servers) +- **Production** (`--profile production`): 1 server + console + PostgreSQL + channel (single node) +- **Cluster** (`--profile cluster`): 10-node server fleet + PostgreSQL + channel + console (multi-node) - **ddgCluster** (`--profile ddgCluster`): Cluster + DuckDuckGo Search MCP sidecar (web search via MCP, no API key needed) ## Environment Variables (.env) @@ -84,10 +82,6 @@ For commercial providers (OpenAI, Anthropic-via-proxy), use the real key. - `POSTGRES_USER` — PostgreSQL username (default: turnstone) - `POSTGRES_PASSWORD` — PostgreSQL password (required for production/cluster) -### Redis -- `REDIS_PASSWORD` — Redis password (optional but recommended) -- `REDIS_PORT` — Redis port (default: 6379) - ### Authentication - `TURNSTONE_AUTH_ENABLED` — Enable auth (`true`/empty) - `TURNSTONE_JWT_SECRET` — JWT signing secret (required if auth enabled) @@ -121,7 +115,6 @@ The sidecar uses MCP streamable-http transport with DNS rebinding protection dis Safe search is disabled by default. ### Cluster -- `HEARTBEAT_TTL` — Bridge heartbeat TTL in seconds (default: 60) - `APPROVAL_TIMEOUT` — Tool approval timeout in seconds (default: 3600) ## Auth Setup Flow @@ -142,7 +135,7 @@ reasoning_effort, tool timeout, rate limiting, health probes, judge config, memo config, etc.) are configurable via the admin Settings tab in the console — no \ config.toml edits or restarts needed for most changes. These settings are stored in \ the database and apply cluster-wide. The `.env` file only needs bootstrap-critical \ -settings (database, Redis, auth, ports, API keys). Tell users they can fine-tune \ +settings (database, auth, ports, API keys). Tell users they can fine-tune \ model and behavioral settings after deployment through the admin panel. ## Built-in Roles @@ -170,7 +163,7 @@ Walk the user through setting up their deployment step by step: 4. **Database**: SQLite (dev/simple) vs PostgreSQL (production/cluster). \ PostgreSQL is required for cluster mode. 5. **Security**: Recommend enabling auth for any non-local deployment. \ -Use `generate_secret` for JWT secret, Redis password, auth token, and Postgres password. \ +Use `generate_secret` for JWT secret, auth token, and Postgres password. \ Ask for initial admin username and password. \ If the user's deployment will use an external identity provider (Okta, Azure AD, Google, etc.), \ offer to configure OIDC SSO. Ask for the issuer URL, client ID, and client secret. \ diff --git a/turnstone/channels/_config.py b/turnstone/channels/_config.py index 7d53970a..ff3f0884 100644 --- a/turnstone/channels/_config.py +++ b/turnstone/channels/_config.py @@ -13,11 +13,7 @@ class ChannelConfig: guild IDs, etc.). """ - redis_host: str = "localhost" - redis_port: int = 6379 - redis_db: int = 0 - redis_password: str | None = None - prefix: str = "turnstone" + server_url: str = "http://localhost:8080" model: str = "" auto_approve: bool = False auto_approve_tools: list[str] = field(default_factory=list) diff --git a/turnstone/channels/_http.py b/turnstone/channels/_http.py index a62e8dbd..ca876aef 100644 --- a/turnstone/channels/_http.py +++ b/turnstone/channels/_http.py @@ -1,7 +1,7 @@ """Lightweight HTTP server for the channel gateway. Runs alongside the channel adapters (Discord, etc.) to receive notification -requests from the bridge. Exposes ``POST /v1/api/notify`` and ``GET /health``. +requests from the server. Exposes ``POST /v1/api/notify`` and ``GET /health``. """ from __future__ import annotations diff --git a/turnstone/channels/_routing.py b/turnstone/channels/_routing.py index 12396933..5cc38134 100644 --- a/turnstone/channels/_routing.py +++ b/turnstone/channels/_routing.py @@ -1,28 +1,20 @@ """Channel router -- maps external channels/threads to turnstone workstreams. -:class:`ChannelRouter` uses the async Redis broker for MQ communication and -the storage backend for persistent channel-to-workstream mappings. +:class:`ChannelRouter` uses direct HTTP calls to the turnstone server API +and the storage backend for persistent channel-to-workstream mappings. """ from __future__ import annotations import asyncio -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any + +import httpx from turnstone.core.log import get_logger -from turnstone.mq.protocol import ( - ApproveMessage, - CreateWorkstreamMessage, - OutboundEvent, - PlanFeedbackMessage, - SendMessage, - WorkstreamClosedEvent, - WorkstreamCreatedEvent, -) if TYPE_CHECKING: from turnstone.core.storage import StorageBackend - from turnstone.mq.async_broker import AsyncRedisBroker log = get_logger(__name__) @@ -30,92 +22,75 @@ _WS_CREATE_TIMEOUT = 30.0 # seconds class ChannelRouter: - """Manage channel-to-workstream routing and MQ message dispatch. + """Manage channel-to-workstream routing via the server REST API. Parameters ---------- - broker: - An :class:`AsyncRedisBroker` used for pub/sub and queue operations. + server_url: + Base URL of the turnstone server (e.g. ``http://localhost:8080/v1``). storage: A :class:`StorageBackend` instance for persistent route lookups. All storage calls are synchronous and will be wrapped in :func:`asyncio.to_thread`. + api_token: + Optional bearer token for authenticating with the server API. """ def __init__( self, - broker: AsyncRedisBroker, + server_url: str, storage: StorageBackend, *, auto_approve: bool = False, auto_approve_tools: list[str] | None = None, skill: str = "", + api_token: str = "", ) -> None: - self._broker = broker + self._server_url = server_url.rstrip("/") self._storage = storage self._auto_approve = auto_approve self._auto_approve_tools: list[str] = auto_approve_tools or [] self._skill = skill - self._pending: dict[str, asyncio.Event] = {} - self._pending_results: dict[str, str] = {} - self._global_task: asyncio.Task[None] | None = None self._create_locks: dict[str, asyncio.Lock] = {} + headers: dict[str, str] = {} + if api_token: + headers["Authorization"] = f"Bearer {api_token}" + self._client = httpx.AsyncClient( + base_url=self._server_url, + headers=headers, + timeout=_WS_CREATE_TIMEOUT, + ) # -- lifecycle ----------------------------------------------------------- - async def start(self) -> None: - """Subscribe to global events for workstream lifecycle.""" - channel = f"{self._broker._prefix}:events:global" - await self._broker.subscribe(channel, self._on_global_event) - log.info("channel_router.started", channel=channel) + async def aclose(self) -> None: + """Close the underlying HTTP client.""" + await self._client.aclose() + log.info("channel_router.closed") - async def stop(self) -> None: - """Unsubscribe and clean up pending state.""" - channel = f"{self._broker._prefix}:events:global" - await self._broker.unsubscribe(channel) - # Wake any waiters so they don't hang forever. - for evt in self._pending.values(): - evt.set() - self._pending.clear() - self._pending_results.clear() - log.info("channel_router.stopped") + # -- internal helpers ---------------------------------------------------- - # -- event handler ------------------------------------------------------- + async def _post(self, path: str, body: dict[str, Any]) -> httpx.Response: + """POST JSON to the server and return the response.""" + resp = await self._client.post(path, json=body) + resp.raise_for_status() + return resp - async def _on_global_event(self, raw: str) -> None: - """Handle events on the global pub/sub channel. + async def _is_ws_alive(self, ws_id: str) -> bool: + """Check whether *ws_id* is a known workstream. - Exceptions are caught so the broker listener task stays alive. + Uses an O(1) storage lookup (primary-key query) instead of + fetching the full workstream list from the server. If the + workstream exists in the database it is considered alive. A + false positive (exists in DB but not loaded on any server node) + is harmless -- the subsequent ``send_message`` call will receive + a 404 and the adapter will handle reconnection. """ try: - event = OutboundEvent.from_json(raw) - - if isinstance(event, WorkstreamCreatedEvent): - cid = event.correlation_id - if cid in self._pending: - self._pending_results[cid] = event.ws_id - self._pending[cid].set() - log.debug( - "channel_router.ws_created", - ws_id=event.ws_id, - correlation_id=cid, - ) - - elif isinstance(event, WorkstreamClosedEvent): - ws_id = event.ws_id - route = await asyncio.to_thread(self._storage.get_channel_route_by_ws, ws_id) - if route: - # Don't delete the route — the workstream may have been evicted - # and the thread can reactivate it. Route cleanup only happens - # via explicit /close or delete_route(). - log.info( - "channel_router.ws_closed_route_kept", - ws_id=ws_id, - channel_type=route["channel_type"], - channel_id=route["channel_id"], - ) + resolved = await asyncio.to_thread(self._storage.resolve_workstream, ws_id) + return resolved is not None except Exception: - log.exception("channel_router.global_event_error") + return False # -- workstream management ----------------------------------------------- @@ -148,9 +123,8 @@ class ChannelRouter: self._storage.get_channel_route, channel_type, channel_id ) if route: - # Verify the workstream is still alive (owned by a bridge node). - owner = await self._broker.get_ws_owner(route["ws_id"]) - if owner: + # Verify the workstream is still alive on the server. + if await self._is_ws_alive(route["ws_id"]): return route["ws_id"], False # Workstream was evicted/closed — capture old ws_id for # resume, then remove the stale route. @@ -165,44 +139,37 @@ class ChannelRouter: channel_id=channel_id, ) - # 2. Create via MQ with atomic resume (reuse old ws_id directly). + # 2. Create via HTTP API with atomic resume. + # Note: auto_approve_tools is not passed here because the server's + # create endpoint does not accept it. Per-tool auto-approve is + # handled channel-side in the adapter's _should_auto_approve(). resume_ws = old_ws_id or "" - msg = CreateWorkstreamMessage( - name=name, - model=model, - initial_message="" if resume_ws else initial_message, - resume_ws=resume_ws, - auto_approve=self._auto_approve, - auto_approve_tools=list(self._auto_approve_tools), - skill=self._skill, - ) - cid = msg.correlation_id - waiter = asyncio.Event() - self._pending[cid] = waiter - - await self._broker.push_inbound(msg.to_json()) + body: dict[str, Any] = { + "name": name, + "model": model, + "resume_ws": resume_ws, + "skill": self._skill, + "auto_approve": self._auto_approve, + } log.info( "channel_router.creating_workstream", - correlation_id=cid, channel_type=channel_type, channel_id=channel_id, resume_ws=resume_ws or None, ) - try: - await asyncio.wait_for(waiter.wait(), timeout=_WS_CREATE_TIMEOUT) - except TimeoutError: - self._pending.pop(cid, None) - self._pending_results.pop(cid, None) - raise - - ws_id = self._pending_results.pop(cid, "") - self._pending.pop(cid, None) + resp = await self._post("/api/workstreams/new", body) + data = resp.json() + ws_id: str = data.get("ws_id", "") if not ws_id: msg_err = "workstream creation returned empty ws_id" raise RuntimeError(msg_err) + # 3. Send the initial message if this is a brand-new workstream. + if initial_message and not resume_ws: + await self._post("/api/send", {"ws_id": ws_id, "message": initial_message}) + # 4. Persist the route. await asyncio.to_thread( self._storage.create_channel_route, channel_type, channel_id, ws_id @@ -232,20 +199,10 @@ class ChannelRouter: # -- message dispatch ---------------------------------------------------- - async def send_message(self, ws_id: str, message: str) -> str: - """Push a :class:`SendMessage` to the broker. - - Returns the ``correlation_id`` of the submitted message. - """ - msg = SendMessage( - ws_id=ws_id, - message=message, - auto_approve=self._auto_approve, - auto_approve_tools=list(self._auto_approve_tools), - ) - await self._broker.push_inbound(msg.to_json()) - log.debug("channel_router.send_message", ws_id=ws_id, correlation_id=msg.correlation_id) - return msg.correlation_id + async def send_message(self, ws_id: str, message: str) -> None: + """Send a user message to a workstream via the server API.""" + await self._post("/api/send", {"ws_id": ws_id, "message": message}) + log.debug("channel_router.send_message", ws_id=ws_id) async def send_approval( self, @@ -255,15 +212,15 @@ class ChannelRouter: feedback: str = "", always: bool = False, ) -> None: - """Push an :class:`ApproveMessage` to the broker response queue.""" - msg = ApproveMessage( - ws_id=ws_id, - request_id=correlation_id, - approved=approved, - feedback=feedback or None, - always=always, - ) - await self._broker.push_response(correlation_id, msg.to_json()) + """Approve or deny a pending tool call via the server API.""" + body: dict[str, Any] = { + "ws_id": ws_id, + "approved": approved, + "always": always, + } + if feedback: + body["feedback"] = feedback + await self._post("/api/approve", body) log.debug( "channel_router.send_approval", ws_id=ws_id, @@ -272,13 +229,8 @@ class ChannelRouter: ) async def send_plan_feedback(self, ws_id: str, correlation_id: str, feedback: str) -> None: - """Push a :class:`PlanFeedbackMessage` to the broker response queue.""" - msg = PlanFeedbackMessage( - ws_id=ws_id, - request_id=correlation_id, - feedback=feedback, - ) - await self._broker.push_response(correlation_id, msg.to_json()) + """Respond to a plan review via the server API.""" + await self._post("/api/plan", {"ws_id": ws_id, "feedback": feedback}) log.debug( "channel_router.send_plan_feedback", ws_id=ws_id, @@ -298,3 +250,15 @@ class ChannelRouter: channel_id=channel_id, deleted=deleted, ) + + async def close_workstream(self, ws_id: str) -> None: + """Close a workstream via the server API.""" + try: + await self._post("/api/workstreams/close", {"ws_id": ws_id}) + log.info("channel_router.close_workstream", ws_id=ws_id) + except httpx.HTTPStatusError as exc: + log.warning( + "channel_router.close_workstream_failed", + ws_id=ws_id, + status=exc.response.status_code, + ) diff --git a/turnstone/channels/cli.py b/turnstone/channels/cli.py index 9aa368ff..21e7a8a5 100644 --- a/turnstone/channels/cli.py +++ b/turnstone/channels/cli.py @@ -1,8 +1,8 @@ """Unified channel gateway entry point. Launches one or more channel adapters (Discord, Slack, etc.) connected to -the turnstone cluster via Redis MQ. An HTTP server runs alongside for -inbound notification delivery from the server. +the turnstone server via HTTP. An HTTP server runs alongside for inbound +notification delivery from the server. Run as: ``turnstone-channel --discord-token $TURNSTONE_DISCORD_TOKEN`` """ @@ -15,17 +15,19 @@ import sys def main() -> None: - """Parse arguments, initialize storage and broker, and run adapters.""" + """Parse arguments, initialize storage, and run adapters.""" import argparse parser = argparse.ArgumentParser( description="turnstone channel gateway — bridges messaging platforms to the turnstone cluster" ) - # -- Redis --------------------------------------------------------------- - from turnstone.mq.broker import add_redis_args - - add_redis_args(parser) + # -- Server connection --------------------------------------------------- + parser.add_argument( + "--server-url", + default=os.environ.get("TURNSTONE_SERVER_URL", "http://localhost:8080"), + help="Turnstone server URL (default: $TURNSTONE_SERVER_URL or http://localhost:8080)", + ) # -- Discord ------------------------------------------------------------- parser.add_argument( @@ -115,10 +117,7 @@ def main() -> None: auth_token = args.auth_token jwt_secret = os.environ.get("TURNSTONE_JWT_SECRET", "").strip() - # -- Broker -------------------------------------------------------------- - from turnstone.mq.broker import async_broker_from_args - - broker = async_broker_from_args(args) + server_url: str = args.server_url # -- Adapter selection --------------------------------------------------- adapters_configured = False @@ -150,10 +149,7 @@ def main() -> None: ] config = DiscordConfig( - redis_host=args.redis_host, - redis_port=args.redis_port, - redis_db=args.redis_db, - redis_password=args.redis_password, + server_url=server_url, model=args.model, auto_approve=args.auto_approve, bot_token=args.discord_token, @@ -162,7 +158,7 @@ def main() -> None: ) storage = get_storage() - bot = TurnstoneBot(config, broker, storage) + bot = TurnstoneBot(config, server_url, storage, api_token=auth_token) adapters = {"discord": bot} # Create HTTP app for notification delivery @@ -178,6 +174,7 @@ def main() -> None: adapter="discord", guild_id=config.guild_id, http_port=args.http_port, + server_url=server_url, ) async def _run_all() -> None: diff --git a/turnstone/channels/discord/bot.py b/turnstone/channels/discord/bot.py index 4d9c05fe..c7fb66d3 100644 --- a/turnstone/channels/discord/bot.py +++ b/turnstone/channels/discord/bot.py @@ -1,29 +1,36 @@ """Discord bot adapter — connects Discord threads to turnstone workstreams. :class:`TurnstoneBot` extends ``discord.ext.commands.Bot`` and manages the -lifecycle of event subscriptions, streaming message edits, and interactive +lifecycle of SSE event subscriptions, streaming message edits, and interactive approval / plan-review views. + +Events are consumed from the server's per-workstream SSE endpoint +(``GET /v1/api/events?ws_id=X``) using httpx-sse, replacing the previous +Redis MQ pub/sub transport. """ from __future__ import annotations import asyncio +import contextlib +import json import time from dataclasses import dataclass, field from typing import TYPE_CHECKING +import httpx + from turnstone.channels._formatter import chunk_message from turnstone.channels._routing import ChannelRouter from turnstone.core.log import get_logger -from turnstone.mq.protocol import ( - ApprovalRequestEvent, +from turnstone.sdk.events import ( + ApproveRequestEvent, ContentEvent, ErrorEvent, IntentVerdictEvent, - OutboundEvent, PlanReviewEvent, - TurnCompleteEvent, - WorkstreamResumedEvent, + ServerEvent, + StreamEndEvent, ) if TYPE_CHECKING: @@ -32,10 +39,13 @@ if TYPE_CHECKING: from turnstone.channels.discord.config import DiscordConfig from turnstone.core.storage._protocol import StorageBackend - from turnstone.mq.async_broker import AsyncRedisBroker log = get_logger(__name__) +# SSE reconnection parameters +_SSE_RECONNECT_DELAY: float = 2.0 +_SSE_MAX_RECONNECT_DELAY: float = 30.0 + # --------------------------------------------------------------------------- # StreamingMessage helper @@ -117,10 +127,12 @@ class TurnstoneBot: ---------- config: Discord-specific configuration. - broker: - Async Redis broker for MQ communication. + server_url: + Base URL of the turnstone server API (e.g. ``http://localhost:8080/v1``). storage: Storage backend for persistent route / user lookups. + api_token: + Optional bearer token for authenticating with the server API. """ channel_type: str = "discord" @@ -129,40 +141,51 @@ class TurnstoneBot: def __init__( self, config: DiscordConfig, - broker: AsyncRedisBroker, + server_url: str, storage: StorageBackend, + *, + api_token: str = "", ) -> None: import discord from discord.ext import commands self.config = config - self.broker = broker + self._server_url = server_url.rstrip("/") + self._api_token = api_token self.storage = storage self.router = ChannelRouter( - broker, + server_url, storage, auto_approve=config.auto_approve, auto_approve_tools=list(config.auto_approve_tools), skill=config.skill, + api_token=api_token, ) self._subscribed_ws: set[str] = set() + self._sse_tasks: dict[str, asyncio.Task[None]] = {} self._streaming: dict[str, StreamingMessage] = {} # Track the Discord message containing the pending approval embed per # workstream so that IntentVerdictEvent can update it with LLM judge # results. self._pending_approval_msgs: dict[str, discord.Message] = {} - # Notification reply tracking: maps Discord message ID → + # Notification reply tracking: maps Discord message ID -> # (ws_id, target_discord_user_id) so that DM replies can be routed # back to the originating workstream. The target user ID is checked # on reply to prevent cross-user message injection. self._notify_ws_map: dict[int, tuple[str, str]] = {} - # Temporary DM forwarding: maps ws_id → (DM channel, target_user_id) + # Temporary DM forwarding: maps ws_id -> (DM channel, target_user_id) # for forwarding the workstream's next response back to the # notification reply DM. The target_user_id is carried so the # response message can be re-tracked for multi-turn DM conversations. self._notify_reply_channels: dict[str, tuple[discord.abc.Messageable, str]] = {} + # Shared HTTP client for SSE connections (long-lived, no timeout). + headers: dict[str, str] = {} + if api_token: + headers["Authorization"] = f"Bearer {api_token}" + self._http_client = httpx.AsyncClient(headers=headers, timeout=None) + intents = discord.Intents.default() intents.message_content = True @@ -189,9 +212,6 @@ class TurnstoneBot: from turnstone.channels.discord.cog import MessageCog from turnstone.channels.discord.views import ApprovalView, PlanReviewView - await self.broker.connect() - await self.router.start() - msg_cog = MessageCog(self._bot) await self._bot.add_cog(msg_cog._cog) @@ -223,7 +243,7 @@ class TurnstoneBot: """Re-subscribe to event channels for existing discord routes. Queries the storage backend for all channel routes of type ``discord`` - and subscribes to each workstream's event channel. + and opens SSE connections for each workstream. """ routes = await asyncio.to_thread(self.storage.list_channel_routes_by_type, "discord") for route in routes: @@ -247,23 +267,25 @@ class TurnstoneBot: ws_id: str, thread: discord.abc.Messageable, ) -> None: - """Subscribe to workstream events and dispatch them to *thread*.""" + """Subscribe to workstream events via SSE and dispatch them to *thread*.""" if ws_id in self._subscribed_ws: return - channel = f"{self.broker._prefix}:events:{ws_id}" - - async def _callback(raw: str) -> None: - await self._on_ws_event(ws_id, thread, raw) - - await self.broker.subscribe(channel, _callback) + task = asyncio.create_task( + self._sse_listener(ws_id, thread), + name=f"sse:{ws_id}", + ) + self._sse_tasks[ws_id] = task self._subscribed_ws.add(ws_id) log.info("discord.subscribed", ws_id=ws_id) async def unsubscribe_ws(self, ws_id: str) -> None: - """Cancel the subscription for *ws_id* and clean up streaming state.""" - channel = f"{self.broker._prefix}:events:{ws_id}" - await self.broker.unsubscribe(channel) + """Cancel the SSE listener for *ws_id* and clean up streaming state.""" + task = self._sse_tasks.pop(ws_id, None) + if task is not None: + task.cancel() + with contextlib.suppress(asyncio.CancelledError, Exception): + await task self._subscribed_ws.discard(ws_id) self._streaming.pop(ws_id, None) self._pending_approval_msgs.pop(ws_id, None) @@ -274,15 +296,63 @@ class TurnstoneBot: del self._notify_ws_map[mid] log.info("discord.unsubscribed", ws_id=ws_id) + # -- SSE listener -------------------------------------------------------- + + async def _sse_listener(self, ws_id: str, thread: discord.abc.Messageable) -> None: + """SSE listener task for a workstream. + + Connects to the server's per-workstream SSE endpoint, parses events + using :meth:`ServerEvent.from_dict`, and dispatches to + :meth:`_on_ws_event`. Reconnects with exponential backoff on + connection failures. + """ + import httpx_sse + + url = f"{self._server_url}/api/events" + delay = _SSE_RECONNECT_DELAY + + while True: + try: + async with httpx_sse.aconnect_sse( + self._http_client, "GET", url, params={"ws_id": ws_id} + ) as event_source: + delay = _SSE_RECONNECT_DELAY # reset on successful connect + async for sse in event_source.aiter_sse(): + if sse.event == "message" or not sse.event: + try: + data = json.loads(sse.data) + except json.JSONDecodeError: + log.debug( + "discord.sse_invalid_json", + ws_id=ws_id, + data=sse.data[:200], + ) + continue + event = ServerEvent.from_dict(data) + await self._on_ws_event(ws_id, thread, event) + if isinstance(event, StreamEndEvent): + return # clean end, do not reconnect + except httpx.RemoteProtocolError: + # Server closed connection (normal on stream_end or shutdown). + log.debug("discord.sse_remote_closed", ws_id=ws_id) + except asyncio.CancelledError: + return # unsubscribe or shutdown + except Exception: + log.warning("discord.sse_error", ws_id=ws_id, exc_info=True) + + # Exponential backoff before reconnecting. + await asyncio.sleep(delay) + delay = min(delay * 2, _SSE_MAX_RECONNECT_DELAY) + # -- event dispatch ------------------------------------------------------ async def _on_ws_event( self, ws_id: str, thread: discord.abc.Messageable, - raw: str, + event: ServerEvent, ) -> None: - """Handle an outbound event for a subscribed workstream.""" + """Handle a typed server event for a subscribed workstream.""" import discord from turnstone.channels._formatter import ( @@ -292,8 +362,6 @@ class TurnstoneBot: ) from turnstone.channels.discord.views import ApprovalView, PlanReviewView - event = OutboundEvent.from_json(raw) - if isinstance(event, ContentEvent): sm = self._streaming.get(ws_id) if sm is None: @@ -305,7 +373,7 @@ class TurnstoneBot: self._streaming[ws_id] = sm await sm.append(event.text) - elif isinstance(event, ApprovalRequestEvent): + elif isinstance(event, ApproveRequestEvent): # Evaluate admin tool policies before auto-approve. _policy_handled = False if self.storage is not None: @@ -328,7 +396,7 @@ class TurnstoneBot: denied = [n for n, v in verdicts.items() if v == "deny"] await self.router.send_approval( ws_id, - event.correlation_id, + "", approved=False, feedback=f"Blocked by tool policy: {', '.join(denied)}", ) @@ -339,7 +407,7 @@ class TurnstoneBot: elif all(verdicts.get(n) == "allow" for n in _tool_names): await self.router.send_approval( ws_id, - event.correlation_id, + "", approved=True, ) await thread.send("*Tool approved by policy.*") @@ -349,9 +417,12 @@ class TurnstoneBot: if not _policy_handled and ( self.config.auto_approve or self._should_auto_approve(event) ): - await self.router.send_approval(ws_id, event.correlation_id, approved=True) + # correlation_id is empty because the server's /api/approve + # endpoint resolves approvals by ws_id alone (one pending + # approval per workstream at a time). + await self.router.send_approval(ws_id, "", approved=True) await thread.send("*Tool auto-approved.*") - else: + elif not _policy_handled: text = format_approval_request(event.items) embed = discord.Embed( title="Tool Approval Required", @@ -368,7 +439,7 @@ class TurnstoneBot: value=format_verdict(verdict), inline=False, ) - embed.set_footer(text=f"{ws_id}|{event.correlation_id}") + embed.set_footer(text=f"{ws_id}|") msg = await thread.send(embed=embed, view=ApprovalView(self)._view) self._pending_approval_msgs[ws_id] = msg @@ -379,7 +450,7 @@ class TurnstoneBot: description=text, color=discord.Color.blue(), ) - embed.set_footer(text=f"{ws_id}|{event.correlation_id}") + embed.set_footer(text=f"{ws_id}|") await thread.send(embed=embed, view=PlanReviewView(self)._view) elif isinstance(event, IntentVerdictEvent): @@ -414,45 +485,37 @@ class TurnstoneBot: except Exception: log.debug("discord.verdict_embed_edit_failed", ws_id=ws_id) - elif isinstance(event, TurnCompleteEvent): + elif isinstance(event, StreamEndEvent): sm = self._streaming.pop(ws_id, None) if sm is not None: await sm.finalize() - elif event.content: - # Catch-up: content events were missed (race between global - # SSE and per-ws SSE) — send the full response directly. - for chunk in chunk_message(event.content, self.config.max_message_length): - await thread.send(chunk) - # Forward response to notification reply DM if active. + # Forward accumulated response to notification reply DM if active. dm_entry = self._notify_reply_channels.pop(ws_id, None) - if dm_entry is not None and event.content: - dm_channel, target_user_id = dm_entry - last_msg: discord.Message | None = None - for chunk in chunk_message(event.content, self.config.max_message_length): - try: - last_msg = await dm_channel.send(chunk) - except Exception: - log.debug("discord.notify_reply_dm_failed", ws_id=ws_id) - break - # Track the response message so the user can reply again - # for multi-turn DM conversations. - if last_msg is not None: - self._track_notification(last_msg.id, ws_id, target_user_id) + if dm_entry is not None and sm is not None: + content = "".join(sm._buffer) + if content: + dm_channel, target_user_id = dm_entry + last_msg: discord.Message | None = None + for dm_chunk in chunk_message(content, self.config.max_message_length): + try: + last_msg = await dm_channel.send(dm_chunk) + except Exception: + log.debug("discord.notify_reply_dm_failed", ws_id=ws_id) + break + # Track the response message so the user can reply again + # for multi-turn DM conversations. + if last_msg is not None: + self._track_notification(last_msg.id, ws_id, target_user_id) # Clean up pending approval message tracking. self._pending_approval_msgs.pop(ws_id, None) - elif isinstance(event, WorkstreamResumedEvent): - name = event.name or "previous workstream" - count = event.message_count - await thread.send(f"*Resumed: {name} ({count} messages restored)*") - elif isinstance(event, ErrorEvent): safe_msg = event.message[:500] if event.message else "An error occurred" await thread.send(f"**Error:** {safe_msg}") # -- helpers ------------------------------------------------------------- - def _should_auto_approve(self, event: ApprovalRequestEvent) -> bool: + def _should_auto_approve(self, event: ApproveRequestEvent) -> bool: """Return True if all tools in *event.items* are in the auto-approve list.""" allowed = self.config.auto_approve_tools if not allowed or not event.items: @@ -541,9 +604,9 @@ class TurnstoneBot: return msg_id_str async def stop(self) -> None: - """Disconnect the bot and clean up subscriptions.""" + """Disconnect the bot, cancel SSE tasks, and clean up.""" for ws_id in list(self._subscribed_ws): await self.unsubscribe_ws(ws_id) - await self.router.stop() - await self.broker.close() + await self.router.aclose() + await self._http_client.aclose() await self._bot.close() diff --git a/turnstone/channels/discord/cog.py b/turnstone/channels/discord/cog.py index 0f5ce9a6..0c720f92 100644 --- a/turnstone/channels/discord/cog.py +++ b/turnstone/channels/discord/cog.py @@ -10,7 +10,6 @@ import asyncio from typing import TYPE_CHECKING from turnstone.core.log import get_logger -from turnstone.mq.protocol import CloseWorkstreamMessage if TYPE_CHECKING: import discord @@ -184,10 +183,9 @@ class MessageCog: ) # Create workstream WITHOUT initial_message — subscribe to events - # first, then send the message. Sending initial_message through - # the bridge races with subscription: Redis pub/sub is fire-and- - # forget, so response events published before subscribe completes - # are silently dropped. + # first, then send the message. With SSE the event stream is + # reliable once connected, but we still subscribe first for + # consistency. ws_id, _is_new = await self.ts.router.get_or_create_workstream( channel_type="discord", channel_id=str(thread.id), @@ -253,7 +251,7 @@ class MessageCog: # Register the DM channel for response forwarding. The bot's # _on_ws_event handler will send the next turn's response here, # track the response for further replies, and clean up on - # TurnCompleteEvent. + # StreamEndEvent. self.ts._notify_reply_channels[ws_id] = (message.channel, target_user_id) log.info( @@ -446,9 +444,8 @@ class MessageCog: ws_id = route["ws_id"] - # Close via MQ. - msg = CloseWorkstreamMessage(ws_id=ws_id) - await self.ts.broker.push_inbound(msg.to_json()) + # Close via server API. + await self.ts.router.close_workstream(ws_id) # Delete route and unsubscribe. await self.ts.router.delete_route("discord", str(channel.id)) diff --git a/turnstone/channels/discord/views.py b/turnstone/channels/discord/views.py index 86e265b4..db20f41b 100644 --- a/turnstone/channels/discord/views.py +++ b/turnstone/channels/discord/views.py @@ -131,7 +131,7 @@ class ApprovalView: ws_id, correlation_id = parsed # Verify user is linked. Scope enforcement (approve) happens - # server-side when the bridge executes the tool approval. + # server-side when the tool approval is executed. user_id = await self.bot.router.resolve_user("discord", str(interaction.user.id)) if user_id is None: await interaction.response.send_message( diff --git a/turnstone/console/collector.py b/turnstone/console/collector.py index 50037dd0..3a180c8d 100644 --- a/turnstone/console/collector.py +++ b/turnstone/console/collector.py @@ -1,8 +1,7 @@ """Cluster state collector — aggregates data from all turnstone nodes. -Discovers nodes via Redis heartbeat keys, polls each node's /v1/api/dashboard -endpoint for workstream data, and subscribes to the cluster event channel -for real-time state changes. +Discovers nodes via the service registry (StorageBackend), polls each +node's /v1/api/dashboard endpoint for workstream data. """ from __future__ import annotations @@ -21,7 +20,7 @@ import httpx if TYPE_CHECKING: from turnstone.core.auth import ServiceTokenManager - from turnstone.mq.broker import RedisBroker + from turnstone.core.storage._protocol import StorageBackend log = logging.getLogger("turnstone.console.collector") @@ -42,18 +41,16 @@ class NodeSnapshot: class ClusterCollector: - """Aggregates cluster state from Redis and per-node HTTP APIs. + """Aggregates cluster state from the service registry and per-node HTTP APIs. - Three daemon threads: - 1. Event subscriber — real-time state changes from {prefix}:events:cluster - 2. Node discovery — scans heartbeat keys every ``discovery_interval`` seconds - 3. Poll loop — fetches /v1/api/dashboard from each node every ``poll_interval`` seconds + Two daemon threads: + 1. Node discovery — queries the service registry every ``discovery_interval`` seconds + 2. Poll loop — fetches /v1/api/dashboard from each node every ``poll_interval`` seconds """ def __init__( self, - broker: RedisBroker, - prefix: str = "turnstone", + storage: StorageBackend, poll_interval: float = 15.0, discovery_interval: float = 15.0, max_poll_workers: int = 200, @@ -63,8 +60,7 @@ class ClusterCollector: tls_verify: Any = True, tls_cert: tuple[str, str] | None = None, ): - self._broker = broker - self._prefix = prefix + self._storage = storage self._poll_interval = poll_interval self._discovery_interval = discovery_interval self._max_poll_workers = max_poll_workers @@ -121,7 +117,6 @@ class ClusterCollector: """Start background threads.""" self._running = True for target, name in [ - (self._event_loop, "console-events"), (self._discovery_loop, "console-discovery"), (self._poll_loop, "console-poll"), ]: @@ -137,74 +132,6 @@ class ClusterCollector: self._http_client.close() log.info("ClusterCollector stopped") - # -- event subscription -------------------------------------------------- - - def _event_loop(self) -> None: - """Subscribe to cluster events for real-time updates.""" - while self._running: - try: - self._broker.subscribe_cluster(self._on_cluster_event) - while self._running: - time.sleep(1) - except Exception: - log.exception("Cluster subscription error, reconnecting in 5s") - time.sleep(5) - - def _on_cluster_event(self, raw: str) -> None: - """Handle a cluster event from Redis pub/sub.""" - try: - data = json.loads(raw) - except (json.JSONDecodeError, TypeError): - return - - etype = data.get("type", "") - ws_id = data.get("ws_id", "") - node_id = data.get("node_id", "") - - with self._lock: - if etype == "cluster_state" and node_id in self._nodes: - node = self._nodes[node_id] - if ws_id in node.workstreams: - ws = node.workstreams[ws_id] - ws["state"] = data.get("state", ws.get("state", "idle")) - if "tokens" in data: - ws["tokens"] = data["tokens"] - if "context_ratio" in data: - ws["context_ratio"] = data["context_ratio"] - if "activity" in data: - ws["activity"] = data["activity"] - if "activity_state" in data: - ws["activity_state"] = data["activity_state"] - - elif etype == "ws_created" and node_id: - if node_id in self._nodes: - node = self._nodes[node_id] - node.workstreams[ws_id] = { - "id": ws_id, - "name": data.get("name", ""), - "state": "idle", - "node": node_id, - "server_url": node.server_url, - "title": data.get("title", ""), - "tokens": 0, - "context_ratio": 0.0, - "activity": "", - "activity_state": "", - "tool_calls": 0, - } - - elif etype == "ws_closed": - for node in self._nodes.values(): - node.workstreams.pop(ws_id, None) - - elif etype == "ws_rename": - for node in self._nodes.values(): - if ws_id in node.workstreams: - node.workstreams[ws_id]["name"] = data.get("name", "") - - # Fan out to SSE listeners - self._fanout(data) - def _fanout(self, event: dict[str, Any]) -> None: """Copy an event to all registered SSE listener queues.""" with self._listeners_lock: @@ -215,7 +142,7 @@ class ClusterCollector: # -- node discovery ------------------------------------------------------ def _discovery_loop(self) -> None: - """Periodically scan Redis for active nodes.""" + """Periodically scan the service registry for active nodes.""" while self._running: try: self._discover_nodes() @@ -224,29 +151,35 @@ class ClusterCollector: time.sleep(self._discovery_interval) def _discover_nodes(self) -> None: - """Scan heartbeat keys and update the node map.""" - active = self._broker.list_nodes() + """Query the service registry and update the node map.""" + raw_services = self._storage.list_services("server", max_age_seconds=120) active_ids = set() - pending_events = [] + pending_events: list[dict[str, Any]] = [] + with self._lock: - for meta in active: - nid = meta.get("node_id", "") + for svc in raw_services: + nid = svc.get("service_id", "") if not nid: continue active_ids.add(nid) + # Parse optional metadata JSON for max_ws, started + meta: dict[str, Any] = {} + raw_meta = svc.get("metadata", "") + if raw_meta: + with contextlib.suppress(json.JSONDecodeError, TypeError): + meta = json.loads(raw_meta) + url = svc.get("url", "") if nid not in self._nodes: self._nodes[nid] = NodeSnapshot( node_id=nid, - server_url=meta.get("server_url", ""), + server_url=url, started=meta.get("started", 0.0), max_ws=meta.get("max_ws", 10), ) pending_events.append({"type": "node_joined", "node_id": nid}) log.info("Discovered node: %s", nid) else: - self._nodes[nid].server_url = meta.get( - "server_url", self._nodes[nid].server_url - ) + self._nodes[nid].server_url = url or self._nodes[nid].server_url self._nodes[nid].max_ws = meta.get("max_ws", self._nodes[nid].max_ws) # Remove nodes whose heartbeats expired diff --git a/turnstone/console/scheduler.py b/turnstone/console/scheduler.py index b754a9d5..df25897d 100644 --- a/turnstone/console/scheduler.py +++ b/turnstone/console/scheduler.py @@ -1,25 +1,28 @@ """Background task scheduler for timed workstream dispatch. Runs as a daemon thread inside the console process. Checks for due tasks -every ``check_interval`` seconds and dispatches them as -``CreateWorkstreamMessage`` via the MQ broker. +every ``check_interval`` seconds and dispatches them via HTTP POST to +server nodes' ``/v1/api/workstreams/new`` endpoint. -Uses Redis ``SET NX EX`` for distributed locking in multi-console deployments. +Uses a ``system_settings`` row for distributed locking in multi-console +deployments. """ from __future__ import annotations +import json import threading import uuid from datetime import UTC, datetime from typing import TYPE_CHECKING, Any +import httpx import structlog if TYPE_CHECKING: from turnstone.console.collector import ClusterCollector + from turnstone.core.auth import ServiceTokenManager from turnstone.core.storage._protocol import StorageBackend - from turnstone.mq.broker import RedisBroker log = structlog.get_logger(__name__) @@ -44,18 +47,16 @@ class TaskScheduler: def __init__( self, - broker: RedisBroker, collector: ClusterCollector, storage: StorageBackend, - prefix: str = "turnstone", check_interval: float = 15.0, lock_ttl: int = 60, max_fan_out: int = 20, + api_token: str = "", + token_manager: ServiceTokenManager | None = None, ) -> None: - self._broker = broker self._collector = collector self._storage = storage - self._prefix = prefix self._check_interval = check_interval self._lock_ttl = lock_ttl self._max_fan_out = max_fan_out @@ -63,6 +64,10 @@ class TaskScheduler: self._thread: threading.Thread | None = None self._tick_count = 0 self._prune_every = 240 # ~1 hour at 15s intervals + self._lock_owner = uuid.uuid4().hex + self._api_token = api_token + self._token_manager = token_manager + self._http_client = httpx.Client(timeout=30) def start(self) -> None: """Start the scheduler daemon thread.""" @@ -76,6 +81,7 @@ class TaskScheduler: self._stop_event.set() if self._thread is not None: self._thread.join(timeout=5) + self._http_client.close() log.info("scheduler.stopped") def _loop(self) -> None: @@ -87,18 +93,71 @@ class TaskScheduler: log.exception("scheduler.tick_error") self._stop_event.wait(self._check_interval) - # Lua script for safe lock release — only delete if we still own the lock - _UNLOCK_SCRIPT = "if redis.call('get',KEYS[1])==ARGV[1] then return redis.call('del',KEYS[1]) else return 0 end" + def _try_acquire_lock(self) -> bool: + """Try to acquire the scheduler lock via system_settings. + + Uses a row with key ``scheduler_lock``. The value is a JSON + object ``{"owner": "", "acquired": ""}``. Another + instance's lock is considered expired when its timestamp is + older than ``_lock_ttl`` seconds. + + To reduce the TOCTOU window of a read-then-write approach, this + method writes unconditionally and reads back to verify ownership. + If two schedulers race, one write wins and the loser sees the + winner's value on read-back. The race window is microseconds + (write + read-back) which is acceptable for 15s tick intervals. + """ + now = datetime.now(UTC) + now_str = now.strftime("%Y-%m-%dT%H:%M:%S") + + # Check if another instance holds a non-expired lock before + # attempting to overwrite it. + existing = self._storage.get_system_setting("scheduler_lock") + if existing is not None: + try: + lock_data = json.loads(existing.get("value", "{}")) + except (json.JSONDecodeError, TypeError): + lock_data = {} + owner = lock_data.get("owner", "") + acquired_str = lock_data.get("acquired", "") + if owner != self._lock_owner and acquired_str: + try: + acquired_dt = datetime.strptime(acquired_str, "%Y-%m-%dT%H:%M:%S").replace( + tzinfo=UTC + ) + if (now - acquired_dt).total_seconds() < self._lock_ttl: + return False # Another instance holds a valid lock + except ValueError: + pass # Malformed timestamp — take the lock + + # Write our lock and read back to verify we won any concurrent race. + lock_value = json.dumps({"owner": self._lock_owner, "acquired": now_str}) + self._storage.upsert_system_setting("scheduler_lock", lock_value) + stored = self._storage.get_system_setting("scheduler_lock") + if stored is not None: + try: + data = json.loads(stored.get("value", "{}")) + except (json.JSONDecodeError, TypeError): + return False + return bool(data.get("owner") == self._lock_owner) + return False + + def _release_lock(self) -> None: + """Release the scheduler lock if we still own it.""" + existing = self._storage.get_system_setting("scheduler_lock") + if existing is not None: + try: + lock_data = json.loads(existing.get("value", "{}")) + except (json.JSONDecodeError, TypeError): + lock_data = {} + if lock_data.get("owner") == self._lock_owner: + self._storage.delete_system_setting("scheduler_lock") def _tick(self) -> None: """Single scheduler iteration: acquire lock, query due tasks, dispatch.""" now = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S") - # Distributed lock with unique owner — prevents releasing another instance's lock - lock_key = f"{self._prefix}:scheduler:lock" - lock_value = uuid.uuid4().hex - acquired = self._broker._redis.set(lock_key, lock_value, nx=True, ex=self._lock_ttl) - if not acquired: + if not self._try_acquire_lock(): return try: @@ -125,10 +184,7 @@ class TaskScheduler: except Exception: log.warning("scheduler.prune_audit_error", exc_info=True) finally: - # Only release our own lock (safe even if TTL expired and another took it) - self._broker._redis.eval( # type: ignore[no-untyped-call] - self._UNLOCK_SCRIPT, 1, lock_key, lock_value - ) + self._release_lock() def _dispatch_task(self, task: dict[str, Any], now: str) -> None: """Dispatch a single task as one or more CreateWorkstreamMessages.""" @@ -196,58 +252,73 @@ class TaskScheduler: raw = task.get("auto_approve_tools", "") return [t.strip() for t in raw.split(",") if t.strip()] - def _dispatch_to_node(self, task: dict[str, Any], node_id: str, now: str) -> None: - """Send a CreateWorkstreamMessage to a specific node.""" - from turnstone.mq.protocol import CreateWorkstreamMessage + def _auth_headers(self) -> dict[str, str]: + """Build auth headers for HTTP dispatch. - msg = CreateWorkstreamMessage( - name=task["name"], - model=task.get("model", ""), - target_node=node_id, - initial_message=task["initial_message"], - auto_approve=bool(task.get("auto_approve", 0)), - auto_approve_tools=self._parse_tools(task), - user_id=task.get("created_by", ""), - skill=task.get("skill", ""), - ) - self._broker.push_inbound(msg.to_json(), node_id=node_id) + Prefers a :class:`ServiceTokenManager` (auto-rotating JWT) over a + static API token. Returns an empty dict when neither is configured. + """ + if self._token_manager is not None: + return dict(self._token_manager.bearer_header) + if self._api_token: + return {"Authorization": f"Bearer {self._api_token}"} + return {} + + def _get_node_url(self, node_id: str) -> str: + """Resolve a node_id to its server URL via the collector.""" + detail = self._collector.get_node_detail(node_id) + if detail: + url: str = detail.get("server_url", "") + return url + return "" + + def _dispatch_to_node(self, task: dict[str, Any], node_id: str, now: str) -> None: + """POST to a specific node's /v1/api/workstreams/new endpoint.""" + server_url = self._get_node_url(node_id) + if not server_url: + self._record_failure(task, now, f"No URL for node {node_id}") + return + + correlation_id = uuid.uuid4().hex + body: dict[str, Any] = { + "name": task["name"], + "model": task.get("model", ""), + "initial_message": task["initial_message"], + "auto_approve": bool(task.get("auto_approve", 0)), + "auto_approve_tools": ",".join(self._parse_tools(task)), + "user_id": task.get("created_by", ""), + "skill": task.get("skill", ""), + } + try: + resp = self._http_client.post( + f"{server_url.rstrip('/')}/v1/api/workstreams/new", + json=body, + headers=self._auth_headers(), + ) + resp.raise_for_status() + except Exception: + self._record_failure(task, now, f"HTTP dispatch to {node_id} failed") + log.warning("scheduler.http_dispatch_failed", node_id=node_id, exc_info=True) + return self._storage.record_task_run( run_id=uuid.uuid4().hex, task_id=task["task_id"], node_id=node_id, ws_id="", - correlation_id=msg.correlation_id, + correlation_id=correlation_id, started=now, status="dispatched", error="", ) def _dispatch_to_pool(self, task: dict[str, Any], now: str) -> None: - """Send a CreateWorkstreamMessage to the shared pool queue.""" - from turnstone.mq.protocol import CreateWorkstreamMessage - - msg = CreateWorkstreamMessage( - name=task["name"], - model=task.get("model", ""), - initial_message=task["initial_message"], - auto_approve=bool(task.get("auto_approve", 0)), - auto_approve_tools=self._parse_tools(task), - user_id=task.get("created_by", ""), - skill=task.get("skill", ""), - ) - self._broker.push_inbound(msg.to_json()) - - self._storage.record_task_run( - run_id=uuid.uuid4().hex, - task_id=task["task_id"], - node_id="pool", - ws_id="", - correlation_id=msg.correlation_id, - started=now, - status="dispatched", - error="", - ) + """Dispatch to any available server node (pool mode).""" + node_id = _pick_best_node(self._collector) + if not node_id: + self._record_failure(task, now, "No reachable nodes for pool dispatch") + return + self._dispatch_to_node(task, node_id, now) def _record_failure(self, task: dict[str, Any], now: str, error: str) -> None: """Record a failed dispatch attempt.""" diff --git a/turnstone/console/server.py b/turnstone/console/server.py index 25e38e4a..32dd3ef9 100644 --- a/turnstone/console/server.py +++ b/turnstone/console/server.py @@ -4,7 +4,7 @@ Serves the cluster-level dashboard UI and provides REST/SSE APIs backed by the ClusterCollector. Uses Starlette/ASGI with uvicorn. Also provides: -- Workstream creation via MQ dispatch to target nodes +- Workstream creation via HTTP dispatch to target server nodes - Reverse proxy for server UIs so users only need console port access """ @@ -44,8 +44,6 @@ if TYPE_CHECKING: from starlette.requests import Request - from turnstone.mq.broker import RedisBroker - log = logging.getLogger("turnstone.console.server") # --------------------------------------------------------------------------- @@ -390,12 +388,12 @@ async def list_available_models(request: Request) -> JSONResponse: async def create_workstream(request: Request) -> JSONResponse: - """POST /v1/api/cluster/workstreams/new — create a workstream via MQ. + """POST /v1/api/cluster/workstreams/new — create a workstream via HTTP. Three targeting modes: - - ``node_id`` set to a specific node ID → directed to that node's queue + - ``node_id`` set to a specific node ID → POST to that node - ``node_id`` omitted or ``"auto"`` → console picks the node with most headroom - - ``node_id`` set to ``"pool"`` → pushed to the shared queue for any bridge + - ``node_id`` set to ``"pool"`` → console picks any available node """ from turnstone.core.web_helpers import read_json_or_400 @@ -403,7 +401,6 @@ async def create_workstream(request: Request) -> JSONResponse: if isinstance(body, JSONResponse): return body - broker: RedisBroker = request.app.state.broker collector: ClusterCollector = request.app.state.collector raw_node_id = body.get("node_id", "") @@ -445,30 +442,14 @@ async def create_workstream(request: Request) -> JSONResponse: skill = raw_skill[:256] resume_ws = raw_resume_ws[:64] - from turnstone.mq.protocol import CreateWorkstreamMessage - auth = getattr(getattr(request, "state", None), "auth_result", None) uid: str = getattr(auth, "user_id", "") or "" - # General pool — push to shared queue, any bridge picks it up + # Pool — pick any available node if node_id == "pool": - msg = CreateWorkstreamMessage( - name=name, - model=model, - initial_message=initial_message, - skill=skill, - resume_ws=resume_ws, - user_id=uid, - ) - broker.push_inbound(msg.to_json()) - log.debug("Pool dispatch: correlation_id=%s name=%r", msg.correlation_id, name) - return JSONResponse( - { - "status": "ok", - "correlation_id": msg.correlation_id, - "target_node": "pool", - } - ) + node_id = _pick_best_node(collector) + if not node_id: + return JSONResponse({"error": "No reachable nodes available"}, status_code=503) # Auto-select node by most available capacity if not node_id or node_id == "auto": @@ -476,26 +457,41 @@ async def create_workstream(request: Request) -> JSONResponse: if not node_id: return JSONResponse({"error": "No reachable nodes available"}, status_code=503) - # Validate node exists + # Validate node exists and get its URL detail = collector.get_node_detail(node_id) if not detail: return JSONResponse({"error": "Node not found"}, status_code=404) - msg = CreateWorkstreamMessage( - name=name, - model=model, - target_node=node_id, - initial_message=initial_message, - skill=skill, - resume_ws=resume_ws, - user_id=uid, - ) - broker.push_inbound(msg.to_json(), node_id=node_id) + server_url = detail.get("server_url", "") + if not server_url: + return JSONResponse({"error": "Node has no URL"}, status_code=502) + + ws_body = { + "name": name, + "model": model, + "initial_message": initial_message, + "skill": skill, + "resume_ws": resume_ws, + "user_id": uid, + } + + client: httpx.AsyncClient = request.app.state.proxy_client + headers = _proxy_auth_headers(request) + try: + resp = await client.post( + f"{server_url.rstrip('/')}/v1/api/workstreams/new", + json=ws_body, + headers=headers, + ) + resp.raise_for_status() + except httpx.HTTPError as exc: + log.warning("Workstream dispatch to %s failed: %s", node_id, exc) + return JSONResponse({"error": f"Dispatch to node {node_id} failed"}, status_code=502) return JSONResponse( { "status": "ok", - "correlation_id": msg.correlation_id, + "correlation_id": resp.json().get("ws_id", ""), "target_node": node_id, } ) @@ -878,7 +874,6 @@ async def _lifespan(app: Starlette) -> AsyncGenerator[None, None]: await app.state.proxy_sse_client.aclose() await app.state.proxy_client.aclose() app.state.collector.stop() - app.state.broker.close() # --------------------------------------------------------------------------- @@ -5337,7 +5332,6 @@ def _seed_config_from_env(config_store: Any, storage: Any) -> None: def create_app( *, collector: ClusterCollector, - broker: RedisBroker, auth_config: Any, jwt_secret: str = "", auth_storage: Any = None, @@ -5638,7 +5632,6 @@ def create_app( lifespan=_lifespan, ) app.state.collector = collector - app.state.broker = broker app.state.auth_config = auth_config app.state.jwt_secret = jwt_secret app.state.auth_storage = auth_storage @@ -5670,9 +5663,10 @@ def create_app( from turnstone.console.scheduler import TaskScheduler scheduler = TaskScheduler( - broker=broker, collector=collector, storage=auth_storage, + api_token=proxy_auth_token, + token_manager=proxy_token_mgr, ) app.state.scheduler = scheduler else: @@ -5703,9 +5697,8 @@ def main() -> None: formatter_class=argparse.RawDescriptionHelpFormatter, epilog=textwrap.dedent("""\ Examples: - turnstone-console # default Redis on localhost + turnstone-console # default settings turnstone-console --port 9090 # custom port - turnstone-console --redis-host redis.internal # remote Redis """), ) parser.add_argument( @@ -5719,9 +5712,6 @@ def main() -> None: default=8090, help="Port to listen on (default: 8090)", ) - from turnstone.mq.broker import add_redis_args - - add_redis_args(parser) parser.add_argument( "--poll-interval", type=float, @@ -5740,16 +5730,44 @@ def main() -> None: from turnstone.core.config import add_config_arg, apply_config add_config_arg(parser) - apply_config(parser, ["console", "redis", "auth"]) + apply_config(parser, ["console", "auth"]) args = parser.parse_args() from turnstone.core.log import configure_logging_from_args configure_logging_from_args(args, "console") - from turnstone.mq.broker import broker_from_args + from turnstone.core.auth import load_auth_config, load_jwt_secret - broker = broker_from_args(args) + auth_config = load_auth_config() + jwt_secret = load_jwt_secret() if auth_config.enabled else "" + + # Initialize storage early — the collector needs it for service discovery. + auth_storage = None + try: + from turnstone.core.storage import init_storage + + db_backend = os.environ.get("TURNSTONE_DB_BACKEND", "sqlite") + db_url = os.environ.get("TURNSTONE_DB_URL", "") + db_path = os.environ.get("TURNSTONE_DB_PATH", "") + auth_storage = init_storage( + db_backend, + path=db_path, + url=db_url, + sslmode=os.environ.get("TURNSTONE_DB_SSLMODE", ""), + sslrootcert=os.environ.get("TURNSTONE_DB_SSLROOTCERT", ""), + sslcert=os.environ.get("TURNSTONE_DB_SSLCERT", ""), + sslkey=os.environ.get("TURNSTONE_DB_SSLKEY", ""), + ) + except Exception: + log.info("Console storage not available — admin API disabled, JWT-only auth") + + if auth_storage is None: + log.error( + "Storage backend is required for the console (service discovery). " + "Set TURNSTONE_DB_PATH or TURNSTONE_DB_URL." + ) + raise SystemExit(1) # If no explicit auth token is provided, use a ServiceTokenManager # so collector JWTs auto-rotate. A shared JWT secret is required for @@ -5778,7 +5796,7 @@ def main() -> None: log.info("console.collector_token_manager_created") collector = ClusterCollector( - broker=broker, + storage=auth_storage, poll_interval=args.poll_interval, auth_token=collector_token if collector_token_mgr is None else "", token_manager=collector_token_mgr, @@ -5787,31 +5805,6 @@ def main() -> None: _load_static() - from turnstone.core.auth import load_auth_config, load_jwt_secret - - auth_config = load_auth_config() - jwt_secret = load_jwt_secret() if auth_config.enabled else "" - - # Initialize storage for user/token management (optional — requires DB config) - auth_storage = None - try: - from turnstone.core.storage import init_storage - - db_backend = os.environ.get("TURNSTONE_DB_BACKEND", "sqlite") - db_url = os.environ.get("TURNSTONE_DB_URL", "") - db_path = os.environ.get("TURNSTONE_DB_PATH", "") - auth_storage = init_storage( - db_backend, - path=db_path, - url=db_url, - sslmode=os.environ.get("TURNSTONE_DB_SSLMODE", ""), - sslrootcert=os.environ.get("TURNSTONE_DB_SSLROOTCERT", ""), - sslcert=os.environ.get("TURNSTONE_DB_SSLCERT", ""), - sslkey=os.environ.get("TURNSTONE_DB_SSLKEY", ""), - ) - except Exception: - log.info("Console storage not available — admin API disabled, JWT-only auth") - # If no explicit auth token is provided, use a ServiceTokenManager # so proxy JWTs auto-rotate. proxy_token = args.auth_token @@ -5897,7 +5890,6 @@ def main() -> None: app = create_app( collector=collector, - broker=broker, auth_config=auth_config, jwt_secret=jwt_secret, auth_storage=auth_storage, diff --git a/turnstone/core/config.py b/turnstone/core/config.py index d89e5048..71671829 100644 --- a/turnstone/core/config.py +++ b/turnstone/core/config.py @@ -113,23 +113,6 @@ _CONFIG_MAP: dict[str, dict[str, str]] = { "workstream_idle_timeout": "workstream_idle_timeout", "max_workstreams": "max_workstreams", }, - "bridge": { - "server_url": "server_url", - "node_id": "node_id", - "approval_timeout": "approval_timeout", - "heartbeat_ttl": "heartbeat_ttl", - "log_level": "log_level", - }, - "redis": { - "host": "redis_host", - "port": "redis_port", - "password": "redis_password", - "db": "redis_db", - "tls": "redis_tls", - "tls_ca": "redis_tls_ca", - "tls_cert": "redis_tls_cert", - "tls_key": "redis_tls_key", - }, "console": { "host": "host", "port": "port", diff --git a/turnstone/core/settings_registry.py b/turnstone/core/settings_registry.py index 6569dc3c..b1b3b3bc 100644 --- a/turnstone/core/settings_registry.py +++ b/turnstone/core/settings_registry.py @@ -2,7 +2,7 @@ Every setting that can be stored in the ``system_settings`` table must have an entry here. Unknown keys are rejected at the API boundary. -Bootstrap settings (database, Redis, auth, server/console bind) are +Bootstrap settings (database, auth, server/console bind) are excluded — they are needed before storage is available. """ @@ -537,7 +537,7 @@ def _build_registry() -> dict[str, SettingDef]: "tls", restart_required=True, help="When enabled, the console runs an internal Certificate Authority and " - "ACME server. All cluster services (servers, bridge, channels) auto-provision " + "ACME server. All cluster services (servers, channels) auto-provision " "short-lived certificates for mutual TLS. Requires lacme: pip install turnstone[tls]", ), SettingDef( @@ -562,9 +562,7 @@ BOOTSTRAP_SECTIONS: frozenset[str] = frozenset( { "api", "database", - "redis", "auth", - "bridge", "console", }, ) diff --git a/turnstone/core/workstream.py b/turnstone/core/workstream.py index d34b71aa..4e0caaff 100644 --- a/turnstone/core/workstream.py +++ b/turnstone/core/workstream.py @@ -132,6 +132,7 @@ class WorkstreamManager: skill: str | None = None, skill_id: str = "", skill_version: int = 0, + ws_id: str = "", ) -> Workstream: """Create a new workstream. Returns the new ws. @@ -145,6 +146,8 @@ class WorkstreamManager: skill: Optional skill name. skill_id: Template ID of the skill (for lineage tracking). skill_version: Version of the skill at creation time. + ws_id: Optional workstream ID. If non-empty, used as-is instead of + generating a new UUID. """ # Fast-fail capacity check (avoids expensive ChatSession creation when full). first_evicted: Workstream | None = None @@ -164,7 +167,7 @@ class WorkstreamManager: # Create workstream and ChatSession outside the lock (construction is # expensive — involves LLM client setup and DB writes). - ws = Workstream(name=name) + ws = Workstream(id=ws_id, name=name) if ws_id else Workstream(name=name) if ui_factory: ws.ui = ui_factory(ws.id) ws.session = self._session_factory(ws.ui, model, ws.id, skill=skill) diff --git a/turnstone/mq/__init__.py b/turnstone/mq/__init__.py deleted file mode 100644 index e25ba2c6..00000000 --- a/turnstone/mq/__init__.py +++ /dev/null @@ -1,26 +0,0 @@ -"""Message queue integration for turnstone. - -Provides a bridge service (turnstone-bridge) that connects message queues to the -turnstone-server HTTP API, and a client library for external systems to publish -commands and subscribe to progress. -""" - -from turnstone.mq.broker import MessageBroker, RedisBroker -from turnstone.mq.client import TurnResult, TurnstoneClient - -__all__ = [ - "AsyncRedisBroker", - "MessageBroker", - "RedisBroker", - "TurnstoneClient", - "TurnResult", -] - - -def __getattr__(name: str) -> object: - if name == "AsyncRedisBroker": - from turnstone.mq.async_broker import AsyncRedisBroker - - return AsyncRedisBroker - msg = f"module {__name__!r} has no attribute {name!r}" - raise AttributeError(msg) diff --git a/turnstone/mq/async_broker.py b/turnstone/mq/async_broker.py deleted file mode 100644 index dc4989a7..00000000 --- a/turnstone/mq/async_broker.py +++ /dev/null @@ -1,336 +0,0 @@ -"""Async Redis message broker. - -Provides :class:`AsyncRedisBroker`, an asyncio-native counterpart to -:class:`~turnstone.mq.broker.RedisBroker`. Uses ``redis.asyncio`` for all I/O -and manages pub/sub listeners as :class:`asyncio.Task` instances. -""" - -from __future__ import annotations - -import asyncio -import contextlib -import json -from typing import TYPE_CHECKING, Any - -if TYPE_CHECKING: - from collections.abc import Callable - - import redis.asyncio as _aredis_t - - -class AsyncRedisBroker: - """Async Redis-backed message broker using lists (queues) and pub/sub. - - This is the asyncio equivalent of :class:`~turnstone.mq.broker.RedisBroker`. - All methods are coroutines and must be awaited. - - Queue keys: - ``{prefix}:inbound`` — shared inbound command queue - ``{prefix}:inbound:{node_id}`` — per-node directed queue - ``{prefix}:resp:{request_id}`` — per-request response queues - - Routing keys: - ``{prefix}:ws:{ws_id}`` — workstream ownership (string) - ``{prefix}:node:{node_id}`` — node heartbeat + metadata (string/JSON) - - Pub/sub channels: - ``{prefix}:events:global`` — global event channel - ``{prefix}:events:{ws_id}`` — per-workstream event channel - ``{prefix}:events:cluster`` — cluster-wide state changes - """ - - def __init__( - self, - host: str = "localhost", - port: int = 6379, - db: int = 0, - prefix: str = "turnstone", - password: str | None = None, - response_ttl: int = 600, - ssl: bool = False, - ssl_ca_certs: str | None = None, - ssl_certfile: str | None = None, - ssl_keyfile: str | None = None, - ) -> None: - self._host = host - self._port = port - self._db = db - self._password = password - self._prefix = prefix - self._response_ttl = response_ttl - self._ssl_kwargs: dict[str, Any] = {} - if ssl: - self._ssl_kwargs["ssl"] = True - if ssl_ca_certs: - self._ssl_kwargs["ssl_ca_certs"] = ssl_ca_certs - if ssl_certfile: - self._ssl_kwargs["ssl_certfile"] = ssl_certfile - if ssl_keyfile: - self._ssl_kwargs["ssl_keyfile"] = ssl_keyfile - self._redis: _aredis_t.Redis[str] | None = None - self._pubsub: _aredis_t.client.PubSub | None = None - self._tasks: dict[str, asyncio.Task[None]] = {} - self._callbacks: dict[str, Callable[[str], Any]] = {} - self._queues: dict[str, asyncio.Queue[str]] = {} - self._workers: dict[str, asyncio.Task[None]] = {} - self._listener_task: asyncio.Task[None] | None = None - - # -- connection ---------------------------------------------------------- - - async def connect(self) -> None: - """Create the async Redis connection. - - This is called lazily before first use if the connection has not yet - been established. - """ - if self._redis is not None: - return - - import redis.asyncio as aioredis - - self._redis = aioredis.Redis( - host=self._host, - port=self._port, - db=self._db, - password=self._password, - decode_responses=True, - retry_on_timeout=True, - **self._ssl_kwargs, - max_connections=200, - ) - self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True) - - async def _ensure_connected(self) -> None: - """Ensure the Redis connection is established.""" - if self._redis is None: - await self.connect() - - @property - def _r(self) -> _aredis_t.Redis[str]: - """Return the Redis client, assuming it is connected.""" - if self._redis is None: - msg = "Broker not connected — call connect() first" - raise RuntimeError(msg) - return self._redis - - @property - def _ps(self) -> _aredis_t.client.PubSub: - """Return the pub/sub client, assuming it is connected.""" - if self._pubsub is None: - msg = "Broker not connected — call connect() first" - raise RuntimeError(msg) - return self._pubsub - - # -- inbound queue ------------------------------------------------------- - - async def push_inbound(self, message: str, node_id: str = "") -> None: - """Push a message onto the inbound queue. - - If *node_id* is set, pushes to the per-node queue for directed - routing. Otherwise pushes to the shared queue. - """ - await self._ensure_connected() - if node_id: - await self._r.rpush(f"{self._prefix}:inbound:{node_id}", message) - else: - await self._r.rpush(f"{self._prefix}:inbound", message) - - # -- outbound pub/sub ---------------------------------------------------- - - async def publish_outbound(self, channel: str, event: str) -> None: - """Publish an event to an outbound channel.""" - await self._ensure_connected() - await self._r.publish(channel, event) - - async def subscribe(self, channel: str, callback: Callable[[str], Any]) -> None: - """Subscribe to a pub/sub channel. - - The *callback* receives the message string for each published event. - It may be a regular function or an async coroutine. - - All subscriptions share a single listener task that dispatches - messages to the correct callback based on the channel name. - """ - await self._ensure_connected() - await self._ps.subscribe(channel) - self._callbacks[channel] = callback - - # Per-channel queue + worker ensures ordered delivery within a channel - # while allowing different channels to process concurrently. - q: asyncio.Queue[str] = asyncio.Queue() - self._queues[channel] = q - self._workers[channel] = asyncio.create_task(self._channel_worker(channel, q)) - - # Start the shared listener task if not already running. - if self._listener_task is None or self._listener_task.done(): - self._listener_task = asyncio.create_task(self._dispatch_loop()) - - async def _dispatch_loop(self) -> None: - """Single listener that routes pub/sub messages to per-channel queues. - - Each channel has its own queue + worker task, ensuring ordered - delivery within a channel while allowing different channels to - process concurrently. - """ - import logging - - _log = logging.getLogger("turnstone.mq.async_broker") - try: - while self._callbacks: - msg = await self._ps.get_message( - ignore_subscribe_messages=True, - timeout=0.1, - ) - if msg is None: - # Yield control so cancellation can be delivered. - await asyncio.sleep(0) - continue - if msg["type"] == "message": - ch = msg.get("channel", "") - q = self._queues.get(ch) - if q is not None: - q.put_nowait(msg["data"]) - _log.debug("Dispatch loop exiting — no active callbacks") - except asyncio.CancelledError: - return - - async def _channel_worker(self, channel: str, q: asyncio.Queue[str]) -> None: - """Process messages for a single channel sequentially.""" - import logging - - _log = logging.getLogger("turnstone.mq.async_broker") - try: - while True: - data = await q.get() - cb = self._callbacks.get(channel) - if cb is not None: - try: - result = cb(data) - if asyncio.iscoroutine(result): - await result - except Exception: - _log.exception("Listener callback error on %s", channel) - except asyncio.CancelledError: - return - - async def unsubscribe(self, channel: str) -> None: - """Unsubscribe from a channel and cancel its worker.""" - await self._ensure_connected() - self._callbacks.pop(channel, None) - self._queues.pop(channel, None) - worker = self._workers.pop(channel, None) - if worker is not None: - worker.cancel() - with contextlib.suppress(asyncio.CancelledError): - await worker - # Legacy per-channel task cleanup (in case any remain). - task = self._tasks.pop(channel, None) - if task is not None: - task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await task - await self._ps.unsubscribe(channel) - - # -- response queues ----------------------------------------------------- - - async def push_response(self, queue_name: str, message: str) -> None: - """Push a response onto a named response queue.""" - await self._ensure_connected() - key = f"{self._prefix}:resp:{queue_name}" - await self._r.rpush(key, message) - await self._r.expire(key, self._response_ttl) - - async def pop_response(self, queue_name: str, timeout: float = 300) -> str | None: - """Pop from a named response queue. Returns ``None`` on timeout.""" - await self._ensure_connected() - key = f"{self._prefix}:resp:{queue_name}" - result = await self._r.blpop(key, timeout=int(timeout)) - return result[1] if result else None - - # -- routing primitives -------------------------------------------------- - - async def get_ws_owner(self, ws_id: str) -> str | None: - """Look up the node that owns a workstream.""" - await self._ensure_connected() - return await self._r.get(f"{self._prefix}:ws:{ws_id}") - - async def set_ws_owner(self, ws_id: str, node_id: str, ttl: int = 0) -> None: - """Register which node owns a workstream.""" - await self._ensure_connected() - key = f"{self._prefix}:ws:{ws_id}" - if ttl > 0: - await self._r.set(key, node_id, ex=ttl) - else: - await self._r.set(key, node_id) - - async def del_ws_owner(self, ws_id: str) -> None: - """Remove workstream ownership.""" - await self._ensure_connected() - await self._r.delete(f"{self._prefix}:ws:{ws_id}") - - async def register_node(self, node_id: str, metadata: dict[str, Any], ttl: int = 60) -> None: - """Register or refresh a node's heartbeat with metadata.""" - await self._ensure_connected() - key = f"{self._prefix}:node:{node_id}" - await self._r.set(key, json.dumps(metadata), ex=ttl) - - async def list_nodes(self) -> list[dict[str, Any]]: - """List all active nodes (those with unexpired heartbeats).""" - await self._ensure_connected() - pattern = f"{self._prefix}:node:*" - prefix_len = len(f"{self._prefix}:node:") - # Collect all keys first, then batch-fetch with MGET to avoid - # N+1 round-trips (1 GET per node). - keys: list[str] = [] - async for key in self._r.scan_iter(match=pattern, count=100): - keys.append(key) - if not keys: - return [] - values = await self._r.mget(keys) - nodes: list[dict[str, Any]] = [] - for key, raw in zip(keys, values, strict=True): - if raw: - try: - meta: dict[str, Any] = json.loads(raw) - except json.JSONDecodeError: - meta = {} - meta["node_id"] = key[prefix_len:] - nodes.append(meta) - return nodes - - # -- lifecycle ----------------------------------------------------------- - - async def close(self) -> None: - """Cancel all listener tasks and close the Redis connection.""" - self._callbacks.clear() - self._queues.clear() - # Cancel per-channel workers. - for worker in self._workers.values(): - worker.cancel() - for worker in self._workers.values(): - with contextlib.suppress(asyncio.CancelledError): - await worker - self._workers.clear() - # Cancel the shared dispatch loop. - if self._listener_task is not None: - if not self._listener_task.done(): - self._listener_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await self._listener_task - self._listener_task = None - # Legacy per-channel tasks. - for task in self._tasks.values(): - task.cancel() - for task in self._tasks.values(): - with contextlib.suppress(asyncio.CancelledError): - await task - self._tasks.clear() - - if self._pubsub is not None: - with contextlib.suppress(Exception): - await self._pubsub.close() - self._pubsub = None - - if self._redis is not None: - await self._redis.close() - self._redis = None diff --git a/turnstone/mq/bridge.py b/turnstone/mq/bridge.py deleted file mode 100644 index 3329d411..00000000 --- a/turnstone/mq/bridge.py +++ /dev/null @@ -1,1194 +0,0 @@ -"""Bridge service — connects message queues to turnstone-server's HTTP API. - -The bridge (listener+speaker) reads commands from an inbound queue, drives -workstreams on the turnstone-server via HTTP, consumes SSE for progress, and -publishes events to outbound pub/sub channels. - -Run as: ``turnstone-bridge --server-url http://localhost:8080`` -""" - -from __future__ import annotations - -import json -import logging -import os -import threading -import time -import uuid -from typing import TYPE_CHECKING, Any - -import httpx - -from turnstone.mq.broker import RedisBroker -from turnstone.mq.protocol import ( - AckEvent, - ApprovalRequestEvent, - ClusterStateEvent, - ContentEvent, - ErrorEvent, - HealthResponseEvent, - InboundMessage, - InfoEvent, - IntentVerdictEvent, - NodeListEvent, - OutboundEvent, - OutputWarningEvent, - PlanReviewEvent, - ReasoningEvent, - StateChangeEvent, - StatusEvent, - StreamEndEvent, - ToolInfoEvent, - ToolOutputChunkEvent, - ToolResultEvent, - TurnCompleteEvent, - WorkstreamClosedEvent, - WorkstreamCreatedEvent, - WorkstreamListEvent, - WorkstreamRenameEvent, - WorkstreamResumedEvent, -) - -if TYPE_CHECKING: - from collections.abc import Callable, Iterator - -log = logging.getLogger("turnstone.mq.bridge") - -# Server's default safe tools (auto-approved without user confirmation) -DEFAULT_SAFE_TOOLS = frozenset(["read_file", "search", "man", "memory", "recall"]) - - -class Bridge: - """Connects a message broker to turnstone-server's HTTP API. - - Threading model:: - - Main Thread: Inbound loop (BLPOP on broker) - Global SSE Thread: GET /api/events/global - Per-WS SSE Thread × N: GET /api/events?ws_id=X - """ - - def __init__( - self, - server_url: str = "http://localhost:8080", - broker: RedisBroker | None = None, - approval_timeout: float = 3600, - prefix: str = "turnstone", - node_id: str = "", - heartbeat_ttl: int = 60, - auth_token: str = "", - token_manager: Any = None, - tls_verify: Any = True, - tls_cert: tuple[str, str] | None = None, - ) -> None: - self._server_url = server_url.rstrip("/") - self._broker = broker or RedisBroker() - self._approval_timeout = approval_timeout - self._prefix = prefix - self._node_id = node_id # resolved in run() from server /health - self._server_max_ws: int = 10 # resolved in _fetch_node_id() from server /health - self._heartbeat_ttl = heartbeat_ttl - self._started_at = time.time() - self._auth_token = auth_token - self._token_manager = token_manager # ServiceTokenManager (auto-rotating) - self._tls_verify = tls_verify # CA cert path or ssl.SSLContext or True - self._tls_cert = tls_cert # (cert_path, key_path) for mTLS - - # Shared httpx client for short-lived POST requests (main thread only). - # Auth headers refreshed per-request via event hook so auto-rotating - # tokens are picked up transparently. - self._http = httpx.Client( - base_url=self._server_url, - timeout=30, - event_hooks={"request": [self._inject_auth]}, - verify=self._tls_verify, - cert=self._tls_cert, - ) - - # Protected by _lock — accessed from main, global SSE, and per-ws SSE threads - self._lock = threading.Lock() - self._ws_threads: dict[str, threading.Thread] = {} - self._ws_auto_approve: dict[str, bool] = {} - self._ws_approve_tools: dict[str, set[str]] = {} - self._active_sends: dict[str, str] = {} # ws_id → correlation_id - # Tombstone dicts: value is (request_id, resolved_at). - # resolved_at == 0.0 means active; > 0 means resolved at that - # monotonic time. Entries stay until cleaned up by ws_state / - # ws_closed events (or TTL fallback) to prevent TOCTOU races. - self._pending_approvals: dict[str, tuple[str, float]] = {} - self._pending_plan_reviews: dict[str, tuple[str, float]] = {} - self._running = True - - @property - def _auth_headers(self) -> dict[str, str]: - """Return current Authorization header, auto-rotating if managed.""" - if self._token_manager is not None: - return dict(self._token_manager.bearer_header) - if self._auth_token: - return {"Authorization": f"Bearer {self._auth_token}"} - return {} - - def _inject_auth(self, request: httpx.Request) -> None: - """httpx event hook: inject current auth header into each request.""" - headers = self._auth_headers - for k, v in headers.items(): - request.headers[k] = v - - # -- thread context helper ------------------------------------------------ - - def _run_in_context(self, fn: Callable[..., Any], *args: Any) -> Callable[[], None]: - """Return a callable that sets ``ctx_node_id`` before invoking *fn*.""" - node_id = self._node_id - - def _wrapper() -> None: - from turnstone.core.log import ctx_node_id - - ctx_node_id.set(node_id) - fn(*args) - - return _wrapper - - # -- public entry point -------------------------------------------------- - - def _fetch_server_metadata(self) -> None: - """Fetch max_ws from server /health (best-effort, non-blocking).""" - try: - resp = self._http.get("/health") - if resp.status_code == 200: - self._server_max_ws = resp.json().get("max_ws", 10) - except Exception: - log.debug("Failed to fetch server metadata", exc_info=True) - - def _fetch_node_id(self) -> str: - """Retrieve node_id from server /health with capped exponential backoff. - - Retries indefinitely so the bridge recovers when a server comes - back after a transient outage. 4xx responses (auth/config errors) - still fail fast. - """ - attempt = 0 - delay = 1.0 - max_delay = 60.0 - while True: - attempt += 1 - try: - resp = self._http.get("/health") - if 400 <= resp.status_code < 500: - log.critical("Server returned %d — check auth_token/config", resp.status_code) - raise SystemExit(1) - resp.raise_for_status() - data = resp.json() - nid = data.get("node_id", "") - if nid: - self._server_max_ws = data.get("max_ws", 10) - return str(nid) - log.warning("Server /health missing node_id (attempt %d)", attempt) - except SystemExit: - raise - except Exception as exc: - log.warning( - "Failed to fetch node_id from server (attempt %d): %s", - attempt, - exc, - ) - time.sleep(delay) - delay = min(delay * 2, max_delay) - - def run(self) -> None: - """Block until shutdown (KeyboardInterrupt).""" - if not self._node_id: - self._node_id = self._fetch_node_id() - else: - # node_id was pre-set — still need to fetch max_ws from server - self._fetch_server_metadata() - from turnstone.core.log import ctx_node_id - - ctx_node_id.set(self._node_id) - log.info( - "Bridge starting — node=%s server=%s max_ws=%d", - self._node_id, - self._server_url, - self._server_max_ws, - ) - self._recover_workstreams() - - heartbeat_t = threading.Thread( - target=self._run_in_context(self._heartbeat_loop), daemon=True - ) - heartbeat_t.start() - - global_t = threading.Thread(target=self._run_in_context(self._global_sse_loop), daemon=True) - global_t.start() - - try: - self._inbound_loop() - except KeyboardInterrupt: - log.info("Bridge shutting down") - finally: - self._running = False - self._http.close() - self._broker.close() - - # -- recovery ------------------------------------------------------------ - - def _recover_workstreams(self) -> None: - """Discover active workstreams on startup and register ownership.""" - try: - resp = self._http.get("/v1/api/workstreams") - data = resp.json() - for ws in data.get("workstreams", []): - ws_id = ws["id"] - ws_name = ws.get("name", "") - log.info("Recovered workstream %s (%s)", ws_id, ws_name) - self._broker.set_ws_owner(ws_id, self._node_id) - self._start_ws_sse(ws_id) - self._publish_cluster( - WorkstreamCreatedEvent( - ws_id=ws_id, - name=ws_name, - node_id=self._node_id, - ) - ) - except Exception as exc: - log.warning("Could not recover workstreams: %s", exc) - - # -- inbound loop -------------------------------------------------------- - - def _inbound_loop(self) -> None: - while self._running: - raw = self._broker.pop_inbound(timeout=5, node_id=self._node_id) - if raw is None: - continue - try: - msg = InboundMessage.from_json(raw) - self._dispatch(msg) - except Exception as exc: - log.error("Failed to process inbound message: %s", exc) - self._publish_global(ErrorEvent(message=f"Failed to process message: {exc}")) - - def _dispatch(self, msg: InboundMessage) -> None: - # Messages that need routing (have ws_id or target_node) - routed_handlers = { - "send": self._handle_send, - "command": self._handle_command, - "create_workstream": self._handle_create_ws, - "close_workstream": self._handle_close_ws, - "cancel": self._handle_cancel, - } - # Messages that are always local (no routing needed) - local_handlers = { - "approve": self._handle_approve, - "plan_feedback": self._handle_plan_feedback, - "list_workstreams": self._handle_list_ws, - "health": self._handle_health, - "list_nodes": self._handle_list_nodes, - } - - if msg.type in routed_handlers: - self._route_or_process(msg, routed_handlers[msg.type]) - elif msg.type in local_handlers: - local_handlers[msg.type](msg) - else: - self._publish_global( - ErrorEvent( - correlation_id=msg.correlation_id, - message=f"Unknown message type: {msg.type!r}", - ) - ) - - def _route_or_process( - self, msg: InboundMessage, handler: Callable[[InboundMessage], None] - ) -> None: - """Route a message to the correct node, or process locally.""" - target = getattr(msg, "target_node", "") - ws_id = getattr(msg, "ws_id", "") - - # Directed to a different node? - if target and target != self._node_id: - log.debug("Routing to node %s: %s", target, msg.type) - self._broker.push_inbound(msg.to_json(), node_id=target) - return - - # Existing workstream owned by another node? - if ws_id: - owner = self._broker.get_ws_owner(ws_id) - if owner and owner != self._node_id: - log.debug("Re-routing to owner %s for ws %s", owner, ws_id) - self._broker.push_inbound(msg.to_json(), node_id=owner) - return - - handler(msg) - - # -- handlers ------------------------------------------------------------ - - def _handle_send(self, msg: InboundMessage) -> None: - ws_id = getattr(msg, "ws_id", "") - message = getattr(msg, "message", "") - auto_approve = getattr(msg, "auto_approve", False) - auto_approve_tools = getattr(msg, "auto_approve_tools", []) - name = getattr(msg, "name", "") - - # Auto-create workstream if needed - if not ws_id: - ws_id, _resumed = self._create_ws_on_server( - name=name, - auto_approve=auto_approve, - auto_approve_tools=auto_approve_tools, - correlation_id=msg.correlation_id, - ) - if not ws_id: - return # error already published - else: - # Update approval settings for existing workstream - with self._lock: - if auto_approve: - self._ws_auto_approve[ws_id] = True - if auto_approve_tools: - self._ws_approve_tools[ws_id] = set(auto_approve_tools) - - with self._lock: - self._active_sends[ws_id] = msg.correlation_id - - try: - resp = self._http.post("/v1/api/send", json={"message": message, "ws_id": ws_id}) - data = resp.json() - except Exception: - with self._lock: - self._active_sends.pop(ws_id, None) - raise - - if data.get("status") != "ok": - with self._lock: - self._active_sends.pop(ws_id, None) - - self._publish_ws( - ws_id, - AckEvent( - ws_id=ws_id, - correlation_id=msg.correlation_id, - status="ok" if data.get("status") == "ok" else "error", - detail=data.get("error", ""), - ), - ) - - def _handle_approve(self, msg: InboundMessage) -> None: - request_id = getattr(msg, "request_id", "") - if request_id: - self._broker.push_response(request_id, msg.to_json()) - - def _handle_plan_feedback(self, msg: InboundMessage) -> None: - request_id = getattr(msg, "request_id", "") - if request_id: - self._broker.push_response(request_id, msg.to_json()) - - def _handle_cancel(self, msg: InboundMessage) -> None: - ws_id = getattr(msg, "ws_id", "") - resp = self._http.post("/v1/api/cancel", json={"ws_id": ws_id}) - data = resp.json() - self._publish_ws( - ws_id, - AckEvent( - ws_id=ws_id, - correlation_id=msg.correlation_id, - status="ok" if data.get("status") == "ok" else "error", - detail=data.get("error", ""), - ), - ) - - def _handle_command(self, msg: InboundMessage) -> None: - ws_id = getattr(msg, "ws_id", "") - command = getattr(msg, "command", "") - resp = self._http.post("/v1/api/command", json={"command": command, "ws_id": ws_id}) - data = resp.json() - self._publish_ws( - ws_id, - AckEvent( - ws_id=ws_id, - correlation_id=msg.correlation_id, - status="ok" if data.get("status") == "ok" else "error", - detail=data.get("error", ""), - ), - ) - - def _handle_create_ws(self, msg: InboundMessage) -> None: - name = getattr(msg, "name", "") - auto_approve = getattr(msg, "auto_approve", False) - auto_approve_tools = getattr(msg, "auto_approve_tools", []) - model = getattr(msg, "model", "") - initial_message = getattr(msg, "initial_message", "") - resume_ws = getattr(msg, "resume_ws", "") - user_id = getattr(msg, "user_id", "") - skill = getattr(msg, "skill", "") - if user_id: - log.info("bridge.create_ws user_id=%s name=%s model=%s", user_id, name, model) - ws_id, resumed = self._create_ws_on_server( - name=name, - auto_approve=auto_approve, - auto_approve_tools=auto_approve_tools, - correlation_id=msg.correlation_id, - model=model, - resume_ws=resume_ws, - skill=skill, - user_id=user_id, - ) - # Send initial_message only when no workstream was actually resumed. - # Use the server's `resumed` response (not just the intent) so that - # a pruned/missing workstream falls back to sending the initial message. - if ws_id and initial_message and not resumed: - # Track the send so the global SSE handler emits TurnCompleteEvent - # when the workstream returns to idle. - with self._lock: - self._active_sends[ws_id] = msg.correlation_id - try: - resp = self._http.post( - "/v1/api/send", json={"message": initial_message, "ws_id": ws_id} - ) - data = resp.json() - if data.get("error"): - log.warning("Initial message failed for ws %s: %s", ws_id, data["error"]) - with self._lock: - self._active_sends.pop(ws_id, None) - except Exception as exc: - log.warning("Initial message send failed for ws %s: %s", ws_id, exc) - with self._lock: - self._active_sends.pop(ws_id, None) - - def _handle_close_ws(self, msg: InboundMessage) -> None: - ws_id = getattr(msg, "ws_id", "") - resp = self._http.post("/v1/api/workstreams/close", json={"ws_id": ws_id}) - data = resp.json() - self._publish_ws( - ws_id, - AckEvent( - ws_id=ws_id, - correlation_id=msg.correlation_id, - status="ok" if data.get("status") == "ok" else "error", - detail=data.get("error", ""), - ), - ) - - def _handle_list_ws(self, msg: InboundMessage) -> None: - resp = self._http.get("/v1/api/workstreams") - data = resp.json() - self._publish_global( - WorkstreamListEvent( - correlation_id=msg.correlation_id, - workstreams=data.get("workstreams", []), - ) - ) - - def _handle_health(self, msg: InboundMessage) -> None: - resp = self._http.get("/health") - data = resp.json() - self._publish_global( - HealthResponseEvent( - correlation_id=msg.correlation_id, - data=data, - ) - ) - - # -- workstream creation helper ------------------------------------------ - - def _create_ws_on_server( - self, - name: str, - auto_approve: bool, - auto_approve_tools: list[str], - correlation_id: str, - model: str = "", - resume_ws: str = "", - skill: str = "", - user_id: str = "", - ) -> tuple[str, bool]: - """Create a workstream on the server. Returns (ws_id, resumed).""" - try: - payload: dict[str, Any] = {"name": name, "auto_approve": auto_approve} - if model: - payload["model"] = model - if resume_ws: - payload["resume_ws"] = resume_ws - if skill: - payload["skill"] = skill - if user_id: - payload["user_id"] = user_id - resp = self._http.post( - "/v1/api/workstreams/new", - json=payload, - ) - data = resp.json() - if "error" in data: - self._publish_global( - AckEvent( - correlation_id=correlation_id, - status="error", - detail=data["error"], - ) - ) - return "", False - ws_id: str = data["ws_id"] - ws_name = data.get("name", "") - resumed = data.get("resumed", False) - - self._broker.set_ws_owner(ws_id, self._node_id) - - with self._lock: - if auto_approve: - self._ws_auto_approve[ws_id] = True - if auto_approve_tools: - self._ws_approve_tools[ws_id] = set(auto_approve_tools) - - self._start_ws_sse(ws_id) - - self._publish_global( - WorkstreamCreatedEvent( - ws_id=ws_id, - name=ws_name, - correlation_id=correlation_id, - resumed=resumed, - message_count=data.get("message_count", 0), - ) - ) - self._publish_cluster( - WorkstreamCreatedEvent( - ws_id=ws_id, - name=ws_name, - correlation_id=correlation_id, - node_id=self._node_id, - ) - ) - - # Emit per-workstream resume confirmation. - if resumed: - self._publish_ws( - ws_id, - WorkstreamResumedEvent( - ws_id=ws_id, - correlation_id=correlation_id, - message_count=data.get("message_count", 0), - name=ws_name, - ), - ) - - return ws_id, resumed - except Exception as exc: - self._publish_global( - AckEvent( - correlation_id=correlation_id, - status="error", - detail=str(exc), - ) - ) - return "", False - - # -- SSE consumption ----------------------------------------------------- - - def _start_ws_sse(self, ws_id: str) -> None: - with self._lock: - if ws_id in self._ws_threads and self._ws_threads[ws_id].is_alive(): - return - t = threading.Thread(target=self._run_in_context(self._ws_sse_loop, ws_id), daemon=True) - self._ws_threads[ws_id] = t - t.start() - - def _ws_sse_loop(self, ws_id: str) -> None: - """Consume per-workstream SSE and forward events.""" - # Each SSE thread gets its own httpx client (not thread-safe to share). - # Use event hooks so auth headers refresh on each reconnect. - with httpx.Client( - base_url=self._server_url, - timeout=None, - event_hooks={"request": [self._inject_auth]}, - verify=self._tls_verify, - cert=self._tls_cert, - ) as sse_client: - while self._running: - # Stop if workstream was closed (thread removed from registry) - with self._lock: - if ws_id not in self._ws_threads: - return - try: - with sse_client.stream("GET", f"/v1/api/events?ws_id={ws_id}") as resp: - for data in _iter_sse_data(resp): - if not self._running: - break - with self._lock: - if ws_id not in self._ws_threads: - return - self._handle_ws_event(ws_id, data) - except Exception as exc: - if self._running: - log.debug("WS SSE reconnecting (%s): %s", ws_id, exc) - time.sleep(2) - - def _handle_ws_event(self, ws_id: str, data: dict[str, Any]) -> None: - etype = data.get("type", "") - - if etype == "content": - self._publish_ws(ws_id, ContentEvent(ws_id=ws_id, text=data.get("text", ""))) - elif etype == "reasoning": - self._publish_ws(ws_id, ReasoningEvent(ws_id=ws_id, text=data.get("text", ""))) - elif etype == "tool_info": - self._publish_ws(ws_id, ToolInfoEvent(ws_id=ws_id, items=data.get("items", []))) - elif etype == "approve_request": - self._handle_approval(ws_id, data) - elif etype == "plan_review": - self._handle_plan_review(ws_id, data) - elif etype == "tool_output_chunk": - self._publish_ws( - ws_id, - ToolOutputChunkEvent( - ws_id=ws_id, - call_id=data.get("call_id", ""), - chunk=data.get("chunk", ""), - ), - ) - elif etype == "tool_result": - self._publish_ws( - ws_id, - ToolResultEvent( - ws_id=ws_id, - call_id=data.get("call_id", ""), - name=data.get("name", ""), - output=data.get("output", ""), - ), - ) - elif etype == "status": - self._publish_ws( - ws_id, - StatusEvent( - ws_id=ws_id, - prompt_tokens=data.get("prompt_tokens", 0), - completion_tokens=data.get("completion_tokens", 0), - total_tokens=data.get("total_tokens", 0), - context_window=data.get("context_window", 0), - pct=data.get("pct", 0), - effort=data.get("effort", ""), - cache_creation_tokens=data.get("cache_creation_tokens", 0), - cache_read_tokens=data.get("cache_read_tokens", 0), - tool_calls_this_turn=data.get("tool_calls_this_turn", 0), - turn_count=data.get("turn_count", 0), - ), - ) - elif etype == "error": - self._publish_ws(ws_id, ErrorEvent(ws_id=ws_id, message=data.get("message", ""))) - elif etype == "info": - self._publish_ws(ws_id, InfoEvent(ws_id=ws_id, message=data.get("message", ""))) - elif etype == "intent_verdict": - self._publish_ws( - ws_id, - IntentVerdictEvent( - ws_id=ws_id, - call_id=data.get("call_id", ""), - func_name=data.get("func_name", ""), - intent_summary=data.get("intent_summary", ""), - risk_level=data.get("risk_level", ""), - confidence=float(data.get("confidence", 0.0)), - recommendation=data.get("recommendation", ""), - reasoning=data.get("reasoning", ""), - evidence=json.dumps(data.get("evidence", [])), - tier=data.get("tier", ""), - judge_model=data.get("judge_model", ""), - verdict_id=data.get("verdict_id", ""), - latency_ms=int(data.get("latency_ms", 0)), - ), - ) - elif etype == "output_warning": - self._publish_ws( - ws_id, - OutputWarningEvent( - ws_id=ws_id, - call_id=data.get("call_id", ""), - func_name=data.get("func_name", ""), - risk_level=data.get("risk_level", "none"), - flags=json.dumps(data.get("flags", [])), - annotations=json.dumps(data.get("annotations", [])), - redacted=int(data.get("redacted", False)), - ), - ) - elif etype == "stream_end": - self._publish_ws(ws_id, StreamEndEvent(ws_id=ws_id)) - - def _handle_approval(self, ws_id: str, data: dict[str, Any]) -> None: - """Handle an approval request — auto-approve or forward to client.""" - items = data.get("items", []) - - # Read flags under lock, then release before any HTTP calls. - with self._lock: - auto = self._ws_auto_approve.get(ws_id, False) - approve_set = self._ws_approve_tools.get(ws_id, DEFAULT_SAFE_TOOLS) - - tool_names = { - it.get("approval_label", "") or it.get("func_name", "") - for it in items - if it.get("needs_approval") and it.get("func_name") and not it.get("error") - } - - # Evaluate admin tool policies — deny overrides auto-approve. - if tool_names: - try: - from turnstone.core.policy import evaluate_tool_policies_batch - from turnstone.core.storage import _registry as _stor_reg - - # Use _storage directly — get_storage() auto-initializes - # a local SQLite DB which would have no admin policies. - storage = _stor_reg._storage - if storage is not None: - verdicts = evaluate_tool_policies_batch(storage, list(tool_names)) - if any(v == "deny" for v in verdicts.values()): - denied = [n for n, v in verdicts.items() if v == "deny"] - self._api_approve( - ws_id, - approved=False, - feedback=f"Blocked by tool policy: {', '.join(denied)}", - ) - return - if all(verdicts.get(n) == "allow" for n in tool_names): - self._api_approve(ws_id, approved=True) - return - except Exception: - log.debug("Tool policy evaluation failed for ws %s", ws_id, exc_info=True) - - if auto: - self._api_approve(ws_id, approved=True) - return - - if tool_names and tool_names.issubset(approve_set): - self._api_approve(ws_id, approved=True) - return - - # Skip if an approval is already pending (or recently resolved) for - # this workstream. Resolved tombstones are cleaned up by ws_state - # events, preventing TOCTOU races from SSE reconnect re-injection. - # TTL fallback: if the global SSE is lagging, allow re-entry after - # 3x the approval timeout so the workstream doesn't hang. - with self._lock: - entry = self._pending_approvals.get(ws_id) - if entry is not None: - _, resolved_at = entry - ttl = 3 * self._approval_timeout - if not resolved_at or time.monotonic() - resolved_at < ttl: - log.debug("Skipping duplicate approval for ws %s", ws_id) - return - self._pending_approvals.pop(ws_id) - request_id = uuid.uuid4().hex[:12] - self._pending_approvals[ws_id] = (request_id, 0.0) - - # Forward to client — spawn a thread so we don't block SSE consumption - self._publish_ws( - ws_id, - ApprovalRequestEvent( - ws_id=ws_id, - correlation_id=request_id, - items=items, - ), - ) - - def _wait_approval() -> None: - try: - raw_resp = self._broker.pop_response(request_id, timeout=self._approval_timeout) - if raw_resp: - resp_msg = InboundMessage.from_json(raw_resp) - approved = getattr(resp_msg, "approved", False) - feedback = getattr(resp_msg, "feedback", None) - always = getattr(resp_msg, "always", False) - self._api_approve(ws_id, approved=approved, feedback=feedback) - if always and approved: - tool_names = { - it.get("approval_label", "") or it.get("func_name", "") - for it in items - if it.get("needs_approval") - and it.get("func_name") - and not it.get("error") - } - tool_names.discard("") - tool_names.discard("__budget_override__") - if tool_names: - with self._lock: - existing = self._ws_approve_tools.get( - ws_id, set(DEFAULT_SAFE_TOOLS) - ) - self._ws_approve_tools[ws_id] = existing | tool_names - else: - log.warning("Approval timeout for ws %s — denying", ws_id) - self._api_approve(ws_id, approved=False, feedback="Approval timed out") - finally: - with self._lock: - entry = self._pending_approvals.get(ws_id) - if entry is not None and entry[0] == request_id: - self._pending_approvals[ws_id] = (request_id, time.monotonic()) - - threading.Thread(target=self._run_in_context(_wait_approval), daemon=True).start() - - def _handle_plan_review(self, ws_id: str, data: dict[str, Any]) -> None: - """Handle a plan review request — auto-approve or forward to client.""" - with self._lock: - auto = self._ws_auto_approve.get(ws_id, False) - # Skip if a plan review is already pending or recently resolved - # (SSE reconnect guard). Resolved tombstones are cleaned up by - # ws_state events, allowing the refinement loop to re-enter. - # TTL fallback handles global SSE lag. - if not auto and ws_id in self._pending_plan_reviews: - _, resolved_at = self._pending_plan_reviews[ws_id] - ttl = 3 * self._approval_timeout - if not resolved_at or time.monotonic() - resolved_at < ttl: - log.debug("Skipping duplicate plan review for ws %s", ws_id) - return - self._pending_plan_reviews.pop(ws_id) - if not auto: - request_id = uuid.uuid4().hex[:12] - self._pending_plan_reviews[ws_id] = (request_id, 0.0) - - if auto: - self._http.post("/v1/api/plan", json={"feedback": "", "ws_id": ws_id}) - return - - self._publish_ws( - ws_id, - PlanReviewEvent( - ws_id=ws_id, - correlation_id=request_id, - content=data.get("content", ""), - ), - ) - - def _wait_plan() -> None: - def _mark_resolved() -> None: - # Unlike _wait_approval (which uses finally), we mark - # resolved explicitly because it must happen BEFORE the - # HTTP POST — the refinement loop needs the tombstone - # in place while the server processes feedback. - with self._lock: - entry = self._pending_plan_reviews.get(ws_id) - if entry is not None and entry[0] == request_id: - self._pending_plan_reviews[ws_id] = (request_id, time.monotonic()) - - try: - raw_resp = self._broker.pop_response(request_id, timeout=self._approval_timeout) - # Mark resolved instead of popping — the tombstone blocks - # stale SSE re-injections. Cleanup happens on the next - # ws_state event (global SSE), which the server emits - # before starting the next LLM turn. The refinement - # plan_review event (per-WS SSE) therefore arrives after - # cleanup, assuming global SSE is not delayed by more - # than one LLM turn. - _mark_resolved() - if raw_resp: - resp_msg = InboundMessage.from_json(raw_resp) - feedback = getattr(resp_msg, "feedback", "") - self._http.post("/v1/api/plan", json={"feedback": feedback, "ws_id": ws_id}) - else: - log.warning("Plan review timeout for ws %s — rejecting", ws_id) - self._http.post("/v1/api/plan", json={"feedback": "reject", "ws_id": ws_id}) - except Exception: - _mark_resolved() - # Best-effort rejection so the server doesn't hang - try: - self._http.post( - "/v1/api/plan", - json={"feedback": "reject", "ws_id": ws_id}, - ) - except Exception: - log.warning("Failed to reject plan for ws=%s", ws_id, exc_info=True) - raise - - threading.Thread(target=self._run_in_context(_wait_plan), daemon=True).start() - - def _api_approve( - self, - ws_id: str, - approved: bool, - feedback: str | None = None, - ) -> None: - body: dict[str, Any] = {"approved": approved, "ws_id": ws_id} - if feedback: - body["feedback"] = feedback - self._http.post("/v1/api/approve", json=body) - - # -- global SSE ---------------------------------------------------------- - - def _global_sse_loop(self) -> None: - # Own httpx client for the long-lived SSE connection. - # Use event hooks so auth headers refresh on each reconnect. - with httpx.Client( - base_url=self._server_url, - timeout=None, - event_hooks={"request": [self._inject_auth]}, - verify=self._tls_verify, - cert=self._tls_cert, - ) as sse_client: - while self._running: - try: - with sse_client.stream("GET", "/v1/api/events/global") as resp: - for data in _iter_sse_data(resp): - if not self._running: - break - self._handle_global_event(data) - except Exception as exc: - if self._running: - log.debug("Global SSE reconnecting: %s", exc) - time.sleep(2) - - def _handle_global_event(self, data: dict[str, Any]) -> None: - etype = data.get("type", "") - ws_id = data.get("ws_id", "") - - if etype == "ws_state": - state = data.get("state", "") - self._publish_ws(ws_id, StateChangeEvent(ws_id=ws_id, state=state)) - self._publish_global(StateChangeEvent(ws_id=ws_id, state=state)) - self._publish_cluster( - ClusterStateEvent( - ws_id=ws_id, - state=state, - node_id=self._node_id, - tokens=data.get("tokens", 0), - context_ratio=data.get("context_ratio", 0.0), - activity=data.get("activity", ""), - activity_state=data.get("activity_state", ""), - ) - ) - - # Clean up resolved pending tombstones on any state transition. - # This allows legitimate new approval/plan_review events to - # proceed while stale SSE re-injections remain blocked. - with self._lock: - a_entry = self._pending_approvals.get(ws_id) - if a_entry is not None and a_entry[1] > 0: - self._pending_approvals.pop(ws_id, None) - p_entry = self._pending_plan_reviews.get(ws_id) - if p_entry is not None and p_entry[1] > 0: - self._pending_plan_reviews.pop(ws_id, None) - - # Completion detection — emit for all idle transitions so - # channel adapters can finalize streaming messages even when - # the turn was initiated from the server UI (no correlation_id). - if state == "idle": - with self._lock: - cid = self._active_sends.pop(ws_id, None) - self._publish_ws( - ws_id, - TurnCompleteEvent( - ws_id=ws_id, - correlation_id=cid or "", - content=data.get("content", ""), - ), - ) - - elif etype == "ws_rename": - self._publish_global(WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", ""))) - self._publish_cluster(WorkstreamRenameEvent(ws_id=ws_id, name=data.get("name", ""))) - - elif etype == "ws_closed": - self._publish_global(WorkstreamClosedEvent(ws_id=ws_id)) - self._publish_cluster(WorkstreamClosedEvent(ws_id=ws_id)) - self._broker.del_ws_owner(ws_id) - with self._lock: - self._ws_threads.pop(ws_id, None) - self._ws_auto_approve.pop(ws_id, None) - self._ws_approve_tools.pop(ws_id, None) - self._active_sends.pop(ws_id, None) - self._pending_approvals.pop(ws_id, None) - self._pending_plan_reviews.pop(ws_id, None) - - # -- heartbeat ----------------------------------------------------------- - - def _heartbeat_loop(self) -> None: - """Periodically register this node in the broker. - - An initial jitter (derived from the node_id) staggers heartbeats - across cluster nodes so they don't all hit Redis at the same instant. - """ - # Deterministic per-node jitter: spread across first quarter of TTL - h = hash(self._node_id) & 0x7FFFFFFF - jitter = (h % 2147483647) / 2147483647 * (self._heartbeat_ttl / 4) - if jitter > 0.1: - time.sleep(jitter) - while self._running: - self._broker.register_node( - self._node_id, - { - "server_url": self._server_url, - "started": self._started_at, - "max_ws": self._server_max_ws, - }, - ttl=self._heartbeat_ttl, - ) - time.sleep(self._heartbeat_ttl / 2) - - # -- node listing -------------------------------------------------------- - - def _handle_list_nodes(self, msg: InboundMessage) -> None: - nodes = self._broker.list_nodes() - self._publish_global(NodeListEvent(correlation_id=msg.correlation_id, nodes=nodes)) - - # -- publish helpers ----------------------------------------------------- - - def _publish_ws(self, ws_id: str, event: OutboundEvent) -> None: - channel = f"{self._prefix}:events:{ws_id}" - self._broker.publish_outbound(channel, event.to_json()) - - def _publish_global(self, event: OutboundEvent) -> None: - self._broker.publish_outbound(f"{self._prefix}:events:global", event.to_json()) - - def _publish_cluster(self, event: OutboundEvent) -> None: - self._broker.publish_outbound(f"{self._prefix}:events:cluster", event.to_json()) - - -# --------------------------------------------------------------------------- -# SSE parsing helper -# --------------------------------------------------------------------------- - - -def _iter_sse_data(resp: httpx.Response) -> Iterator[dict[str, Any]]: - """Yield parsed JSON dicts from an SSE stream.""" - from httpx_sse import EventSource - - source = EventSource(resp) - for sse in source.iter_sse(): - if sse.data: - try: - data: dict[str, Any] = json.loads(sse.data) - yield data - except json.JSONDecodeError: - log.debug("Skipping malformed SSE data: %.200s", sse.data) - - -# --------------------------------------------------------------------------- -# CLI entry point -# --------------------------------------------------------------------------- - - -def main() -> None: - import argparse - - parser = argparse.ArgumentParser( - description="turnstone message queue bridge — connects Redis queues to turnstone-server" - ) - parser.add_argument( - "--server-url", - default="http://localhost:8080", - help="turnstone-server URL (default: %(default)s)", - ) - from turnstone.mq.broker import add_redis_args - - add_redis_args(parser) - parser.add_argument( - "--approval-timeout", - type=float, - default=3600, - help="Seconds to wait for approval responses (default: %(default)s)", - ) - parser.add_argument( - "--node-id", - default="", - help="Node identifier for multi-node routing (default: hostname)", - ) - parser.add_argument( - "--heartbeat-ttl", - type=int, - default=60, - help="Heartbeat TTL in seconds (default: %(default)s)", - ) - from turnstone.core.log import add_log_args - - add_log_args(parser) - parser.add_argument( - "--auth-token", - default=os.environ.get("TURNSTONE_AUTH_TOKEN", ""), - help="Bearer token for authenticating to turnstone-server (default: $TURNSTONE_AUTH_TOKEN)", - ) - from turnstone.core.config import add_config_arg, apply_config - - add_config_arg(parser) - apply_config(parser, ["bridge", "redis", "auth"]) - args = parser.parse_args() - - from turnstone.core.log import configure_logging_from_args - - configure_logging_from_args(args, "bridge") - - from turnstone.mq.broker import broker_from_args - - broker = broker_from_args(args) - # If no explicit auth token is provided, use a ServiceTokenManager - # so bridge JWTs auto-rotate (1-hour expiry, refreshed at 80%). - # A shared JWT secret is required for multi-service deployments — - # ephemeral secrets differ per process and break inter-service auth. - auth_token = args.auth_token - token_manager = None - if not auth_token: - jwt_secret = os.environ.get("TURNSTONE_JWT_SECRET", "") - if not jwt_secret: - log.error( - "TURNSTONE_JWT_SECRET is not set and no --auth-token provided. " - "The bridge cannot authenticate to the server. Set TURNSTONE_JWT_SECRET " - "to a shared secret (at least 32 characters) or pass --auth-token." - ) - raise SystemExit(1) - from turnstone.core.auth import JWT_AUD_SERVER, ServiceTokenManager - - token_manager = ServiceTokenManager( - user_id="bridge", - scopes=frozenset({"approve"}), - source="bridge", - secret=jwt_secret, - audience=JWT_AUD_SERVER, - expiry_hours=1, - ) - log.info("bridge.jwt_minted") - - # TLS: request cert from console ACME if enabled - tls_verify: Any = True - tls_cert: tuple[str, str] | None = None - if os.environ.get("TURNSTONE_TLS_ENABLED", "").lower() in ("true", "1", "yes"): - try: - import asyncio - import socket - - from turnstone.core.storage import init_storage - from turnstone.core.tls import TLSClient - - db_backend = os.environ.get("TURNSTONE_DB_BACKEND", "sqlite") - db_url = os.environ.get("TURNSTONE_DB_URL", "") - db_path = os.environ.get("TURNSTONE_DB_PATH", "") - storage = init_storage(db_backend, path=db_path, url=db_url) - - hostname = socket.getfqdn() - hostnames = [hostname, "localhost", "127.0.0.1"] - extra_sans = os.environ.get("TURNSTONE_TLS_SANS", "") - if extra_sans: - hostnames.extend(s.strip() for s in extra_sans.split(",") if s.strip()) - tls_client = TLSClient( - storage=storage, - hostnames=hostnames, - ) - asyncio.run(tls_client.init()) - ssl_ctx = tls_client.get_client_ssl_context() - if ssl_ctx: - # SSLContext has both CA (verify server) and client cert - # (present to server) loaded — full mTLS in one object - tls_verify = ssl_ctx - if args.server_url.startswith("http://"): - args.server_url = args.server_url.replace("http://", "https://") - log.info("bridge.tls.enabled: %s", args.server_url) - except ImportError: - log.warning("TLS enabled but lacme not installed") - except Exception: - log.warning("bridge.tls.init_failed", exc_info=True) - - bridge = Bridge( - server_url=args.server_url, - broker=broker, - approval_timeout=args.approval_timeout, - node_id=args.node_id, - heartbeat_ttl=args.heartbeat_ttl, - auth_token=auth_token, - token_manager=token_manager, - tls_verify=tls_verify, - tls_cert=tls_cert, - ) - bridge.run() - - -if __name__ == "__main__": - main() diff --git a/turnstone/mq/broker.py b/turnstone/mq/broker.py deleted file mode 100644 index 604f1e67..00000000 --- a/turnstone/mq/broker.py +++ /dev/null @@ -1,342 +0,0 @@ -"""Abstract message broker and Redis implementation. - -The MessageBroker protocol defines the interface for inbound queuing, outbound -pub/sub, per-request response queues, and multi-node routing primitives. -RedisBroker is the default provider. -""" - -from __future__ import annotations - -import contextlib -import json -from typing import TYPE_CHECKING, Any, Protocol, cast - -if TYPE_CHECKING: - from collections.abc import Callable - - import redis as _redis_t - - -class MessageBroker(Protocol): - """Abstract message broker for inbound/outbound communication. - - Implementations must provide: - - Reliable inbound queue (FIFO, at-least-once delivery) - - Outbound pub/sub channels (fan-out to all subscribers) - - Per-request response queues for approval request/response correlation - - Workstream ownership tracking (ws_id → node_id) - - Node registry with heartbeat - """ - - def push_inbound(self, message: str, node_id: str = "") -> None: - """Push a message onto the inbound queue. - - If *node_id* is set, pushes to the per-node queue for directed - routing. Otherwise pushes to the shared queue. - """ - ... - - def pop_inbound(self, timeout: float = 0, node_id: str = "") -> str | None: - """Pop next message from the inbound queue (bridge side). - - If *node_id* is set, BLPOPs from both the per-node queue (priority) - and the shared queue. Otherwise BLPOPs from the shared queue only. - Returns None on timeout. - """ - ... - - def publish_outbound(self, channel: str, event: str) -> None: - """Publish an event to an outbound channel.""" - ... - - def subscribe_outbound(self, channel: str, callback: Callable[[str], None]) -> None: - """Subscribe to an outbound channel.""" - ... - - def unsubscribe_outbound(self, channel: str) -> None: - """Unsubscribe from an outbound channel.""" - ... - - def push_response(self, queue_name: str, message: str) -> None: - """Push a response onto a named response queue.""" - ... - - def pop_response(self, queue_name: str, timeout: float = 300) -> str | None: - """Pop from a named response queue. Returns None on timeout.""" - ... - - # -- routing primitives -------------------------------------------------- - - def set_ws_owner(self, ws_id: str, node_id: str, ttl: int = 0) -> None: - """Register which node owns a workstream.""" - ... - - def get_ws_owner(self, ws_id: str) -> str | None: - """Look up the node that owns a workstream. Returns None if unowned.""" - ... - - def del_ws_owner(self, ws_id: str) -> None: - """Remove workstream ownership (on close).""" - ... - - def register_node(self, node_id: str, metadata: dict[str, Any], ttl: int = 60) -> None: - """Register or refresh a node's heartbeat with metadata.""" - ... - - def list_nodes(self) -> list[dict[str, Any]]: - """List all active nodes (those with unexpired heartbeats).""" - ... - - def subscribe_cluster(self, callback: Callable[[str], None]) -> None: - """Subscribe to the cluster-wide event channel.""" - ... - - def close(self) -> None: - """Clean up connections.""" - ... - - -class RedisBroker: - """Redis-backed MessageBroker using lists (queues) and pub/sub (events). - - Queue keys: - ``{prefix}:inbound`` — shared inbound command queue - ``{prefix}:inbound:{node_id}`` — per-node directed queue - ``{prefix}:resp:{request_id}`` — per-request response queues - - Routing keys: - ``{prefix}:ws:{ws_id}`` — workstream ownership (string) - ``{prefix}:node:{node_id}`` — node heartbeat + metadata (string/JSON) - - Pub/sub channels: - ``{prefix}:events:global`` — global event channel - ``{prefix}:events:{ws_id}`` — per-workstream event channel - ``{prefix}:events:cluster`` — cluster-wide state changes - """ - - def __init__( - self, - host: str = "localhost", - port: int = 6379, - db: int = 0, - prefix: str = "turnstone", - password: str | None = None, - response_ttl: int = 600, - ssl: bool = False, - ssl_ca_certs: str | None = None, - ssl_certfile: str | None = None, - ssl_keyfile: str | None = None, - ) -> None: - import redis - - self._prefix = prefix - self._response_ttl = response_ttl - pool_kwargs: dict[str, Any] = {} - if ssl: - pool_kwargs["connection_class"] = redis.SSLConnection - if ssl_ca_certs: - pool_kwargs["ssl_ca_certs"] = ssl_ca_certs - if ssl_certfile: - pool_kwargs["ssl_certfile"] = ssl_certfile - if ssl_keyfile: - pool_kwargs["ssl_keyfile"] = ssl_keyfile - self._pool: _redis_t.ConnectionPool = redis.ConnectionPool( - host=host, - port=port, - db=db, - password=password, - decode_responses=True, - retry_on_timeout=True, - max_connections=200, - **pool_kwargs, - ) - self._redis: _redis_t.Redis[str] = cast( - "_redis_t.Redis[str]", - redis.Redis(connection_pool=self._pool), - ) - self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True) - self._listener_thread: Any = None - self._running = True - - # -- inbound queue ------------------------------------------------------- - - def push_inbound(self, message: str, node_id: str = "") -> None: - if node_id: - self._redis.rpush(f"{self._prefix}:inbound:{node_id}", message) - else: - self._redis.rpush(f"{self._prefix}:inbound", message) - - def pop_inbound(self, timeout: float = 0, node_id: str = "") -> str | None: - t = int(timeout) if timeout > 0 else 0 - if node_id: - # Per-node queue first (priority), then shared queue - result = self._redis.blpop( - [f"{self._prefix}:inbound:{node_id}", f"{self._prefix}:inbound"], - timeout=t, - ) - else: - result = self._redis.blpop(f"{self._prefix}:inbound", timeout=t) - return result[1] if result else None - - # -- outbound pub/sub ---------------------------------------------------- - - def publish_outbound(self, channel: str, event: str) -> None: - self._redis.publish(channel, event) - - def subscribe_outbound(self, channel: str, callback: Callable[[str], None]) -> None: - def _handler(msg: dict[str, Any]) -> None: - callback(msg["data"]) - - self._pubsub.subscribe(**{channel: _handler}) - if self._listener_thread is None or not self._listener_thread.is_alive(): - self._listener_thread = self._pubsub.run_in_thread(sleep_time=0.1, daemon=True) - - def unsubscribe_outbound(self, channel: str) -> None: - self._pubsub.unsubscribe(channel) - - # -- response queues ----------------------------------------------------- - - def push_response(self, queue_name: str, message: str) -> None: - key = f"{self._prefix}:resp:{queue_name}" - self._redis.rpush(key, message) - self._redis.expire(key, self._response_ttl) - - def pop_response(self, queue_name: str, timeout: float = 300) -> str | None: - key = f"{self._prefix}:resp:{queue_name}" - result = self._redis.blpop(key, timeout=int(timeout)) - return result[1] if result else None - - # -- routing primitives -------------------------------------------------- - - def set_ws_owner(self, ws_id: str, node_id: str, ttl: int = 0) -> None: - key = f"{self._prefix}:ws:{ws_id}" - if ttl > 0: - self._redis.set(key, node_id, ex=ttl) - else: - self._redis.set(key, node_id) - - def get_ws_owner(self, ws_id: str) -> str | None: - return self._redis.get(f"{self._prefix}:ws:{ws_id}") - - def del_ws_owner(self, ws_id: str) -> None: - self._redis.delete(f"{self._prefix}:ws:{ws_id}") - - def register_node(self, node_id: str, metadata: dict[str, Any], ttl: int = 60) -> None: - key = f"{self._prefix}:node:{node_id}" - self._redis.set(key, json.dumps(metadata), ex=ttl) - - def list_nodes(self) -> list[dict[str, Any]]: - pattern = f"{self._prefix}:node:*" - prefix_len = len(f"{self._prefix}:node:") - # Collect all keys first, then batch-fetch with MGET to avoid - # N+1 round-trips (1 GET per node). - keys = list(self._redis.scan_iter(match=pattern, count=100)) - if not keys: - return [] - values = self._redis.mget(keys) - nodes: list[dict[str, Any]] = [] - for key, raw in zip(keys, values, strict=True): - if raw: - try: - meta: dict[str, Any] = json.loads(raw) - except json.JSONDecodeError: - meta = {} - meta["node_id"] = key[prefix_len:] - nodes.append(meta) - return nodes - - # -- cluster event channel ----------------------------------------------- - - def subscribe_cluster(self, callback: Callable[[str], None]) -> None: - """Subscribe to the cluster-wide event channel.""" - channel = f"{self._prefix}:events:cluster" - self.subscribe_outbound(channel, callback) - - # -- lifecycle ----------------------------------------------------------- - - def close(self) -> None: - self._running = False - if self._listener_thread is not None: - self._listener_thread.stop() - self._listener_thread = None - with contextlib.suppress(Exception): - self._pubsub.close() - self._pool.disconnect() - - -# --------------------------------------------------------------------------- -# CLI helpers (shared across bridge, console, channels) -# --------------------------------------------------------------------------- - - -def add_redis_args(parser: Any) -> None: - """Add Redis CLI arguments including TLS options.""" - import os - - parser.add_argument( - "--redis-host", - default=os.environ.get("REDIS_HOST", "localhost"), - help="Redis host (default: $REDIS_HOST or localhost)", - ) - parser.add_argument( - "--redis-port", - type=int, - default=int(os.environ.get("REDIS_PORT", "6379")), - help="Redis port (default: %(default)s)", - ) - parser.add_argument( - "--redis-password", - default=os.environ.get("REDIS_PASSWORD"), - help="Redis password (default: $REDIS_PASSWORD)", - ) - parser.add_argument( - "--redis-db", - type=int, - default=0, - help="Redis DB number (default: %(default)s)", - ) - parser.add_argument("--redis-tls", action="store_true", help="Enable Redis TLS") - parser.add_argument("--redis-tls-ca", default=None, help="Redis CA cert path") - parser.add_argument("--redis-tls-cert", default=None, help="Redis client cert path") - parser.add_argument("--redis-tls-key", default=None, help="Redis client key path") - - -def _redis_tls_kwargs(args: Any) -> dict[str, Any]: - """Extract Redis TLS kwargs from parsed args.""" - kwargs: dict[str, Any] = {} - if getattr(args, "redis_tls", False): - kwargs["ssl"] = True - ca = getattr(args, "redis_tls_ca", None) - if ca: - kwargs["ssl_ca_certs"] = ca - cert = getattr(args, "redis_tls_cert", None) - if cert: - kwargs["ssl_certfile"] = cert - key = getattr(args, "redis_tls_key", None) - if key: - kwargs["ssl_keyfile"] = key - return kwargs - - -def broker_from_args(args: Any) -> RedisBroker: - """Create a :class:`RedisBroker` from parsed CLI arguments.""" - return RedisBroker( - host=args.redis_host, - port=args.redis_port, - db=args.redis_db, - password=args.redis_password or None, - **_redis_tls_kwargs(args), - ) - - -def async_broker_from_args(args: Any) -> Any: - """Create an :class:`AsyncRedisBroker` from parsed CLI arguments.""" - from turnstone.mq.async_broker import AsyncRedisBroker - - return AsyncRedisBroker( - host=args.redis_host, - port=args.redis_port, - db=args.redis_db, - password=args.redis_password or None, - **_redis_tls_kwargs(args), - ) diff --git a/turnstone/mq/client.py b/turnstone/mq/client.py deleted file mode 100644 index 7612caa0..00000000 --- a/turnstone/mq/client.py +++ /dev/null @@ -1,316 +0,0 @@ -"""Client library for interacting with turnstone through a message broker. - -Usage:: - - from turnstone.mq.client import TurnstoneClient - - client = TurnstoneClient() - result = client.send_and_wait( - "What files are in the current directory?", - auto_approve=True, - ) - print(result.content) - client.close() -""" - -from __future__ import annotations - -import threading -from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any - -from turnstone.mq.broker import MessageBroker, RedisBroker -from turnstone.mq.protocol import ( - ApproveMessage, - CloseWorkstreamMessage, - CommandMessage, - ContentEvent, - CreateWorkstreamMessage, - ErrorEvent, - HealthMessage, - ListWorkstreamsMessage, - OutboundEvent, - PlanFeedbackMessage, - ReasoningEvent, - SendMessage, - ToolResultEvent, - TurnCompleteEvent, - WorkstreamCreatedEvent, -) - -if TYPE_CHECKING: - from collections.abc import Callable - - -@dataclass -class TurnResult: - """Aggregated result of a send_and_wait call.""" - - correlation_id: str = "" - ws_id: str = "" - content_parts: list[str] = field(default_factory=list) - reasoning_parts: list[str] = field(default_factory=list) - tool_results: list[tuple[str, str]] = field(default_factory=list) - errors: list[str] = field(default_factory=list) - timed_out: bool = False - - @property - def content(self) -> str: - return "".join(self.content_parts) - - @property - def reasoning(self) -> str: - return "".join(self.reasoning_parts) - - @property - def ok(self) -> bool: - return not self.timed_out and not self.errors - - -class TurnstoneClient: - """Client library for turnstone message queue integration. - - All methods are synchronous. The broker handles background threads - for pub/sub subscriptions. - """ - - def __init__( - self, - broker: MessageBroker | None = None, - prefix: str = "turnstone", - **redis_kwargs: object, - ) -> None: - """Create a client. - - Pass ``broker`` for a custom broker, or provide Redis kwargs - (``host``, ``port``, ``db``, ``password``) to use the default - RedisBroker. - """ - self._broker: MessageBroker = broker or RedisBroker(**redis_kwargs) # type: ignore[arg-type] - self._prefix = prefix - - # -- fire-and-forget commands ------------------------------------------- - - def send( - self, - message: str, - ws_id: str = "", - name: str = "", - auto_approve: bool = False, - auto_approve_tools: list[str] | None = None, - target_node: str = "", - ) -> str: - """Send a message. Returns correlation_id for tracking. - - If *target_node* is set, the message is pushed to that node's - dedicated queue. If *ws_id* is set and *target_node* is not, - the client looks up the workstream's owning node and routes - accordingly. - """ - msg = SendMessage( - message=message, - ws_id=ws_id, - name=name, - auto_approve=auto_approve, - auto_approve_tools=auto_approve_tools or [], - target_node=target_node, - ) - node = target_node or (self._broker.get_ws_owner(ws_id) if ws_id else "") or "" - self._broker.push_inbound(msg.to_json(), node_id=node) - return msg.correlation_id - - def create_workstream( - self, - name: str = "", - auto_approve: bool = False, - auto_approve_tools: list[str] | None = None, - target_node: str = "", - initial_message: str = "", - skill: str = "", - ) -> str: - """Create a workstream. Returns correlation_id.""" - msg = CreateWorkstreamMessage( - name=name, - auto_approve=auto_approve, - auto_approve_tools=auto_approve_tools or [], - target_node=target_node, - initial_message=initial_message, - skill=skill, - ) - self._broker.push_inbound(msg.to_json(), node_id=target_node) - return msg.correlation_id - - def close_workstream(self, ws_id: str) -> str: - """Close a workstream. Returns correlation_id.""" - msg = CloseWorkstreamMessage(ws_id=ws_id) - self._broker.push_inbound(msg.to_json()) - return msg.correlation_id - - def command(self, ws_id: str, command: str) -> str: - """Execute a slash command. Returns correlation_id.""" - msg = CommandMessage(ws_id=ws_id, command=command) - self._broker.push_inbound(msg.to_json()) - return msg.correlation_id - - def list_workstreams(self) -> str: - """Request workstream list. Returns correlation_id.""" - msg = ListWorkstreamsMessage() - self._broker.push_inbound(msg.to_json()) - return msg.correlation_id - - def health(self) -> str: - """Request health status. Returns correlation_id.""" - msg = HealthMessage() - self._broker.push_inbound(msg.to_json()) - return msg.correlation_id - - def list_nodes(self) -> list[dict[str, Any]]: - """List active bridge nodes (reads directly from broker).""" - return self._broker.list_nodes() - - # -- approval / plan response ------------------------------------------- - - def approve( - self, - request_id: str, - ws_id: str = "", - approved: bool = True, - feedback: str | None = None, - always: bool = False, - ) -> None: - """Respond to a tool approval request.""" - msg = ApproveMessage( - ws_id=ws_id, - request_id=request_id, - approved=approved, - feedback=feedback, - always=always, - ) - self._broker.push_response(request_id, msg.to_json()) - - def plan_feedback( - self, - request_id: str, - ws_id: str = "", - feedback: str = "", - ) -> None: - """Respond to a plan review request.""" - msg = PlanFeedbackMessage( - ws_id=ws_id, - request_id=request_id, - feedback=feedback, - ) - self._broker.push_response(request_id, msg.to_json()) - - # -- blocking send ------------------------------------------------------- - - def send_and_wait( - self, - message: str, - ws_id: str = "", - name: str = "", - auto_approve: bool = True, - auto_approve_tools: list[str] | None = None, - target_node: str = "", - timeout: float = 600, - on_event: Callable[[OutboundEvent], None] | None = None, - ) -> TurnResult: - """Send a message and block until the turn completes. - - Returns a TurnResult with aggregated content, tool results, etc. - """ - # Build the message but don't send yet — subscribe first to avoid - # a race where the bridge processes the message before we subscribe. - msg = SendMessage( - message=message, - ws_id=ws_id, - name=name, - auto_approve=auto_approve, - auto_approve_tools=auto_approve_tools or [], - target_node=target_node, - ) - cid = msg.correlation_id - - result = TurnResult(correlation_id=cid, ws_id=ws_id) - done = threading.Event() - actual_ws_id = ws_id - - def _on_global(raw: str) -> None: - nonlocal actual_ws_id - event = OutboundEvent.from_json(raw) - if on_event: - on_event(event) - - if isinstance(event, WorkstreamCreatedEvent) and event.correlation_id == cid: - actual_ws_id = event.ws_id - result.ws_id = event.ws_id - self._broker.subscribe_outbound(f"{self._prefix}:events:{actual_ws_id}", _on_ws) - - def _on_ws(raw: str) -> None: - event = OutboundEvent.from_json(raw) - if on_event: - on_event(event) - - if isinstance(event, ContentEvent): - result.content_parts.append(event.text) - elif isinstance(event, ReasoningEvent): - result.reasoning_parts.append(event.text) - elif isinstance(event, ToolResultEvent): - result.tool_results.append((event.name, event.output)) - elif isinstance(event, ErrorEvent): - result.errors.append(event.message) - elif isinstance(event, TurnCompleteEvent) and event.correlation_id == cid: - done.set() - - # Subscribe BEFORE pushing — ensures we don't miss early events - self._broker.subscribe_outbound(f"{self._prefix}:events:global", _on_global) - if actual_ws_id: - self._broker.subscribe_outbound(f"{self._prefix}:events:{actual_ws_id}", _on_ws) - - # Now push the message (route to target node or ws owner if known) - node = target_node or (self._broker.get_ws_owner(ws_id) if ws_id else "") or "" - self._broker.push_inbound(msg.to_json(), node_id=node) - - done.wait(timeout=timeout) - - # Cleanup - self._broker.unsubscribe_outbound(f"{self._prefix}:events:global") - if actual_ws_id: - self._broker.unsubscribe_outbound(f"{self._prefix}:events:{actual_ws_id}") - - result.ws_id = actual_ws_id - result.timed_out = not done.is_set() - return result - - # -- subscription -------------------------------------------------------- - - def subscribe( - self, - callback: Callable[[OutboundEvent], None], - ws_id: str = "", - ) -> None: - """Subscribe to events for a specific workstream or global events.""" - channel = f"{self._prefix}:events:{ws_id}" if ws_id else f"{self._prefix}:events:global" - - def _cb(raw: str) -> None: - event = OutboundEvent.from_json(raw) - callback(event) - - self._broker.subscribe_outbound(channel, _cb) - - def unsubscribe(self, ws_id: str = "") -> None: - """Unsubscribe from a workstream or global channel.""" - channel = f"{self._prefix}:events:{ws_id}" if ws_id else f"{self._prefix}:events:global" - self._broker.unsubscribe_outbound(channel) - - # -- lifecycle ----------------------------------------------------------- - - def close(self) -> None: - """Clean up broker connection.""" - self._broker.close() - - def __enter__(self) -> TurnstoneClient: - return self - - def __exit__(self, *exc: object) -> None: - self.close() diff --git a/turnstone/mq/protocol.py b/turnstone/mq/protocol.py deleted file mode 100644 index 66d0e041..00000000 --- a/turnstone/mq/protocol.py +++ /dev/null @@ -1,482 +0,0 @@ -"""Message protocol for turnstone message queue integration. - -Defines all structured message types exchanged between the client and bridge. -Inbound messages flow from client → bridge via a reliable queue. -Outbound events flow from bridge → client via pub/sub channels. -""" - -from __future__ import annotations - -import json -import time -import uuid -from dataclasses import asdict, dataclass, field, fields -from typing import Any - -# --------------------------------------------------------------------------- -# Inbound messages (client → bridge) -# --------------------------------------------------------------------------- - - -@dataclass -class InboundMessage: - """Base for all messages sent by clients to the bridge.""" - - type: str = "" - correlation_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12]) - timestamp: float = field(default_factory=time.time) - - def to_json(self) -> str: - return json.dumps(asdict(self)) - - @classmethod - def from_json(cls, raw: str) -> InboundMessage: - data = json.loads(raw) - msg_type = data.get("type", "") - klass = _INBOUND_REGISTRY.get(msg_type) - if klass is None: - raise ValueError(f"Unknown inbound message type: {msg_type!r}") - valid = {f.name for f in fields(klass)} - return klass(**{k: v for k, v in data.items() if k in valid}) - - -@dataclass -class SendMessage(InboundMessage): - """Send a user message to a workstream.""" - - type: str = "send" - ws_id: str = "" - message: str = "" - auto_approve: bool = False - auto_approve_tools: list[str] = field(default_factory=list) - name: str = "" - target_node: str = "" - - -@dataclass -class ApproveMessage(InboundMessage): - """Respond to a tool approval request.""" - - type: str = "approve" - ws_id: str = "" - request_id: str = "" - approved: bool = True - feedback: str | None = None - always: bool = False - - -@dataclass -class PlanFeedbackMessage(InboundMessage): - """Respond to a plan review request.""" - - type: str = "plan_feedback" - ws_id: str = "" - request_id: str = "" - feedback: str = "" - - -@dataclass -class CommandMessage(InboundMessage): - """Execute a slash command.""" - - type: str = "command" - ws_id: str = "" - command: str = "" - - -@dataclass -class CreateWorkstreamMessage(InboundMessage): - """Create a new workstream.""" - - type: str = "create_workstream" - name: str = "" - auto_approve: bool = False - auto_approve_tools: list[str] = field(default_factory=list) - target_node: str = "" - model: str = "" - initial_message: str = "" - resume_ws: str = "" - user_id: str = "" - skill: str = "" - - -@dataclass -class CloseWorkstreamMessage(InboundMessage): - """Close a workstream.""" - - type: str = "close_workstream" - ws_id: str = "" - - -@dataclass -class ListWorkstreamsMessage(InboundMessage): - """Request the list of active workstreams.""" - - type: str = "list_workstreams" - - -@dataclass -class HealthMessage(InboundMessage): - """Request health status.""" - - type: str = "health" - - -@dataclass -class ListNodesMessage(InboundMessage): - """Request the list of active bridge nodes.""" - - type: str = "list_nodes" - - -@dataclass -class CancelMessage(InboundMessage): - """Cancel the active generation in a workstream.""" - - type: str = "cancel" - ws_id: str = "" - - -# --------------------------------------------------------------------------- -# Outbound events (bridge → client) -# --------------------------------------------------------------------------- - - -@dataclass -class OutboundEvent: - """Base for all events published by the bridge.""" - - type: str = "" - ws_id: str = "" - correlation_id: str = "" - timestamp: float = field(default_factory=time.time) - - def to_json(self) -> str: - return json.dumps(asdict(self)) - - @classmethod - def from_json(cls, raw: str) -> OutboundEvent: - data = json.loads(raw) - msg_type = data.get("type", "") - klass = _OUTBOUND_REGISTRY.get(msg_type, OutboundEvent) - valid = {f.name for f in fields(klass)} - return klass(**{k: v for k, v in data.items() if k in valid}) - - -@dataclass -class AckEvent(OutboundEvent): - """Acknowledgment that an inbound message was received.""" - - type: str = "ack" - status: str = "ok" - detail: str = "" - - -@dataclass -class ContentEvent(OutboundEvent): - """Streamed content token from the assistant.""" - - type: str = "content" - text: str = "" - - -@dataclass -class ReasoningEvent(OutboundEvent): - """Streamed reasoning token.""" - - type: str = "reasoning" - text: str = "" - - -@dataclass -class ToolInfoEvent(OutboundEvent): - """Tool call info (auto-approved tools).""" - - type: str = "tool_info" - items: list[dict[str, Any]] = field(default_factory=list) - - -@dataclass -class ApprovalRequestEvent(OutboundEvent): - """Tool approval request forwarded from the server. - - The client must respond with an ApproveMessage whose - request_id matches this event's correlation_id. - """ - - type: str = "approval_request" - items: list[dict[str, Any]] = field(default_factory=list) - - -@dataclass -class ToolOutputChunkEvent(OutboundEvent): - """Incremental streaming output from a bash tool.""" - - type: str = "tool_output_chunk" - call_id: str = "" - chunk: str = "" - - -@dataclass -class ToolResultEvent(OutboundEvent): - """Tool execution result.""" - - type: str = "tool_result" - call_id: str = "" - name: str = "" - output: str = "" - - -@dataclass -class PlanReviewEvent(OutboundEvent): - """Plan review request forwarded from the server. - - The client must respond with a PlanFeedbackMessage whose - request_id matches this event's correlation_id. - """ - - type: str = "plan_review" - content: str = "" - - -@dataclass -class StatusEvent(OutboundEvent): - """Token usage status update.""" - - type: str = "status" - prompt_tokens: int = 0 - completion_tokens: int = 0 - total_tokens: int = 0 - context_window: int = 0 - pct: float = 0.0 - effort: str = "" - cache_creation_tokens: int = 0 - cache_read_tokens: int = 0 - tool_calls_this_turn: int = 0 - turn_count: int = 0 - - -@dataclass -class StateChangeEvent(OutboundEvent): - """Workstream state transition.""" - - type: str = "state_change" - state: str = "" - - -@dataclass -class TurnCompleteEvent(OutboundEvent): - """Emitted when a workstream finishes processing (returns to IDLE). - - This is a synthetic event produced by the bridge when it detects - the ws_state transition to 'idle'. ``correlation_id`` is set for - MQ-initiated turns and empty for turns initiated from the server UI. - - ``content`` carries the full assistant response text piggybacked on - the server's idle SSE event (accumulated server-side in WebUI). - Downstream consumers (e.g. Discord bot) use it for catch-up when the - streaming path missed events, and as the primary delivery path for - bidirectional notification DM forwarding. - """ - - type: str = "turn_complete" - content: str = "" - - -@dataclass -class StreamEndEvent(OutboundEvent): - """LLM stream ended.""" - - type: str = "stream_end" - - -@dataclass -class WorkstreamCreatedEvent(OutboundEvent): - """New workstream created.""" - - type: str = "ws_created" - name: str = "" - node_id: str = "" - resumed: bool = False - message_count: int = 0 - - -@dataclass -class WorkstreamClosedEvent(OutboundEvent): - """Workstream closed.""" - - type: str = "ws_closed" - - -@dataclass -class WorkstreamListEvent(OutboundEvent): - """Workstream list response.""" - - type: str = "ws_list" - workstreams: list[dict[str, Any]] = field(default_factory=list) - - -@dataclass -class WorkstreamRenameEvent(OutboundEvent): - """Workstream renamed.""" - - type: str = "ws_rename" - name: str = "" - - -@dataclass -class HealthResponseEvent(OutboundEvent): - """Health status response.""" - - type: str = "health_response" - data: dict[str, Any] = field(default_factory=dict) - - -@dataclass -class ErrorEvent(OutboundEvent): - """Error event.""" - - type: str = "error" - message: str = "" - - -@dataclass -class InfoEvent(OutboundEvent): - """Informational event.""" - - type: str = "info" - message: str = "" - - -@dataclass -class NodeListEvent(OutboundEvent): - """List of active bridge nodes.""" - - type: str = "node_list" - nodes: list[dict[str, Any]] = field(default_factory=list) - - -@dataclass -class WorkstreamResumedEvent(OutboundEvent): - """Confirmation that a workstream was resumed during creation.""" - - type: str = "ws_resumed" - message_count: int = 0 - name: str = "" - - -@dataclass -class ClusterStateEvent(OutboundEvent): - """Workstream state change with node attribution for cluster dashboard.""" - - type: str = "cluster_state" - ws_id: str = "" - state: str = "" - node_id: str = "" - tokens: int = 0 - context_ratio: float = 0.0 - activity: str = "" - activity_state: str = "" - - -@dataclass -class IntentVerdictEvent(OutboundEvent): - """Intent validation verdict for a pending tool approval.""" - - type: str = "intent_verdict" - call_id: str = "" - func_name: str = "" - intent_summary: str = "" - risk_level: str = "" - confidence: float = 0.0 - recommendation: str = "" - reasoning: str = "" - evidence: str = "[]" # JSON array string - tier: str = "" - judge_model: str = "" - verdict_id: str = "" - latency_ms: int = 0 - - -@dataclass -class OutputWarningEvent(OutboundEvent): - """Output guard warning for tool execution result.""" - - type: str = "output_warning" - call_id: str = "" - func_name: str = "" - risk_level: str = "none" - flags: str = "[]" - annotations: str = "[]" - redacted: int = 0 - - -@dataclass -class ConfigChangeEvent(OutboundEvent): - """System setting changed — nodes should invalidate config cache.""" - - type: str = "config_change" - key: str = "" - node_id: str = "" - action: str = "" # "set" | "delete" - - -# --------------------------------------------------------------------------- -# Type registries (built after all classes are defined) -# --------------------------------------------------------------------------- - - -def _type_default(cls: type[Any]) -> str: - """Return the default value of the 'type' field for a dataclass.""" - for f in fields(cls): - if f.name == "type": - return f.default # type: ignore[return-value] - return "" - - -_INBOUND_REGISTRY: dict[str, type[InboundMessage]] = { - _type_default(cls): cls - for cls in [ - SendMessage, - ApproveMessage, - PlanFeedbackMessage, - CommandMessage, - CreateWorkstreamMessage, - CloseWorkstreamMessage, - ListWorkstreamsMessage, - HealthMessage, - ListNodesMessage, - CancelMessage, - ] -} - -_OUTBOUND_REGISTRY: dict[str, type[OutboundEvent]] = { - _type_default(cls): cls - for cls in [ - AckEvent, - ContentEvent, - ReasoningEvent, - ToolInfoEvent, - ApprovalRequestEvent, - ToolOutputChunkEvent, - ToolResultEvent, - PlanReviewEvent, - StatusEvent, - StateChangeEvent, - TurnCompleteEvent, - StreamEndEvent, - WorkstreamCreatedEvent, - WorkstreamClosedEvent, - WorkstreamListEvent, - WorkstreamRenameEvent, - HealthResponseEvent, - ErrorEvent, - InfoEvent, - NodeListEvent, - WorkstreamResumedEvent, - ClusterStateEvent, - IntentVerdictEvent, - OutputWarningEvent, - ConfigChangeEvent, - ] -} diff --git a/turnstone/sdk/_types.py b/turnstone/sdk/_types.py index 0e98a562..48cba6cb 100644 --- a/turnstone/sdk/_types.py +++ b/turnstone/sdk/_types.py @@ -9,8 +9,8 @@ from dataclasses import dataclass, field class TurnResult: """Aggregated result of a send_and_wait call. - Mirrors the shape of :class:`turnstone.mq.client.TurnResult` but - operates over HTTP/SSE instead of Redis pub/sub. + Collects content, reasoning, tool results, and errors from + an HTTP/SSE event stream into a single result object. """ ws_id: str = "" diff --git a/turnstone/sdk/events.py b/turnstone/sdk/events.py index a24e49ea..ccbcb65c 100644 --- a/turnstone/sdk/events.py +++ b/turnstone/sdk/events.py @@ -1,9 +1,7 @@ """Standalone SSE event dataclasses for the turnstone SDK. These types match the JSON payloads emitted by the server and console -SSE endpoints. They are intentionally decoupled from the MQ protocol -events in ``turnstone.mq.protocol`` so that SDK consumers do not need -the ``redis`` optional dependency. +SSE endpoints. """ from __future__ import annotations @@ -92,6 +90,7 @@ class ToolInfoEvent(ServerEvent): class ApproveRequestEvent(ServerEvent): type: str = "approve_request" items: list[dict[str, Any]] = field(default_factory=list) + judge_pending: bool = False @dataclass @@ -166,6 +165,34 @@ class CancelledEvent(ServerEvent): type: str = "cancelled" +@dataclass +class IntentVerdictEvent(ServerEvent): + type: str = "intent_verdict" + tool_name: str = "" + verdict: str = "" + reason: str = "" + call_id: str = "" + func_name: str = "" + intent_summary: str = "" + risk_level: str = "" + confidence: float = 0.0 + recommendation: str = "" + reasoning: str = "" + tier: str = "" + judge_model: str = "" + verdict_id: str = "" + latency_ms: int = 0 + + +@dataclass +class OutputWarningEvent(ServerEvent): + type: str = "output_warning" + call_id: str = "" + risk_level: str = "" + categories: list[str] = field(default_factory=list) + explanation: str = "" + + # --------------------------------------------------------------------------- # Server global events (/v1/api/events/global) # --------------------------------------------------------------------------- @@ -308,6 +335,8 @@ _SERVER_REGISTRY: dict[str, type[ServerEvent]] = { BusyErrorEvent, ClearUiEvent, CancelledEvent, + IntentVerdictEvent, + OutputWarningEvent, WsStateEvent, WsActivityEvent, WsRenameEvent, diff --git a/turnstone/server.py b/turnstone/server.py index 767a53e9..234fb201 100644 --- a/turnstone/server.py +++ b/turnstone/server.py @@ -18,6 +18,7 @@ import functools import json import os import queue +import re import sys import textwrap import threading @@ -61,6 +62,7 @@ log = get_logger(__name__) _STATIC_DIR = Path(__file__).parent / "ui" / "static" _SHARED_DIR = Path(__file__).parent / "shared_static" _HTML = (_STATIC_DIR / "index.html").read_text(encoding="utf-8") +_VALID_WS_ID = re.compile(r"^[0-9a-f]{32}$") # --------------------------------------------------------------------------- @@ -1520,6 +1522,11 @@ async def create_workstream(request: Request) -> JSONResponse: _st = _get_storage() applied_skill_version = len(_st.list_skill_versions(skill_data["template_id"])) + 1 + requested_ws_id = body.get("ws_id", "") or "" + if not isinstance(requested_ws_id, str): + requested_ws_id = "" + if requested_ws_id and not _VALID_WS_ID.match(requested_ws_id): + return JSONResponse({"error": "invalid ws_id format"}, status_code=400) try: ws = mgr.create( name=body.get("name", ""), @@ -1528,6 +1535,7 @@ async def create_workstream(request: Request) -> JSONResponse: skill=resolved_skill, skill_id=skill_data["template_id"] if skill_data else "", skill_version=applied_skill_version, + ws_id=requested_ws_id, ) assert isinstance(ws.ui, WebUI) if skip or body.get("auto_approve", False): @@ -2117,8 +2125,40 @@ async def _lifespan(app: Starlette) -> AsyncGenerator[None, None]: except Exception: log.warning("TLS auto-renewal startup failed", exc_info=True) + # Register in service registry and start heartbeat + _heartbeat_task: asyncio.Task[None] | None = None + _svc_node_id: str = getattr(app.state, "node_id", "") + _svc_url: str = getattr(app.state, "advertise_url", "") + if _svc_node_id and _svc_url: + from turnstone.core.storage import get_storage as _get_svc_storage + + _svc_storage = _get_svc_storage() + _svc_storage.register_service("server", _svc_node_id, _svc_url) + log.info("server.service_registered", node_id=_svc_node_id, url=_svc_url) + + async def _heartbeat_loop() -> None: + """Periodically update service heartbeat.""" + while True: + await asyncio.sleep(30) + try: + await asyncio.to_thread(_svc_storage.heartbeat_service, "server", _svc_node_id) + except Exception: + log.exception("server.heartbeat_failed") + + _heartbeat_task = asyncio.create_task(_heartbeat_loop()) + yield # Shutdown + if _heartbeat_task is not None: + _heartbeat_task.cancel() + if _svc_node_id and _svc_url: + from turnstone.core.storage import get_storage as _get_svc_dereg + + try: + await asyncio.to_thread(_get_svc_dereg().deregister_service, "server", _svc_node_id) + log.info("server.service_deregistered", node_id=_svc_node_id) + except Exception: + log.exception("server.deregister_failed") tls_client = getattr(app.state, "tls_client", None) if tls_client is not None: await tls_client.stop_renewal() @@ -2177,6 +2217,7 @@ def create_app( watch_runner: Any = None, judge_config: Any = None, config_store: Any = None, + advertise_url: str = "", ) -> Starlette: """Create and configure the Starlette ASGI application.""" _spec = build_server_spec() @@ -2254,6 +2295,7 @@ def create_app( app.state.watch_runner = watch_runner app.state.judge_config = judge_config app.state.config_store = config_store + app.state.advertise_url = advertise_url from turnstone.core.auth import LoginRateLimiter @@ -2703,6 +2745,10 @@ def main() -> None: cors_origins = parse_cors_origins() + # Construct advertise URL for service registration + _advertise_host = socket.gethostname() if args.host in ("0.0.0.0", "::") else args.host + _advertise_url = f"http://{_advertise_host}:{args.port}" + _skip_perms = config_store.get("tools.skip_permissions") app = create_app( workstreams=manager, @@ -2723,6 +2769,7 @@ def main() -> None: watch_runner=_watch_runner, judge_config=judge_config, config_store=config_store, + advertise_url=_advertise_url, ) # Store CLI model args for hot-reload (internal_model_reload reads these) @@ -2797,6 +2844,8 @@ def main() -> None: # Store client on app state for lifespan renewal app.state.tls_client = tls_client + # Update advertise URL to HTTPS now that TLS is active + app.state.advertise_url = f"https://{_advertise_host}:{args.port}" log.info("TLS enabled — serving HTTPS") else: log.warning("TLS enabled but no cert available") diff --git a/turnstone/sim/__init__.py b/turnstone/sim/__init__.py deleted file mode 100644 index b07a025d..00000000 --- a/turnstone/sim/__init__.py +++ /dev/null @@ -1,6 +0,0 @@ -"""Turnstone cluster simulator.""" - -from turnstone.sim.cluster import SimCluster -from turnstone.sim.config import SimConfig - -__all__ = ["SimCluster", "SimConfig"] diff --git a/turnstone/sim/cli.py b/turnstone/sim/cli.py deleted file mode 100644 index f96d6f72..00000000 --- a/turnstone/sim/cli.py +++ /dev/null @@ -1,181 +0,0 @@ -"""CLI entry point for turnstone-sim.""" - -from __future__ import annotations - -import argparse -import asyncio -import json -import logging -import sys -from typing import Any - -from turnstone.sim.cluster import SimCluster -from turnstone.sim.config import SimConfig -from turnstone.sim.scenario import SCENARIOS - -log = logging.getLogger("turnstone.sim") - - -def main() -> None: - parser = argparse.ArgumentParser( - prog="turnstone-sim", - description="Turnstone multi-node cluster simulator", - ) - parser.add_argument( - "--nodes", - type=int, - default=10, - help="Number of simulated nodes (default: 10)", - ) - parser.add_argument( - "--scenario", - choices=list(SCENARIOS.keys()), - default="steady", - help="Scenario to run (default: steady)", - ) - parser.add_argument( - "--duration", - type=int, - default=60, - help="Scenario duration in seconds (default: 60)", - ) - parser.add_argument( - "--mps", - type=float, - default=5.0, - help="Messages per second for steady scenario (default: 5.0)", - ) - parser.add_argument( - "--burst-size", - type=int, - default=100, - help="Message count for burst scenario (default: 100)", - ) - parser.add_argument( - "--llm-latency", - type=float, - default=2.0, - help="Mean LLM response latency in seconds (default: 2.0)", - ) - parser.add_argument( - "--tool-latency", - type=float, - default=0.5, - help="Mean tool execution latency in seconds (default: 0.5)", - ) - parser.add_argument( - "--tool-failure-rate", - type=float, - default=0.02, - help="Tool failure probability 0.0-1.0 (default: 0.02)", - ) - parser.add_argument( - "--node-kill-interval", - type=float, - default=15.0, - help="Seconds between node kills for node_failure scenario (default: 15)", - ) - parser.add_argument( - "--node-kill-count", - type=int, - default=1, - help="Nodes to kill per interval (default: 1)", - ) - from turnstone.mq.broker import add_redis_args - - add_redis_args(parser) - parser.add_argument("--prefix", default="turnstone") - parser.add_argument("--seed", type=int, default=None, help="Random seed for reproducibility") - parser.add_argument("--metrics-file", default="", help="Write JSON metrics to file") - from turnstone.core.log import add_log_args - - add_log_args(parser) - - args = parser.parse_args() - - config = SimConfig( - num_nodes=args.nodes, - scenario=args.scenario, - duration=args.duration, - messages_per_second=args.mps, - burst_size=args.burst_size, - llm_latency_mean=args.llm_latency, - tool_latency_mean=args.tool_latency, - tool_failure_rate=args.tool_failure_rate, - node_kill_interval=args.node_kill_interval, - node_kill_count=args.node_kill_count, - redis_host=args.redis_host, - redis_port=args.redis_port, - redis_password=args.redis_password, - redis_db=args.redis_db, - prefix=args.prefix, - seed=args.seed, - metrics_file=args.metrics_file, - ) - - from turnstone.core.log import configure_logging_from_args - - configure_logging_from_args(args, "sim") - - try: - asyncio.run(_run(config)) - except KeyboardInterrupt: - log.info("Interrupted") - sys.exit(0) - - -async def _run(config: SimConfig) -> None: - cluster = SimCluster(config) - try: - await cluster.start() - log.info( - "Running scenario=%s nodes=%d duration=%ds", - config.scenario, - config.num_nodes, - config.duration, - ) - await cluster.run_scenario() - - report = cluster.report() - _print_report(report, config) - - if config.metrics_file: - with open(config.metrics_file, "w") as f: - json.dump(report, f, indent=2) - log.info("Metrics written to %s", config.metrics_file) - finally: - await cluster.stop() - - -def _print_report(report: dict[str, Any], config: SimConfig) -> None: - lat = report.get("latency", {}) - tp = report.get("throughput", {}) - util = report.get("utilization", {}) - - print("\n" + "=" * 60) - print(" SIMULATION REPORT") - print("=" * 60) - print(f" Scenario: {config.scenario}") - print(f" Nodes: {config.num_nodes}") - print(f" Duration: {report['duration_seconds']}s") - print(f" Total turns: {report['total_turns']}") - print(f" Total errors: {report['total_errors']}") - print(f" Node kills: {report['node_kills']}") - print("-" * 60) - print(" THROUGHPUT") - print(f" Messages/sec: {tp.get('messages_per_sec', 0)}") - print(f" Turns/sec: {tp.get('turns_per_sec', 0)}") - print("-" * 60) - print(" LATENCY (seconds)") - print(f" p50: {lat.get('p50', 0)}") - print(f" p90: {lat.get('p90', 0)}") - print(f" p99: {lat.get('p99', 0)}") - print(f" mean: {lat.get('mean', 0)}") - print(f" max: {lat.get('max', 0)}") - if util: - print("-" * 60) - print(" UTILIZATION") - print(f" Mean ws/node: {util.get('mean_ws_per_node', 0):.1f}") - print(f" Max ws/node: {util.get('max_ws_per_node', 0)}") - print(f" Idle nodes: {util.get('nodes_with_zero_ws', 0)}") - print("=" * 60 + "\n") diff --git a/turnstone/sim/cluster.py b/turnstone/sim/cluster.py deleted file mode 100644 index 0ae4c63f..00000000 --- a/turnstone/sim/cluster.py +++ /dev/null @@ -1,321 +0,0 @@ -"""Cluster orchestration — manages N SimNodes, dispatchers, and metrics.""" - -from __future__ import annotations - -import asyncio -import logging -import math -import time -from concurrent.futures import ThreadPoolExecutor -from typing import TYPE_CHECKING, Any, cast - -import redis - -from turnstone.mq.broker import RedisBroker -from turnstone.sim.metrics import MetricsCollector -from turnstone.sim.node import SimNode - -if TYPE_CHECKING: - from turnstone.sim.config import SimConfig - -log = logging.getLogger("turnstone.sim.cluster") - -# How many node queues a single dispatcher watches via one BLPOP call. -NODES_PER_DISPATCHER = 50 - - -class PooledBroker(RedisBroker): - """RedisBroker that uses a shared external ConnectionPool.""" - - def __init__( - self, - pool: redis.ConnectionPool, - prefix: str = "turnstone", - response_ttl: int = 600, - ) -> None: - # Bypass RedisBroker.__init__ — set up manually with the shared pool. - - self._prefix = prefix - self._response_ttl = response_ttl - self._pool: redis.ConnectionPool = pool - self._redis: redis.Redis[str] = cast("redis.Redis[str]", redis.Redis(connection_pool=pool)) - self._pubsub = self._redis.pubsub(ignore_subscribe_messages=True) - self._listener_thread: Any = None - self._running = True - - def close(self) -> None: - """No-op — the shared pool is managed by SimCluster.""" - self._running = False - - -class InboundDispatcher: - """Watches batches of node queues via a single BLPOP call. - - Instead of one BLPOP per node (which would exhaust Redis connections at - 1000 nodes), a dispatcher batches ~50 node queues into a single BLPOP - on multiple keys. This keeps total Redis connections bounded. - """ - - def __init__( - self, - redis_client: redis.Redis[str], - node_ids: list[str], - nodes: dict[str, SimNode], - prefix: str, - ) -> None: - self._redis = redis_client - self._node_ids = node_ids - self._nodes = nodes - self._prefix = prefix - self._running = True - - # Build BLPOP key list: per-node queues first (priority), shared last - self._keys = [f"{prefix}:inbound:{nid}" for nid in node_ids] - self._keys.append(f"{prefix}:inbound") - - # Pre-compute key → node_id mapping - self._key_to_node: dict[str, str] = {f"{prefix}:inbound:{nid}": nid for nid in node_ids} - - async def run(self) -> None: - while self._running: - # Snapshot keys to avoid race with remove_node() during BLPOP - keys = list(self._keys) - if not keys: - await asyncio.sleep(0.5) - continue - result = await asyncio.to_thread( - self._redis.blpop, - keys, - timeout=1, - ) - if result is None: - continue - - queue_key, raw = result - if isinstance(queue_key, bytes): - queue_key = queue_key.decode() - if isinstance(raw, bytes): - raw = raw.decode() - - node = self._resolve_target(queue_key) - if node and node._running: - await node.handle_message(raw) - - def _resolve_target(self, queue_key: str) -> SimNode | None: - """Determine which SimNode should handle this message.""" - node_id = self._key_to_node.get(queue_key) - if node_id: - return self._nodes.get(node_id) - - # Shared queue — pick running node with fewest workstreams and capacity - if self._nodes: - candidates = [ - n - for n in self._nodes.values() - if n._running and n.workstream_count < n._config.max_ws_per_node - ] - if candidates: - return min(candidates, key=lambda n: n.workstream_count) - # Fall back to any running node if all at capacity - running = [n for n in self._nodes.values() if n._running] - if running: - return min(running, key=lambda n: n.workstream_count) - return None - - def stop(self) -> None: - self._running = False - - def remove_node(self, node_id: str) -> None: - """Remove a node from this dispatcher (for kill simulation).""" - self._nodes.pop(node_id, None) - key = f"{self._prefix}:inbound:{node_id}" - self._key_to_node.pop(key, None) - if key in self._keys: - self._keys.remove(key) - - -class SimCluster: - """Orchestrates N SimNodes, dispatchers, heartbeats, and metrics. - - Usage:: - - cluster = SimCluster(config) - await cluster.start() - await cluster.run_scenario() - report = cluster.report() - await cluster.stop() - """ - - def __init__(self, config: SimConfig) -> None: - self._config = config - self._metrics = MetricsCollector() - self._nodes: dict[str, SimNode] = {} - self._node_order: list[str] = [] - self._dispatchers: list[InboundDispatcher] = [] - self._tasks: list[asyncio.Task[None]] = [] - self._pool: redis.ConnectionPool | None = None - self._redis_client: redis.Redis[str] | None = None - self._executor: ThreadPoolExecutor | None = None - self._running = True - - @property - def metrics(self) -> MetricsCollector: - return self._metrics - - @property - def nodes(self) -> dict[str, SimNode]: - return self._nodes - - @property - def config(self) -> SimConfig: - return self._config - - async def start(self) -> None: - """Create connection pool, nodes, dispatchers; start all tasks.""" - self._executor = ThreadPoolExecutor(max_workers=64) - - # Shared Redis pool - pool: redis.ConnectionPool = redis.ConnectionPool( - host=self._config.redis_host, - port=self._config.redis_port, - db=self._config.redis_db, - password=self._config.redis_password, - decode_responses=True, - retry_on_timeout=True, - max_connections=64, - ) - self._pool = pool - self._redis_client = cast("redis.Redis[str]", redis.Redis(connection_pool=pool)) - - # Create nodes - for i in range(self._config.num_nodes): - node_id = f"sim-{i:04d}" - broker = PooledBroker( - pool, - prefix=self._config.prefix, - ) - node = SimNode(node_id, broker, self._config, self._metrics) - self._nodes[node_id] = node - self._node_order.append(node_id) - - # Create dispatchers (batches of NODES_PER_DISPATCHER) - all_ids = list(self._nodes.keys()) - num_dispatchers = max(1, math.ceil(len(all_ids) / NODES_PER_DISPATCHER)) - for i in range(num_dispatchers): - start = i * NODES_PER_DISPATCHER - batch_ids = all_ids[start : start + NODES_PER_DISPATCHER] - # Each dispatcher gets its own Redis client from the shared pool - client: redis.Redis[str] = cast("redis.Redis[str]", redis.Redis(connection_pool=pool)) - dispatcher = InboundDispatcher( - client, - batch_ids, - dict(self._nodes), - self._config.prefix, - ) - self._dispatchers.append(dispatcher) - self._tasks.append(asyncio.create_task(dispatcher.run())) - - # Start heartbeat task - self._tasks.append(asyncio.create_task(self._heartbeat_loop())) - - # Start utilization snapshot task - self._tasks.append(asyncio.create_task(self._utilization_loop())) - - # Wait for all nodes to register - await self._wait_for_nodes() - log.info( - "Cluster started: %d nodes, %d dispatchers", - len(self._nodes), - len(self._dispatchers), - ) - - async def _heartbeat_loop(self) -> None: - """Register heartbeats for all running nodes concurrently.""" - interval = max(1, self._config.heartbeat_ttl // 2) - loop = asyncio.get_running_loop() - while self._running: - tasks = [ - loop.run_in_executor(self._executor, node.heartbeat_once) - for node in self._nodes.values() - if node._running - ] - if tasks: - await asyncio.gather(*tasks, return_exceptions=True) - await asyncio.sleep(interval) - - async def _utilization_loop(self) -> None: - """Periodically snapshot workstream utilization.""" - while self._running: - await asyncio.sleep(self._config.metrics_interval) - counts = { - nid: node.workstream_count for nid, node in self._nodes.items() if node._running - } - self._metrics.snapshot_utilization(counts) - - async def _wait_for_nodes(self) -> None: - """Do an initial heartbeat and confirm registration.""" - loop = asyncio.get_running_loop() - tasks = [ - loop.run_in_executor(self._executor, node.heartbeat_once) - for node in self._nodes.values() - ] - await asyncio.gather(*tasks, return_exceptions=True) - - assert self._redis_client is not None - redis_client = self._redis_client - registered = 0 - deadline = time.monotonic() + 30 - while time.monotonic() < deadline: - keys = await loop.run_in_executor( - self._executor, - redis_client.keys, - f"{self._config.prefix}:node:sim-*", - ) - registered = len(keys) - if registered >= self._config.num_nodes: - return - await asyncio.sleep(0.5) - raise TimeoutError( - f"Only {registered}/{self._config.num_nodes} nodes registered", - ) - - async def run_scenario(self) -> None: - """Run the configured scenario.""" - from turnstone.sim.scenario import SCENARIOS - - scenario_cls = SCENARIOS.get(self._config.scenario) - if scenario_cls is None: - raise ValueError(f"Unknown scenario: {self._config.scenario!r}") - scenario = scenario_cls() - await scenario.run(self, self._config, self._metrics) - - async def kill_node(self, node_id: str) -> None: - """Simulate a node failure: stop heartbeat, stop processing.""" - node = self._nodes.get(node_id) - if node and node._running: - node.stop() - self._metrics.record_node_kill(node_id) - # Remove from dispatchers - for d in self._dispatchers: - d.remove_node(node_id) - log.info("Killed node %s", node_id) - - def report(self) -> dict[str, Any]: - """Generate final metrics report.""" - return self._metrics.summary() - - async def stop(self) -> None: - """Shutdown all nodes and cancel tasks.""" - self._running = False - for node in self._nodes.values(): - node.stop() - for d in self._dispatchers: - d.stop() - for task in self._tasks: - task.cancel() - await asyncio.gather(*self._tasks, return_exceptions=True) - if self._executor is not None: - self._executor.shutdown(wait=False) - if self._pool is not None: - self._pool.disconnect() - log.info("Cluster stopped") diff --git a/turnstone/sim/config.py b/turnstone/sim/config.py deleted file mode 100644 index 5bd2da03..00000000 --- a/turnstone/sim/config.py +++ /dev/null @@ -1,55 +0,0 @@ -"""Simulation configuration.""" - -from __future__ import annotations - -from dataclasses import dataclass - - -@dataclass(frozen=True) -class SimConfig: - """All parameters controlling a simulation run.""" - - # -- cluster -- - num_nodes: int = 10 - max_ws_per_node: int = 10 - - # -- redis -- - redis_host: str = "localhost" - redis_port: int = 6379 - redis_db: int = 0 - redis_password: str | None = None - prefix: str = "turnstone" - - # -- heartbeat -- - heartbeat_ttl: int = 60 - - # -- LLM simulation -- - llm_latency_mean: float = 2.0 - llm_latency_stddev: float = 0.5 - llm_tokens_mean: int = 200 - llm_tokens_stddev: int = 50 - llm_token_rate: float = 50.0 # tokens/sec streaming speed - context_window: int = 131072 # for computing context ratio - - # -- tool simulation -- - tool_latency_mean: float = 0.5 - tool_latency_stddev: float = 0.2 - tool_failure_rate: float = 0.02 - tool_calls_per_turn_mean: float = 1.5 - tool_calls_per_turn_max: int = 4 - max_tool_rounds: int = 3 - - # -- scenario -- - scenario: str = "steady" - duration: int = 60 - messages_per_second: float = 5.0 - burst_size: int = 100 - node_kill_interval: float = 15.0 - node_kill_count: int = 1 - - # -- metrics -- - metrics_interval: float = 5.0 - metrics_file: str = "" - - # -- reproducibility -- - seed: int | None = None diff --git a/turnstone/sim/engine.py b/turnstone/sim/engine.py deleted file mode 100644 index 02dc2b0a..00000000 --- a/turnstone/sim/engine.py +++ /dev/null @@ -1,139 +0,0 @@ -"""LLM and tool execution simulation engine.""" - -from __future__ import annotations - -import asyncio -import random -from typing import TYPE_CHECKING, Any - -if TYPE_CHECKING: - from turnstone.sim.config import SimConfig - -_WORD_POOL = [ - "the", - "result", - "shows", - "that", - "this", - "file", - "contains", - "function", - "data", - "analysis", - "implementation", - "code", - "completed", - "successfully", - "reviewed", - "output", - "processing", - "module", - "system", - "request", - "response", - "value", - "config", - "status", - "running", - "checked", - "verified", - "found", - "done", -] - -_TOOL_NAMES = [ - "bash", - "read_file", - "search", - "edit_file", - "write_file", - "math", - "web_fetch", -] - - -class ToolSimulationError(Exception): - """Raised when a simulated tool execution fails.""" - - -class SimEngine: - """Simulates LLM responses and tool execution with configurable distributions. - - Stateless — safe to share across workstreams on the same node. - """ - - def __init__(self, config: SimConfig, rng: random.Random | None = None): - self._config = config - self._rng = rng or random.Random(config.seed) - - async def simulate_llm_response(self, first_round: bool) -> tuple[str, list[dict[str, Any]]]: - """Simulate an LLM response. - - Returns ``(content_text, tool_calls)`` where *tool_calls* may be - empty (final answer) or a list of ``{"name": ..., "arguments": ...}`` - dicts. - """ - latency = max( - 0.05, - self._rng.gauss( - self._config.llm_latency_mean, - self._config.llm_latency_stddev, - ), - ) - await asyncio.sleep(latency) - - num_tokens = max( - 10, - int( - self._rng.gauss( - self._config.llm_tokens_mean, - self._config.llm_tokens_stddev, - ) - ), - ) - content = self._generate_content(num_tokens) - - # First round has a higher chance of tool calls; decreasing per round - tool_prob = 0.6 if first_round else 0.3 - if self._rng.random() < tool_prob: - num_calls = min( - max( - 1, - int( - self._rng.expovariate( - 1.0 / self._config.tool_calls_per_turn_mean, - ) - ), - ), - self._config.tool_calls_per_turn_max, - ) - calls = [ - { - "name": self._rng.choice(_TOOL_NAMES), - "arguments": '{"simulated": true}', - } - for _ in range(num_calls) - ] - return content, calls - - return content, [] - - async def simulate_tool_execution(self, tool_name: str) -> str: - """Simulate tool execution with latency and possible failure.""" - latency = max( - 0.01, - self._rng.gauss( - self._config.tool_latency_mean, - self._config.tool_latency_stddev, - ), - ) - await asyncio.sleep(latency) - - if self._rng.random() < self._config.tool_failure_rate: - raise ToolSimulationError(f"Simulated {tool_name} failure") - - return f"[sim] {tool_name} completed successfully" - - def _generate_content(self, num_tokens: int) -> str: - """Generate placeholder content of approximately *num_tokens* tokens.""" - return " ".join(self._rng.choices(_WORD_POOL, k=num_tokens)) diff --git a/turnstone/sim/metrics.py b/turnstone/sim/metrics.py deleted file mode 100644 index 94f795aa..00000000 --- a/turnstone/sim/metrics.py +++ /dev/null @@ -1,111 +0,0 @@ -"""Simulation metrics collection.""" - -from __future__ import annotations - -import threading -import time -from collections import defaultdict -from typing import Any - - -class MetricsCollector: - """Thread-safe metrics collector for simulation runs. - - Uses ``threading.Lock`` so it works from both sync and async contexts. - """ - - def __init__(self) -> None: - self._lock = threading.Lock() - self._turn_latencies: list[float] = [] - self._inject_times: list[float] = [] - self._complete_times: list[float] = [] - self._errors: int = 0 - self._error_details: list[tuple[float, str, str]] = [] - self._node_kills: list[tuple[float, str]] = [] - self._turns_per_node: dict[str, int] = defaultdict(int) - self._ws_counts: list[dict[str, int]] = [] # utilization snapshots - - def record_turn(self, ws_id: str, node_id: str, latency: float) -> None: - with self._lock: - self._turn_latencies.append(latency) - self._complete_times.append(time.monotonic()) - self._turns_per_node[node_id] += 1 - - def record_inject(self) -> None: - with self._lock: - self._inject_times.append(time.monotonic()) - - def record_error(self, node_id: str, message: str) -> None: - with self._lock: - self._errors += 1 - self._error_details.append((time.monotonic(), node_id, message)) - - def record_node_kill(self, node_id: str) -> None: - with self._lock: - self._node_kills.append((time.monotonic(), node_id)) - - def snapshot_utilization(self, ws_counts: dict[str, int]) -> None: - """Record workstream-per-node counts at a point in time.""" - with self._lock: - self._ws_counts.append(dict(ws_counts)) - - def summary(self) -> dict[str, Any]: - """Generate final metrics report with percentiles and aggregates.""" - with self._lock: - latencies = sorted(self._turn_latencies) - n = len(latencies) - - if n > 0 and self._inject_times and self._complete_times: - duration = self._complete_times[-1] - self._inject_times[0] - else: - duration = 0.0 - - # Utilization from latest snapshot - util: dict[str, Any] = {} - if self._ws_counts: - last = self._ws_counts[-1] - counts = list(last.values()) - if counts: - util = { - "mean_ws_per_node": sum(counts) / len(counts), - "max_ws_per_node": max(counts), - "nodes_with_zero_ws": sum(1 for c in counts if c == 0), - } - - return { - "total_turns": n, - "total_errors": self._errors, - "duration_seconds": round(duration, 2), - "throughput": { - "messages_per_sec": round( - len(self._inject_times) / duration, - 2, - ) - if duration > 0 - else 0, - "turns_per_sec": round( - n / duration, - 2, - ) - if duration > 0 - else 0, - }, - "latency": { - "p50": _percentile(latencies, 0.50), - "p90": _percentile(latencies, 0.90), - "p99": _percentile(latencies, 0.99), - "mean": round(sum(latencies) / n, 4) if n else 0, - "max": round(latencies[-1], 4) if n else 0, - }, - "utilization": util, - "node_kills": len(self._node_kills), - "turns_per_node": dict(self._turns_per_node), - } - - -def _percentile(sorted_values: list[float], pct: float) -> float: - if not sorted_values: - return 0.0 - idx = int(len(sorted_values) * pct) - idx = min(idx, len(sorted_values) - 1) - return round(sorted_values[idx], 4) diff --git a/turnstone/sim/node.py b/turnstone/sim/node.py deleted file mode 100644 index b1b21894..00000000 --- a/turnstone/sim/node.py +++ /dev/null @@ -1,443 +0,0 @@ -"""Simulated turnstone node. - -A SimNode replaces Bridge + Server + ChatSession with a lightweight async -coroutine that talks directly to Redis via the real RedisBroker. External -observers (TurnstoneClient, turnstone-console) see identical protocol -behaviour. -""" - -from __future__ import annotations - -import asyncio -import json -import logging -import time -import uuid -from typing import TYPE_CHECKING - -from turnstone.mq.protocol import ( - AckEvent, - ClusterStateEvent, - ContentEvent, - ErrorEvent, - HealthResponseEvent, - InboundMessage, - NodeListEvent, - OutboundEvent, - StateChangeEvent, - StatusEvent, - StreamEndEvent, - ToolResultEvent, - TurnCompleteEvent, - WorkstreamClosedEvent, - WorkstreamCreatedEvent, - WorkstreamListEvent, -) -from turnstone.sim.engine import SimEngine, ToolSimulationError - -if TYPE_CHECKING: - from turnstone.mq.broker import RedisBroker - from turnstone.sim.config import SimConfig - from turnstone.sim.metrics import MetricsCollector - -log = logging.getLogger("turnstone.sim.node") - - -class SimWorkstream: - """Lightweight workstream state machine.""" - - def __init__( - self, - ws_id: str, - name: str, - node: SimNode, - engine: SimEngine, - config: SimConfig, - ): - self.ws_id = ws_id - self.name = name - self.state = "idle" - self._node = node - self._engine = engine - self._config = config - self._turn_count = 0 - self._total_tokens = 0 # accumulated across turns - - async def process_turn(self, message: str, correlation_id: str) -> None: - """Simulate a complete turn: LLM stream -> optional tools -> final.""" - t_start = time.monotonic() - self._turn_count += 1 - - try: - rounds = 0 - while True: - # LLM thinking + streaming - self._set_state("thinking", correlation_id) - content, tool_calls = await self._engine.simulate_llm_response( - rounds == 0, - ) - await self._stream_content(content, correlation_id) - - if not tool_calls or rounds >= self._config.max_tool_rounds: - break - - # Tool execution - self._set_state("running", correlation_id) - for tc in tool_calls: - name = tc["name"] - try: - output = await self._engine.simulate_tool_execution(name) - except ToolSimulationError as exc: - output = f"Error: {exc}" - self._node._metrics.record_error( - self._node.node_id, - str(exc), - ) - self._node._publish_ws( - self.ws_id, - ToolResultEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - name=name, - output=output, - ), - ) - - rounds += 1 - - # Finished — publish status, idle, turn complete - self._publish_status(correlation_id) - self._set_state("idle", correlation_id) - self._node._publish_ws( - self.ws_id, - TurnCompleteEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - ), - ) - self._node._publish_global( - TurnCompleteEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - ), - ) - - except Exception as exc: - self._set_state("error", correlation_id) - self._node._publish_ws( - self.ws_id, - ErrorEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - message=str(exc), - ), - ) - self._node._metrics.record_error(self._node.node_id, str(exc)) - - finally: - latency = time.monotonic() - t_start - self._node._metrics.record_turn( - self.ws_id, - self._node.node_id, - latency, - ) - - async def _stream_content(self, text: str, correlation_id: str) -> None: - """Simulate token-by-token streaming.""" - if not text: - return - # Count tokens (~1 token per word) and accumulate - self._total_tokens += len(text.split()) - chunk_size = max(1, len(text) // 8) - token_delay = 1.0 / max(1, self._config.llm_token_rate) - for i in range(0, len(text), chunk_size): - chunk = text[i : i + chunk_size] - self._node._publish_ws( - self.ws_id, - ContentEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - text=chunk, - ), - ) - await asyncio.sleep(token_delay * len(chunk.split())) - self._node._publish_ws( - self.ws_id, - StreamEndEvent(ws_id=self.ws_id, correlation_id=correlation_id), - ) - - def _set_state(self, state: str, correlation_id: str) -> None: - self.state = state - self._node._publish_global( - StateChangeEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - state=state, - ), - ) - # prompt tokens ~= 2x completion tokens for a realistic ratio - total = self._total_tokens * 3 - ctx_ratio = round(total / self._config.context_window, 3) if total else 0.0 - self._node._publish_cluster( - ClusterStateEvent( - ws_id=self.ws_id, - state=state, - node_id=self._node.node_id, - tokens=total, - context_ratio=ctx_ratio, - ), - ) - - def _publish_status(self, correlation_id: str) -> None: - total = self._total_tokens * 3 # prompt ~= 2x completion - cw = self._config.context_window - self._node._publish_ws( - self.ws_id, - StatusEvent( - ws_id=self.ws_id, - correlation_id=correlation_id, - prompt_tokens=self._total_tokens * 2, - completion_tokens=self._total_tokens, - total_tokens=total, - context_window=cw, - pct=round(total / cw, 3) if cw else 0, - effort="medium", - ), - ) - - -class SimNode: - """A lightweight simulated turnstone node. - - Replaces Bridge + Server + ChatSession with direct Redis protocol - interaction. - """ - - def __init__( - self, - node_id: str, - broker: RedisBroker, - config: SimConfig, - metrics: MetricsCollector, - ): - self.node_id = node_id - self._broker = broker - self._config = config - self._metrics = metrics - # Derive per-node seed so each node has unique RNG sequences - import random - - node_seed = None - if config.seed is not None: - node_seed = hash((config.seed, node_id)) - self._engine = SimEngine(config, rng=random.Random(node_seed)) - self._workstreams: dict[str, SimWorkstream] = {} - self._running = True - self._started_at = time.time() - self._prefix = config.prefix - - @property - def workstream_count(self) -> int: - return len(self._workstreams) - - # -- message handling ---------------------------------------------------- - - async def handle_message(self, raw: str) -> None: - """Parse and dispatch an inbound message.""" - try: - msg = InboundMessage.from_json(raw) - await self._dispatch(msg) - except Exception as exc: - log.error("SimNode %s dispatch error: %s", self.node_id, exc) - self._publish_global(ErrorEvent(message=f"SimNode error: {exc}")) - - async def _dispatch(self, msg: InboundMessage) -> None: - handlers = { - "send": self._handle_send, - "create_workstream": self._handle_create_ws, - "close_workstream": self._handle_close_ws, - "list_workstreams": self._handle_list_ws, - "health": self._handle_health, - "list_nodes": self._handle_list_nodes, - } - handler = handlers.get(msg.type) - if handler: - await handler(msg) - else: - log.debug("SimNode %s ignoring message type: %s", self.node_id, msg.type) - - async def _handle_send(self, msg: InboundMessage) -> None: - ws_id = getattr(msg, "ws_id", "") - message = getattr(msg, "message", "") - cid = msg.correlation_id - - # Find or create workstream - if ws_id and ws_id in self._workstreams: - ws = self._workstreams[ws_id] - elif len(self._workstreams) >= self._config.max_ws_per_node: - self._publish_global( - ErrorEvent( - correlation_id=cid, - message=f"Node {self.node_id} at capacity ({self._config.max_ws_per_node} ws)", - ), - ) - return - else: - ws = self._create_workstream( - name=getattr(msg, "name", ""), - correlation_id=cid, - ) - - self._publish_ws( - ws.ws_id, - AckEvent(ws_id=ws.ws_id, correlation_id=cid, status="ok"), - ) - await ws.process_turn(message, cid) - - async def _handle_create_ws(self, msg: InboundMessage) -> None: - if len(self._workstreams) >= self._config.max_ws_per_node: - self._publish_global( - ErrorEvent( - correlation_id=msg.correlation_id, - message=f"Node {self.node_id} at capacity ({self._config.max_ws_per_node} ws)", - ), - ) - return - name = getattr(msg, "name", "") - ws = self._create_workstream(name=name, correlation_id=msg.correlation_id) - self._publish_ws( - ws.ws_id, - AckEvent(ws_id=ws.ws_id, correlation_id=msg.correlation_id, status="ok"), - ) - - async def _handle_close_ws(self, msg: InboundMessage) -> None: - ws_id = getattr(msg, "ws_id", "") - ws = self._workstreams.pop(ws_id, None) - if ws: - self._broker.del_ws_owner(ws_id) - event = WorkstreamClosedEvent( - ws_id=ws_id, - correlation_id=msg.correlation_id, - ) - self._publish_global(event) - self._publish_cluster(event) - - async def _handle_list_ws(self, msg: InboundMessage) -> None: - ws_list = [ - {"id": ws.ws_id, "name": ws.name, "state": ws.state} - for ws in self._workstreams.values() - ] - self._publish_global( - WorkstreamListEvent( - correlation_id=msg.correlation_id, - workstreams=ws_list, - ), - ) - - async def _handle_health(self, msg: InboundMessage) -> None: - self._publish_global( - HealthResponseEvent( - correlation_id=msg.correlation_id, - data={ - "status": "ok", - "node_id": self.node_id, - "sim": True, - "workstreams": len(self._workstreams), - }, - ), - ) - - async def _handle_list_nodes(self, msg: InboundMessage) -> None: - nodes = self._broker.list_nodes() - self._publish_global( - NodeListEvent(correlation_id=msg.correlation_id, nodes=nodes), - ) - - # -- workstream lifecycle ------------------------------------------------ - - def _create_workstream( - self, - name: str = "", - correlation_id: str = "", - ) -> SimWorkstream: - ws_id = uuid.uuid4().hex[:8] - if not name: - name = f"sim-ws-{ws_id[:4]}" - ws = SimWorkstream(ws_id, name, self, self._engine, self._config) - self._workstreams[ws_id] = ws - self._broker.set_ws_owner(ws_id, self.node_id) - event = WorkstreamCreatedEvent( - ws_id=ws_id, - correlation_id=correlation_id, - name=name, - ) - self._publish_global(event) - # Also publish to cluster channel so the console discovers the ws. - # Include node_id (the console collector keys on it). - self._publish_cluster( - ClusterStateEvent( - ws_id=ws_id, - state="idle", - node_id=self.node_id, - ), - ) - # The cluster channel expects a ws_created with node_id for the - # collector's _on_cluster_event handler. - self._broker.publish_outbound( - f"{self._prefix}:events:cluster", - json.dumps( - { - "type": "ws_created", - "ws_id": ws_id, - "name": name, - "node_id": self.node_id, - "correlation_id": correlation_id, - } - ), - ) - return ws - - # -- heartbeat ----------------------------------------------------------- - - def heartbeat_once(self) -> None: - """Register a single heartbeat with the broker.""" - self._broker.register_node( - self.node_id, - { - "server_url": f"sim://{self.node_id}", - "started": self._started_at, - "sim": True, - "workstreams": len(self._workstreams), - "max_ws": self._config.max_ws_per_node, - }, - ttl=self._config.heartbeat_ttl, - ) - - # -- shutdown ------------------------------------------------------------ - - def stop(self) -> None: - """Mark node as stopped and clean up ownership keys.""" - self._running = False - for ws_id in list(self._workstreams): - self._broker.del_ws_owner(ws_id) - self._workstreams.clear() - - # -- event publishing helpers -------------------------------------------- - - def _publish_global(self, event: OutboundEvent) -> None: - self._broker.publish_outbound( - f"{self._prefix}:events:global", - event.to_json(), - ) - - def _publish_ws(self, ws_id: str, event: OutboundEvent) -> None: - self._broker.publish_outbound( - f"{self._prefix}:events:{ws_id}", - event.to_json(), - ) - - def _publish_cluster(self, event: OutboundEvent) -> None: - self._broker.publish_outbound( - f"{self._prefix}:events:cluster", - event.to_json(), - ) diff --git a/turnstone/sim/scenario.py b/turnstone/sim/scenario.py deleted file mode 100644 index 122d164c..00000000 --- a/turnstone/sim/scenario.py +++ /dev/null @@ -1,226 +0,0 @@ -"""Simulation scenarios — workload patterns for cluster testing.""" - -from __future__ import annotations - -import asyncio -import logging -import time -from typing import TYPE_CHECKING, Any - -from turnstone.mq.broker import RedisBroker -from turnstone.mq.protocol import SendMessage - -if TYPE_CHECKING: - from turnstone.sim.cluster import SimCluster - from turnstone.sim.config import SimConfig - from turnstone.sim.metrics import MetricsCollector - -log = logging.getLogger("turnstone.sim.scenario") - - -class SteadyStateScenario: - """Inject messages at a constant rate for the configured duration.""" - - async def run( - self, - cluster: SimCluster, - config: SimConfig, - metrics: MetricsCollector, - ) -> None: - broker = _make_broker(config) - interval = 1.0 / max(0.01, config.messages_per_second) - deadline = time.monotonic() + config.duration - count = 0 - - try: - while time.monotonic() < deadline: - count += 1 - msg = SendMessage( - message=f"Steady-state message {count}", - auto_approve=True, - ) - broker.push_inbound(msg.to_json()) - metrics.record_inject() - await asyncio.sleep(interval) - finally: - # Allow in-flight turns to finish - await asyncio.sleep(min(10, config.llm_latency_mean * 3)) - broker.close() - log.info("Steady-state scenario complete: %d messages injected", count) - - -class BurstScenario: - """Inject burst_size messages as fast as possible, then wait.""" - - async def run( - self, - cluster: SimCluster, - config: SimConfig, - metrics: MetricsCollector, - ) -> None: - broker = _make_broker(config) - - try: - for i in range(config.burst_size): - msg = SendMessage( - message=f"Burst message {i}", - auto_approve=True, - ) - broker.push_inbound(msg.to_json()) - metrics.record_inject() - - log.info("Burst injected: %d messages", config.burst_size) - # Wait for processing to complete - await asyncio.sleep(config.duration) - finally: - broker.close() - - -class NodeFailureScenario: - """Steady-state load with periodic node kills.""" - - async def run( - self, - cluster: SimCluster, - config: SimConfig, - metrics: MetricsCollector, - ) -> None: - # Start steady injection in background - steady = SteadyStateScenario() - load_task = asyncio.create_task(steady.run(cluster, config, metrics)) - - # Periodically kill nodes - killed = 0 - max_kills = config.num_nodes // 2 # never kill more than half - node_ids = list(cluster.nodes.keys()) - - try: - while killed < max_kills: - await asyncio.sleep(config.node_kill_interval) - for _ in range(config.node_kill_count): - if killed < len(node_ids): - await cluster.kill_node(node_ids[killed]) - killed += 1 - finally: - await load_task - log.info("Node-failure scenario complete: %d nodes killed", killed) - - -class DirectedScenario: - """Send messages targeted to specific nodes.""" - - async def run( - self, - cluster: SimCluster, - config: SimConfig, - metrics: MetricsCollector, - ) -> None: - broker = _make_broker(config) - node_ids = list(cluster.nodes.keys()) - count = min(config.burst_size, len(node_ids)) - - try: - for i in range(count): - target = node_ids[i % len(node_ids)] - msg = SendMessage( - message=f"Directed message to {target}", - auto_approve=True, - target_node=target, - ) - broker.push_inbound(msg.to_json(), node_id=target) - metrics.record_inject() - - log.info("Directed scenario: %d messages sent to specific nodes", count) - await asyncio.sleep(config.duration) - finally: - broker.close() - - -class LifecycleScenario: - """Create, use, and close workstreams across nodes.""" - - async def run( - self, - cluster: SimCluster, - config: SimConfig, - metrics: MetricsCollector, - ) -> None: - from turnstone.mq.protocol import ( - CloseWorkstreamMessage, - CreateWorkstreamMessage, - InboundMessage, - ) - - broker = _make_broker(config) - ws_ids: list[str] = [] - - try: - # Phase 1: Create workstreams - create_count = min(50, config.num_nodes * 2) - for i in range(create_count): - msg: InboundMessage = CreateWorkstreamMessage( - name=f"lifecycle-ws-{i}", - auto_approve=True, - ) - broker.push_inbound(msg.to_json()) - metrics.record_inject() - await asyncio.sleep(0.05) - - # Let creations settle - await asyncio.sleep(3) - - # Phase 2: Send messages to shared queue (will be routed to nodes - # that own workstreams) - for i in range(create_count): - msg = SendMessage( - message=f"Lifecycle message {i}", - auto_approve=True, - ) - broker.push_inbound(msg.to_json()) - metrics.record_inject() - await asyncio.sleep(0.1) - - # Let turns complete - await asyncio.sleep(min(15, config.llm_latency_mean * 5)) - - # Phase 3: Close half the workstreams - # Collect ws_ids from nodes - for node in cluster.nodes.values(): - for ws_id in list(node._workstreams.keys()): - ws_ids.append(ws_id) - - close_count = len(ws_ids) // 2 - for ws_id in ws_ids[:close_count]: - owner = broker.get_ws_owner(ws_id) - msg = CloseWorkstreamMessage(ws_id=ws_id) - broker.push_inbound(msg.to_json(), node_id=owner or "") - await asyncio.sleep(0.05) - - await asyncio.sleep(2) - log.info( - "Lifecycle scenario complete: created %d, closed %d", - create_count, - close_count, - ) - finally: - broker.close() - - -def _make_broker(config: SimConfig) -> RedisBroker: - """Create a RedisBroker for scenario message injection.""" - return RedisBroker( - host=config.redis_host, - port=config.redis_port, - db=config.redis_db, - prefix=config.prefix, - password=config.redis_password, - ) - - -SCENARIOS: dict[str, type[Any]] = { - "steady": SteadyStateScenario, - "burst": BurstScenario, - "node_failure": NodeFailureScenario, - "directed": DirectedScenario, - "lifecycle": LifecycleScenario, -} diff --git a/uv.lock b/uv.lock index 3bfd9788..0f56b8a6 100644 --- a/uv.lock +++ b/uv.lock @@ -185,15 +185,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, ] -[[package]] -name = "async-timeout" -version = "5.0.1" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/a5/ae/136395dfbfe00dfc94da3f3e136d0b13f394cba8f4841120e34226265780/async_timeout-5.0.1.tar.gz", hash = "sha256:d9321a7a3d5a6a5e187e824d2fa0793ce379a202935782d555d6e9d2735677d3", size = 9274, upload-time = "2024-11-06T16:41:39.6Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/fe/ba/e2081de779ca30d473f21f5b30e0e737c438205440784c7dfc81efc2b029/async_timeout-5.0.1-py3-none-any.whl", hash = "sha256:39e3809566ff85354557ec2398b55e096c8364bacac9405a7a1fa429e77fe76c", size = 6233, upload-time = "2024-11-06T16:41:37.9Z" }, -] - [[package]] name = "attrs" version = "26.1.0" @@ -2090,18 +2081,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" }, ] -[[package]] -name = "redis" -version = "7.4.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "async-timeout", marker = "python_full_version < '3.11.3'" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/7b/7f/3759b1d0d72b7c92f0d70ffd9dc962b7b7b5ee74e135f9d7d8ab06b8a318/redis-7.4.0.tar.gz", hash = "sha256:64a6ea7bf567ad43c964d2c30d82853f8df927c5c9017766c55a1d1ed95d18ad", size = 4943913, upload-time = "2026-03-24T09:14:37.53Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/74/3a/95deec7db1eb53979973ebd156f3369a72732208d1391cd2e5d127062a32/redis-7.4.0-py3-none-any.whl", hash = "sha256:a9c74a5c893a5ef8455a5adb793a31bb70feb821c86eccb62eebef5a19c429ec", size = 409772, upload-time = "2026-03-24T09:14:35.968Z" }, -] - [[package]] name = "referencing" version = "0.37.0" @@ -2535,7 +2514,6 @@ all = [ { name = "numpy" }, { name = "psycopg", extra = ["binary"] }, { name = "pytest" }, - { name = "redis" }, { name = "scipy" }, { name = "sympy" }, ] @@ -2544,7 +2522,6 @@ anthropic = [ ] console = [ { name = "croniter" }, - { name = "redis" }, ] ddg = [ { name = "ddgs" }, @@ -2552,14 +2529,9 @@ ddg = [ dev = [ { name = "mypy" }, { name = "ruff" }, - { name = "types-redis" }, ] discord = [ { name = "discord-py" }, - { name = "redis" }, -] -mq = [ - { name = "redis" }, ] postgres = [ { name = "psycopg", extra = ["binary"] }, @@ -2570,9 +2542,6 @@ sandbox = [ { name = "scipy" }, { name = "sympy" }, ] -sim = [ - { name = "redis" }, -] test = [ { name = "croniter" }, { name = "pytest" }, @@ -2605,10 +2574,6 @@ requires-dist = [ { name = "pytest", marker = "extra == 'test'", specifier = ">=9.0" }, { name = "pytest-cov", marker = "extra == 'test'", specifier = ">=6.0" }, { name = "python-frontmatter", specifier = ">=1.0" }, - { name = "redis", marker = "extra == 'console'", specifier = ">=7.2" }, - { name = "redis", marker = "extra == 'discord'", specifier = ">=7.2" }, - { name = "redis", marker = "extra == 'mq'", specifier = ">=7.2" }, - { name = "redis", marker = "extra == 'sim'", specifier = ">=7.2" }, { name = "ruff", marker = "extra == 'dev'", specifier = ">=0.9" }, { name = "scipy", marker = "extra == 'sandbox'", specifier = ">=1.14" }, { name = "sqlalchemy", specifier = ">=2.0" }, @@ -2616,58 +2581,10 @@ requires-dist = [ { name = "starlette", specifier = ">=0.45" }, { name = "structlog", specifier = ">=24.1" }, { name = "sympy", marker = "extra == 'sandbox'", specifier = ">=1.13" }, - { name = "turnstone", extras = ["mq", "console", "sim", "anthropic", "postgres", "discord", "ddg", "tls", "sandbox"], marker = "extra == 'all'" }, - { name = "types-redis", marker = "extra == 'dev'", specifier = ">=4.6" }, + { name = "turnstone", extras = ["console", "anthropic", "postgres", "discord", "ddg", "tls", "sandbox"], marker = "extra == 'all'" }, { name = "uvicorn", specifier = ">=0.34" }, ] -provides-extras = ["test", "dev", "mq", "console", "sim", "anthropic", "postgres", "ddg", "discord", "tls", "sandbox", "all"] - -[[package]] -name = "types-cffi" -version = "2.0.0.20260316" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "types-setuptools" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/07/4c/805b40b094eb3fd60f8d17fa7b3c58a33781311a95d0e6a74da0751ce294/types_cffi-2.0.0.20260316.tar.gz", hash = "sha256:8fb06ed4709675c999853689941133affcd2250cd6121cc11fd22c0d81ad510c", size = 17399, upload-time = "2026-03-16T07:54:43.059Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/81/5e/9f1a709225ad9d0e1d7a6e4366ff285f0113c749e882d6cbeb40eab32e75/types_cffi-2.0.0.20260316-py3-none-any.whl", hash = "sha256:dd504698029db4c580385f679324621cc64d886e6a23e9821d52bc5169251302", size = 20096, upload-time = "2026-03-16T07:54:41.994Z" }, -] - -[[package]] -name = "types-pyopenssl" -version = "24.1.0.20240722" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "cryptography" }, - { name = "types-cffi" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/93/29/47a346550fd2020dac9a7a6d033ea03fccb92fa47c726056618cc889745e/types-pyOpenSSL-24.1.0.20240722.tar.gz", hash = "sha256:47913b4678a01d879f503a12044468221ed8576263c1540dcb0484ca21b08c39", size = 8458, upload-time = "2024-07-22T02:32:22.558Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/98/05/c868a850b6fbb79c26f5f299b768ee0adc1f9816d3461dcf4287916f655b/types_pyOpenSSL-24.1.0.20240722-py3-none-any.whl", hash = "sha256:6a7a5d2ec042537934cfb4c9d4deb0e16c4c6250b09358df1f083682fe6fda54", size = 7499, upload-time = "2024-07-22T02:32:21.232Z" }, -] - -[[package]] -name = "types-redis" -version = "4.6.0.20241004" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "cryptography" }, - { name = "types-pyopenssl" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/3a/95/c054d3ac940e8bac4ca216470c80c26688a0e79e09f520a942bb27da3386/types-redis-4.6.0.20241004.tar.gz", hash = "sha256:5f17d2b3f9091ab75384153bfa276619ffa1cf6a38da60e10d5e6749cc5b902e", size = 49679, upload-time = "2024-10-04T02:43:59.224Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/55/82/7d25dce10aad92d2226b269bce2f85cfd843b4477cd50245d7d40ecf8f89/types_redis-4.6.0.20241004-py3-none-any.whl", hash = "sha256:ef5da68cb827e5f606c8f9c0b49eeee4c2669d6d97122f301d3a55dc6a63f6ed", size = 58737, upload-time = "2024-10-04T02:43:57.968Z" }, -] - -[[package]] -name = "types-setuptools" -version = "82.0.0.20260210" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/4b/90/796ac8c774a7f535084aacbaa6b7053d16fff5c630eff87c3ecff7896c37/types_setuptools-82.0.0.20260210.tar.gz", hash = "sha256:d9719fbbeb185254480ade1f25327c4654f8c00efda3fec36823379cebcdee58", size = 44768, upload-time = "2026-02-10T04:22:02.107Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/3e/54/3489432b1d9bc713c9d8aa810296b8f5b0088403662959fb63a8acdbd4fc/types_setuptools-82.0.0.20260210-py3-none-any.whl", hash = "sha256:5124a7daf67f195c6054e0f00f1d97c69caad12fdcf9113eba33eff0bce8cd2b", size = 68433, upload-time = "2026-02-10T04:22:00.876Z" }, -] +provides-extras = ["test", "dev", "console", "anthropic", "postgres", "ddg", "discord", "tls", "sandbox", "all"] [[package]] name = "typing-extensions"